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 { id, generics, .. }) = kind {
199 let Some(target_params) = self.pm.generic_params.get(&(*id).into()) else {
204 return;
205 };
206 let Some(shifted_generics) =
207 generics.clone().move_from_under_binders(self.binder_depth)
208 else {
209 return;
211 };
212
213 let num_clauses_before_merge = self.params.trait_clauses.len();
215 self.params.merge_predicates_from(
216 target_params
217 .clone()
218 .substitute_explicits(&shifted_generics),
219 );
220
221 self.extracted
223 .trait_refs
224 .extend(shifted_generics.trait_refs);
225
226 for (target_clause_id, tref) in generics.trait_refs.iter_mut_enumerated() {
228 let clause_id = target_clause_id + num_clauses_before_merge;
229 *tref =
230 self.params.trait_clauses[clause_id].identity_tref_at_depth(self.binder_depth);
231 }
232 }
233 }
234 fn enter_region(&mut self, r: &mut Region) {
235 self.replace_with_fresh_var(
236 r,
237 |id| RegionParam::new(id, None, Variance::Unknown),
238 |v| v.into(),
239 );
240 }
241 fn visit_trait_ref(&mut self, _tref: &mut TraitRef) -> ControlFlow<Self::Break> {
248 ControlFlow::Continue(())
251 }
252
253 fn visit_constant_expr(
254 &mut self,
255 _: &mut ConstantExpr,
256 ) -> ::std::ops::ControlFlow<Self::Break> {
257 ControlFlow::Continue(())
258 }
259}
260
261#[derive(Visitor)]
262struct PartialMonomorphizer<'a> {
263 ctx: &'a mut TransformCtx,
264 span: Span,
266 specialize_adts: bool,
268 infected_types: HashSet<TypeDeclId>,
270 generic_params: HashMap<ItemId, GenericParams>,
275 partial_mono_shapes: SeqHashMap<(ItemId, MutabilityShape), ItemId>,
279 reverse_shape_map: HashMap<ItemId, (ItemId, MutabilityShape)>,
281 to_process: VecDeque<ItemId>,
283}
284
285impl<'a> PartialMonomorphizer<'a> {
286 pub fn new(ctx: &'a mut TransformCtx, specialize_adts: bool) -> Self {
287 let infected_types: HashSet<_> = ctx
290 .translated
291 .type_decls
292 .iter()
293 .filter(|tdecl| {
294 tdecl
295 .generics
296 .regions
297 .iter()
298 .any(|r| r.mutability.is_mutable())
299 })
300 .map(|tdecl| tdecl.def_id)
301 .collect();
302
303 let generic_params: HashMap<ItemId, GenericParams> = ctx
305 .translated
306 .all_items()
307 .map(|item| (item.id(), item.generic_params().clone()))
308 .collect();
309
310 let to_process = ctx.translated.all_ids().collect();
312 PartialMonomorphizer {
313 ctx,
314 span: Span::dummy(),
315 specialize_adts,
316 infected_types,
317 generic_params,
318 to_process,
319 partial_mono_shapes: SeqHashMap::default(),
320 reverse_shape_map: Default::default(),
321 }
322 }
323
324 fn is_infected(&self, ty: &Ty) -> bool {
327 match ty.kind() {
328 TyKind::Ref(_, _, RefKind::Mut) => true,
329 TyKind::Ref(_, ty, _)
330 | TyKind::RawPtr(ty, _)
331 | TyKind::Array(ty, ..)
332 | TyKind::Pattern(ty, _)
333 | TyKind::Slice(ty, _) => self.is_infected(ty),
334 TyKind::Adt(tref) => {
335 let ty_infected = self.infected_types.contains(&tref.id);
336 let args_infected = if self.specialize_adts {
337 false
342 } else {
343 tref.generics.types.iter().any(|ty| self.is_infected(ty))
344 };
345 ty_infected || args_infected
346 }
347 TyKind::FnDef(..) | TyKind::FnPtr(..) => false,
351 TyKind::DynTrait(_) => {
352 register_error!(
353 self.ctx,
354 self.span,
355 "`dyn Trait` is unsupported with `--monomorphize-mut`"
356 );
357 false
358 }
359 TyKind::TypeVar(..)
360 | TyKind::Scalar(..)
361 | TyKind::Never
362 | TyKind::TraitType(..)
363 | TyKind::PtrMetadata(..)
364 | TyKind::Error(_) => false,
365 }
366 }
367
368 fn process_generics(&mut self, id: ItemId, generics: &GenericArgs) -> Option<DeclRef<ItemId>> {
372 if !generics.types.iter().any(|ty| self.is_infected(ty)) {
373 return None;
374 }
375
376 let mut new_generics;
379 let (id, generics) = if let Some(&(base_id, ref shape)) = self.reverse_shape_map.get(&id) {
380 new_generics = shape.clone().apply(generics);
381 let _ = self.visit(&mut new_generics); (base_id, &new_generics)
383 } else {
384 (id, generics)
385 };
386
387 let item_params = self.generic_params.get(&id)?;
389 let (shape, shape_args) =
390 MutabilityShapeBuilder::compute_shape(self, item_params, generics);
391
392 let new_params = shape.params.clone();
394 let key: (ItemId, MutabilityShape) = (id, shape);
395 let new_id = *self
396 .partial_mono_shapes
397 .entry(key.clone())
398 .or_insert_with(|| {
399 let new_id = match id {
400 ItemId::Type(_) => {
401 let new_id = self.ctx.translated.type_decls.reserve_slot();
402 self.infected_types.insert(new_id);
403 new_id.into()
404 }
405 ItemId::Fun(_) => self.ctx.translated.fun_decls.reserve_slot().into(),
406 ItemId::Global(_) => self.ctx.translated.global_decls.reserve_slot().into(),
407 ItemId::TraitDecl(_) => self.ctx.translated.trait_decls.reserve_slot().into(),
408 ItemId::TraitImpl(_) => self.ctx.translated.trait_impls.reserve_slot().into(),
409 };
410 self.generic_params.insert(new_id, new_params);
411 self.reverse_shape_map.insert(new_id, key);
412 self.to_process.push_back(new_id);
413 new_id
414 });
415
416 let fmt_ctx = self.ctx.into_fmt();
417 trace!(
418 "processing {}{}\n output: {}{}",
419 id.with_ctx(&fmt_ctx),
420 generics.with_ctx(&fmt_ctx),
421 new_id.with_ctx(&fmt_ctx),
422 shape_args.with_ctx(&fmt_ctx),
423 );
424 Some(DeclRef {
425 id: new_id,
426 generics: Box::new(shape_args),
427 trait_ref: None,
428 })
429 }
430
431 pub fn process_item(&mut self, item: &mut ItemRefMut<'_>) {
435 let _ = item.drive_mut(self);
436 }
437
438 pub fn create_pending_instantiation(&mut self, new_id: ItemId) -> ItemByVal {
443 let (orig_id, shape) = &self.reverse_shape_map[&new_id];
444 let mut decl = self
445 .ctx
446 .translated
447 .get_item(*orig_id)
448 .unwrap()
449 .to_owned()
450 .substitute_with_self(&shape.skip_binder, &TraitRefKind::SelfId);
451
452 let mut decl_mut = decl.as_mut();
453 decl_mut.set_id(new_id);
454 *decl_mut.generic_params() = shape.params.clone();
455
456 let name_ref = &mut decl_mut.item_meta().name;
457 *name_ref = mem::take::<crate::ast::Name>(name_ref).instantiate(shape.clone());
458 self.ctx
459 .translated
460 .item_names
461 .insert(new_id, decl.as_ref().item_meta().name.clone());
462 if let (ItemId::TraitDecl(orig_trait_id), ItemId::TraitDecl(new_trait_id)) =
463 (*orig_id, new_id)
464 {
465 let names = self.ctx.translated.assoc_item_names[orig_trait_id].clone();
466 self.ctx
467 .translated
468 .assoc_item_names
469 .insert(new_trait_id, names);
470 }
471
472 decl
473 }
474}
475
476impl VisitorWithSpan for PartialMonomorphizer<'_> {
477 fn current_span(&mut self) -> &mut Span {
478 &mut self.span
479 }
480}
481impl VisitAstMut for PartialMonomorphizer<'_> {
482 fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
483 VisitWithSpan::new(self).visit(x)
485 }
486
487 fn exit_type_decl_ref(&mut self, x: &mut TypeDeclRef) {
488 if x.is_tuple() && self.ctx.options.no_gen_tuple_structs {
489 return;
490 }
491 if self.specialize_adts
492 && let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics)
493 {
494 x.id = new_decl_ref.id.try_into().unwrap();
495 x.generics = new_decl_ref.generics;
496 }
497 }
498 fn exit_fn_ptr(&mut self, x: &mut FnPtr) {
499 if let FnPtrKind::Fun(id) = *x.kind
502 && let Some(new_decl_ref) = self.process_generics(id.into(), &x.generics)
503 {
504 *x = new_decl_ref.try_into().unwrap()
505 }
506 }
507 fn exit_fun_decl_ref(&mut self, x: &mut FunDeclRef) {
508 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
509 *x = new_decl_ref.try_into().unwrap()
510 }
511 }
512 fn exit_global_decl_ref(&mut self, x: &mut GlobalDeclRef) {
513 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
514 *x = new_decl_ref.try_into().unwrap()
515 }
516 }
517 fn exit_trait_decl_ref(&mut self, x: &mut TraitDeclRef) {
518 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
519 *x = new_decl_ref.try_into().unwrap()
520 }
521 }
522 fn exit_trait_impl_ref(&mut self, x: &mut TraitImplRef) {
523 if let Some(new_decl_ref) = self.process_generics(x.id.into(), &x.generics) {
524 *x = new_decl_ref.try_into().unwrap()
525 }
526 }
527}
528
529pub struct Transform;
530impl TransformPass for Transform {
531 fn transform_ctx(&self, ctx: &mut TransformCtx) {
532 let Some(include_types) = ctx.options.monomorphize_mut else {
533 return;
534 };
535 let mut visitor =
537 PartialMonomorphizer::new(ctx, matches!(include_types, MonomorphizeMut::All));
538 while let Some(id) = visitor.to_process.pop_front() {
539 let mut decl = if visitor.reverse_shape_map.contains_key(&id) {
542 visitor.create_pending_instantiation(id)
544 } else {
545 match visitor.ctx.translated.remove_item_temporarily(id) {
548 Some(decl) => decl,
549 None => continue,
550 }
551 };
552 visitor.process_item(&mut decl.as_mut());
555 visitor.ctx.translated.put_item_back(id, decl);
557 }
558 }
559}