charon_lib/transform/normalize/
skip_trait_refs_when_known.rs1use rustc_hash::FxHashMap as HashMap;
2
3use derive_generic_visitor::*;
4
5use crate::transform::ctx::UllbcPass;
6use crate::{transform::TransformCtx, ullbc_ast::*};
7
8#[derive(Visitor)]
9struct NormalizeFnPtr<'a> {
10 ctx: &'a TransformCtx,
11 visited_tys: HashMap<Ty, Ty>,
14}
15
16impl VisitAstMut for NormalizeFnPtr<'_> {
17 fn enter_fn_ptr(&mut self, fn_ptr: &mut FnPtr) {
18 if let Some(new_fn_ptr) = normalize_default_method_call_on_known_impl(self.ctx, fn_ptr)
19 .or_else(|| normalize_method_call_on_known_impl(self.ctx, fn_ptr))
20 {
21 *fn_ptr = new_fn_ptr;
22 }
23 }
24 fn visit_ty(&mut self, ty: &mut Ty) -> ControlFlow<Self::Break> {
25 if let Some(new_ty) = self.visited_tys.get(ty) {
26 *ty = new_ty.clone();
27 return Continue(());
28 }
29 let old_ty = ty.clone();
30 self.visit_inner(ty)?;
31 self.visited_tys.insert(old_ty, ty.clone());
32 Continue(())
33 }
34}
35
36fn normalize_default_method_call_on_known_impl(
38 ctx: &TransformCtx,
39 fn_ptr: &FnPtr,
40) -> Option<FnPtr> {
41 let fun_id = fn_ptr.kind.as_ref().as_fun()?;
42 let fun_decl = ctx.translated.fun_decls.get(*fun_id)?;
43 let FunSource::TraitDefault {
44 trait_ref,
45 item_id: method_id,
46 } = &fun_decl.src
47 else {
48 return None;
49 };
50 let impl_ref = fn_ptr
52 .generics
53 .trait_refs
54 .get(TraitClauseId::ZERO)
55 .as_ref()?
56 .kind
57 .as_trait_impl()?;
58 let method_generics = {
59 let generics = &fn_ptr.generics;
60 let trait_generics = trait_ref.generics.as_ref();
61 GenericArgs {
62 regions: generics
63 .regions
64 .clone()
65 .split_off(trait_generics.regions.len()),
66 types: generics.types.clone().split_off(trait_generics.types.len()),
67 const_generics: generics
68 .const_generics
69 .clone()
70 .split_off(trait_generics.const_generics.len()),
71 trait_refs: generics
73 .trait_refs
74 .clone()
75 .split_off(trait_generics.trait_refs.len() + 1),
76 }
77 };
78 normalize_method_call(ctx, impl_ref, method_id, &method_generics)
79}
80
81fn normalize_method_call_on_known_impl(ctx: &TransformCtx, fn_ptr: &FnPtr) -> Option<FnPtr> {
83 let FnPtrKind::Trait(trait_ref, method_id) = fn_ptr.kind.as_ref() else {
84 return None;
85 };
86 let TraitRefKind::TraitImpl(impl_ref) = &trait_ref.kind else {
87 return None;
88 };
89 normalize_method_call(ctx, impl_ref, method_id, &fn_ptr.generics)
90}
91
92fn normalize_method_call(
93 ctx: &TransformCtx,
94 impl_ref: &TraitImplRef,
95 method_id: &TraitMethodId,
96 method_generics: &GenericArgs,
97) -> Option<FnPtr> {
98 let trait_impl = &ctx.translated.trait_impls.get(impl_ref.id)?;
99 let bound_fn = trait_impl.methods.get(*method_id)?;
101 if !method_generics.matches(&bound_fn.params) {
102 return None;
103 }
104
105 let fn_ref: Binder<Binder<FunDeclRef>> = Binder::new(
108 BinderKind::Other,
109 trait_impl.generics.clone(),
110 bound_fn.clone(),
111 );
112 let fn_ref = fn_ref.apply(&impl_ref.generics).apply(method_generics);
114 Some(FnPtr::new(FnPtrKind::Fun(fn_ref.id), fn_ref.generics))
115}
116
117pub struct Transform;
118impl UllbcPass for Transform {
119 fn transform_item(&self, ctx: &mut TransformCtx, mut item: ItemRefMut<'_>) {
120 let _ = item.drive_mut(&mut NormalizeFnPtr {
121 ctx,
122 visited_tys: Default::default(),
123 });
124 }
125}