charon_lib/transform/normalize/
partial_monomorphization.rs1use std::collections::{HashMap, HashSet, VecDeque};
14use std::fmt::Display;
15use std::mem;
16
17use derive_generic_visitor::Visitor;
18use index_vec::Idx;
19
20use crate::ast::visitor::{VisitWithBinderDepth, VisitorWithBinderDepth};
21use crate::formatter::IntoFormatter;
22use crate::options::MonomorphizeMut;
23use crate::pretty::FmtWithCtx;
24use crate::register_error;
25use crate::transform::ctx::TransformPass;
26use crate::{transform::TransformCtx, ullbc_ast::*};
27
28type MutabilityShape = Binder<GenericArgs>;
29
30#[derive(Visitor)]
32struct MutabilityShapeBuilder<'pm, 'ctx> {
33 pm: &'pm PartialMonomorphizer<'ctx>,
34 params: GenericParams,
36 extracted: GenericArgs,
38 binder_depth: DeBruijnId,
40}
41
42impl<'pm, 'ctx> MutabilityShapeBuilder<'pm, 'ctx> {
43 fn compute_shape(
64 pm: &'pm PartialMonomorphizer<'ctx>,
65 target_params: &GenericParams,
66 args: &GenericArgs,
67 ) -> (MutabilityShape, GenericArgs) {
68 let mut shape_contents = args.clone();
73 let mut builder = Self {
74 pm,
75 params: GenericParams {
76 regions: IndexVec::new(),
77 types: IndexVec::new(),
78 const_generics: IndexVec::new(),
79 ..target_params.clone()
80 },
81 extracted: GenericArgs {
82 regions: IndexVec::new(),
83 types: IndexVec::new(),
84 const_generics: IndexVec::new(),
85 trait_refs: mem::take(&mut shape_contents.trait_refs),
86 },
87 binder_depth: DeBruijnId::zero(),
88 };
89
90 let _ = VisitWithBinderDepth::new(&mut builder).visit(&mut shape_contents);
93
94 let shape_params = {
95 let mut shape_params = builder.params;
96 shape_params.trait_clauses = shape_params.trait_clauses.map_indexed(|i, x| {
100 if i.index() < target_params.trait_clauses.len() {
101 x.substitute_explicits(&shape_contents)
102 } else {
103 x
104 }
105 });
106 shape_params.trait_type_constraints =
107 shape_params.trait_type_constraints.map_indexed(|i, x| {
108 if i.index() < target_params.trait_type_constraints.len() {
109 x.substitute_explicits(&shape_contents)
110 } else {
111 x
112 }
113 });
114 shape_params.regions_outlive = shape_params
115 .regions_outlive
116 .into_iter()
117 .enumerate()
118 .map(|(i, x)| {
119 if i < target_params.regions_outlive.len() {
120 x.substitute_explicits(&shape_contents)
121 } else {
122 x
123 }
124 })
125 .collect();
126 shape_params.types_outlive = shape_params
127 .types_outlive
128 .into_iter()
129 .enumerate()
130 .map(|(i, x)| {
131 if i < target_params.types_outlive.len() {
132 x.substitute_explicits(&shape_contents)
133 } else {
134 x
135 }
136 })
137 .collect();
138 shape_params
139 };
140
141 shape_contents.trait_refs = shape_params.identity_args().trait_refs;
144 shape_contents
145 .trait_refs
146 .truncate(target_params.trait_clauses.len());
147
148 let shape_args = builder.extracted;
149 let shape = Binder::new(BinderKind::Other, shape_params, shape_contents);
150 (shape, shape_args)
151 }
152
153 fn replace_with_fresh_var<Id, Param, Arg>(
155 &mut self,
156 val: &mut Arg,
157 mk_param: impl FnOnce(Id) -> Param,
158 mk_value: impl FnOnce(DeBruijnVar<Id>) -> Arg,
159 ) where
160 Id: Idx + Display,
161 Arg: TyVisitable + Clone,
162 GenericParams: HasIdxVecOf<Id, Output = Param>,
163 GenericArgs: HasIdxVecOf<Id, Output = Arg>,
164 {
165 let Some(shifted_val) = val.clone().move_from_under_binders(self.binder_depth) else {
166 return;
168 };
169 self.extracted.get_idx_vec_mut().push(shifted_val);
171 let id = self.params.get_idx_vec_mut().push_with(mk_param);
173 *val = mk_value(DeBruijnVar::bound(self.binder_depth, id));
174 }
175}
176
177impl<'pm, 'ctx> VisitorWithBinderDepth for MutabilityShapeBuilder<'pm, 'ctx> {
178 fn binder_depth_mut(&mut self) -> &mut DeBruijnId {
179 &mut self.binder_depth
180 }
181}
182
183impl<'pm, 'ctx> VisitAstMut for MutabilityShapeBuilder<'pm, 'ctx> {
184 fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
185 VisitWithBinderDepth::new(self).visit(x)
186 }
187
188 fn enter_ty(&mut self, ty: &mut Ty) {
189 if !self.pm.is_infected(ty) {
190 self.replace_with_fresh_var(
191 ty,
192 |id| TypeParam::new(id, format!("T{id}"), Variance::Unknown),
193 |v| v.into(),
194 );
195 }
196 }
197 fn exit_ty_kind(&mut self, kind: &mut TyKind) {
198 if let TyKind::Adt(TypeDeclRef {
199 id: TypeId::Adt(id),
200 generics,
201 }) = kind
202 {
203 let Some(target_params) = self.pm.generic_params.get(&(*id).into()) else {
208 return;
209 };
210 let Some(shifted_generics) =
211 generics.clone().move_from_under_binders(self.binder_depth)
212 else {
213 return;
215 };
216
217 let num_clauses_before_merge = self.params.trait_clauses.len();
219 self.params.merge_predicates_from(
220 target_params
221 .clone()
222 .substitute_explicits(&shifted_generics),
223 );
224
225 self.extracted
227 .trait_refs
228 .extend(shifted_generics.trait_refs);
229
230 for (target_clause_id, tref) in generics.trait_refs.iter_mut_enumerated() {
232 let clause_id = target_clause_id + num_clauses_before_merge;
233 *tref =
234 self.params.trait_clauses[clause_id].identity_tref_at_depth(self.binder_depth);
235 }
236 }
237 }
238 fn enter_region(&mut self, r: &mut Region) {
239 self.replace_with_fresh_var(
240 r,
241 |id| RegionParam::new(id, None, Variance::Unknown),
242 |v| v.into(),
243 );
244 }
245 fn visit_trait_ref(&mut self, _tref: &mut TraitRef) -> ControlFlow<Self::Break> {
252 ControlFlow::Continue(())
255 }
256
257 fn visit_constant_expr(
258 &mut self,
259 _: &mut ConstantExpr,
260 ) -> ::std::ops::ControlFlow<Self::Break> {
261 ControlFlow::Continue(())
262 }
263}
264
265#[derive(Visitor)]
266struct PartialMonomorphizer<'a> {
267 ctx: &'a mut TransformCtx,
268 span: Span,
270 specialize_adts: bool,
272 infected_types: HashSet<TypeDeclId>,
274 generic_params: HashMap<ItemId, GenericParams>,
279 partial_mono_shapes: SeqHashMap<(ItemId, MutabilityShape), ItemId>,
283 reverse_shape_map: HashMap<ItemId, (ItemId, MutabilityShape)>,
285 to_process: VecDeque<ItemId>,
287}
288
289impl<'a> PartialMonomorphizer<'a> {
290 pub fn new(ctx: &'a mut TransformCtx, specialize_adts: bool) -> Self {
291 let infected_types: HashSet<_> = ctx
294 .translated
295 .type_decls
296 .iter()
297 .filter(|tdecl| {
298 tdecl
299 .generics
300 .regions
301 .iter()
302 .any(|r| r.mutability.is_mutable())
303 })
304 .map(|tdecl| tdecl.def_id)
305 .collect();
306
307 let generic_params: HashMap<ItemId, GenericParams> = ctx
309 .translated
310 .all_items()
311 .map(|item| (item.id(), item.generic_params().clone()))
312 .collect();
313
314 let to_process = ctx.translated.all_ids().collect();
316 PartialMonomorphizer {
317 ctx,
318 span: Span::dummy(),
319 specialize_adts,
320 infected_types,
321 generic_params,
322 to_process,
323 partial_mono_shapes: SeqHashMap::default(),
324 reverse_shape_map: Default::default(),
325 }
326 }
327
328 fn is_infected(&self, ty: &Ty) -> bool {
331 match ty.kind() {
332 TyKind::Ref(_, _, RefKind::Mut) => true,
333 TyKind::Ref(_, ty, _)
334 | TyKind::RawPtr(ty, _)
335 | TyKind::Array(ty, _)
336 | TyKind::Pattern(ty, _)
337 | TyKind::Slice(ty) => self.is_infected(ty),
338 TyKind::Adt(tref) => match tref.as_adt() {
339 Some(id) => {
340 let ty_infected = self.infected_types.contains(&id);
341 let args_infected = if self.specialize_adts {
342 false
347 } else {
348 tref.generics.types.iter().any(|ty| self.is_infected(ty))
349 };
350 ty_infected || args_infected
351 }
352 None => {
353 tref.generics.types.iter().any(|ty| self.is_infected(ty))
356 }
357 },
358 TyKind::FnDef(..) | TyKind::FnPtr(..) => false,
362 TyKind::DynTrait(_) => {
363 register_error!(
364 self.ctx,
365 self.span,
366 "`dyn Trait` is unsupported with `--monomorphize-mut`"
367 );
368 false
369 }
370 TyKind::TypeVar(..)
371 | TyKind::Literal(..)
372 | TyKind::Never
373 | TyKind::TraitType(..)
374 | TyKind::PtrMetadata(..)
375 | TyKind::Error(_) => false,
376 }
377 }
378
379 fn process_generics(&mut self, id: ItemId, generics: &GenericArgs) -> Option<DeclRef<ItemId>> {
383 if !generics.types.iter().any(|ty| self.is_infected(ty)) {
384 return None;
385 }
386
387 let mut new_generics;
390 let (id, generics) = if let Some(&(base_id, ref shape)) = self.reverse_shape_map.get(&id) {
391 new_generics = shape.clone().apply(generics);
392 let _ = self.visit(&mut new_generics); (base_id, &new_generics)
394 } else {
395 (id, generics)
396 };
397
398 let item_params = self.generic_params.get(&id)?;
400 let (shape, shape_args) =
401 MutabilityShapeBuilder::compute_shape(self, item_params, generics);
402
403 let new_params = shape.params.clone();
405 let key: (ItemId, MutabilityShape) = (id, shape);
406 let new_id = *self
407 .partial_mono_shapes
408 .entry(key.clone())
409 .or_insert_with(|| {
410 let new_id = match id {
411 ItemId::Type(_) => {
412 let new_id = self.ctx.translated.type_decls.reserve_slot();
413 self.infected_types.insert(new_id);
414 new_id.into()
415 }
416 ItemId::Fun(_) => self.ctx.translated.fun_decls.reserve_slot().into(),
417 ItemId::Global(_) => self.ctx.translated.global_decls.reserve_slot().into(),
418 ItemId::TraitDecl(_) => self.ctx.translated.trait_decls.reserve_slot().into(),
419 ItemId::TraitImpl(_) => self.ctx.translated.trait_impls.reserve_slot().into(),
420 };
421 self.generic_params.insert(new_id, new_params);
422 self.reverse_shape_map.insert(new_id, key);
423 self.to_process.push_back(new_id);
424 new_id
425 });
426
427 let fmt_ctx = self.ctx.into_fmt();
428 trace!(
429 "processing {}{}\n output: {}{}",
430 id.with_ctx(&fmt_ctx),
431 generics.with_ctx(&fmt_ctx),
432 new_id.with_ctx(&fmt_ctx),
433 shape_args.with_ctx(&fmt_ctx),
434 );
435 Some(DeclRef {
436 id: new_id,
437 generics: Box::new(shape_args),
438 trait_ref: None,
439 })
440 }
441
442 pub fn process_item(&mut self, item: &mut ItemRefMut<'_>) {
446 let _ = item.drive_mut(self);
447 }
448
449 pub fn create_pending_instantiation(&mut self, new_id: ItemId) -> ItemByVal {
454 let (orig_id, shape) = &self.reverse_shape_map[&new_id];
455 let mut decl = self
456 .ctx
457 .translated
458 .get_item(*orig_id)
459 .unwrap()
460 .to_owned()
461 .substitute_with_self(&shape.skip_binder, &TraitRefKind::SelfId);
462
463 let mut decl_mut = decl.as_mut();
464 decl_mut.set_id(new_id);
465 *decl_mut.generic_params() = shape.params.clone();
466
467 let name_ref = &mut decl_mut.item_meta().name;
468 *name_ref = mem::take::<crate::ast::Name>(name_ref).instantiate(shape.clone());
469 self.ctx
470 .translated
471 .item_names
472 .insert(new_id, decl.as_ref().item_meta().name.clone());
473 if let (ItemId::TraitDecl(orig_trait_id), ItemId::TraitDecl(new_trait_id)) =
474 (*orig_id, new_id)
475 {
476 let names = self.ctx.translated.assoc_item_names[orig_trait_id].clone();
477 self.ctx
478 .translated
479 .assoc_item_names
480 .insert(new_trait_id, names);
481 }
482
483 decl
484 }
485}
486
487impl VisitorWithSpan for PartialMonomorphizer<'_> {
488 fn current_span(&mut self) -> &mut Span {
489 &mut self.span
490 }
491}
492impl VisitAstMut for PartialMonomorphizer<'_> {
493 fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
494 VisitWithSpan::new(self).visit(x)
496 }
497
498 fn exit_type_decl_ref(&mut self, x: &mut TypeDeclRef) {
499 if self.specialize_adts
500 && let Some(id) = x.as_adt()
501 && let Some(new_decl_ref) = self.process_generics(id.into(), &x.generics)
502 {
503 *x = new_decl_ref.try_into().unwrap()
504 }
505 }
506 fn exit_fn_ptr(&mut self, x: &mut FnPtr) {
507 if let FnPtrKind::Fun(FunId::Regular(id)) = *x.kind
510 && let Some(new_decl_ref) = self.process_generics(id.into(), &x.generics)
511 {
512 *x = new_decl_ref.try_into().unwrap()
513 }
514 }
515 fn exit_fun_decl_ref(&mut self, x: &mut FunDeclRef) {
516 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
517 *x = new_decl_ref.try_into().unwrap()
518 }
519 }
520 fn exit_global_decl_ref(&mut self, x: &mut GlobalDeclRef) {
521 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
522 *x = new_decl_ref.try_into().unwrap()
523 }
524 }
525 fn exit_trait_decl_ref(&mut self, x: &mut TraitDeclRef) {
526 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
527 *x = new_decl_ref.try_into().unwrap()
528 }
529 }
530 fn exit_trait_impl_ref(&mut self, x: &mut TraitImplRef) {
531 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
532 *x = new_decl_ref.try_into().unwrap()
533 }
534 }
535}
536
537pub struct Transform;
538impl TransformPass for Transform {
539 fn transform_ctx(&self, ctx: &mut TransformCtx) {
540 let Some(include_types) = ctx.options.monomorphize_mut else {
541 return;
542 };
543 let mut visitor =
545 PartialMonomorphizer::new(ctx, matches!(include_types, MonomorphizeMut::All));
546 while let Some(id) = visitor.to_process.pop_front() {
547 let mut decl = if visitor.reverse_shape_map.contains_key(&id) {
550 visitor.create_pending_instantiation(id)
552 } else {
553 match visitor.ctx.translated.remove_item_temporarily(id) {
556 Some(decl) => decl,
557 None => continue,
558 }
559 };
560 visitor.process_item(&mut decl.as_mut());
563 visitor.ctx.translated.put_item_back(id, decl);
565 }
566 }
567}