Skip to main content

charon_lib/transform/normalize/
skip_trait_refs_when_known.rs

1use 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    /// Types are hash-consed and bodies mention the same types many times; remember what each
12    /// visited type was rewritten to instead of exploring it again.
13    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
36/// Transform `Trait::default_method<X>[impl_trait_for_X]` to a direct method call.
37fn 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    // If the first trait proof (for the self clause) is a known impl.
51    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            // The `+ 1` is for the self clause.
72            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
81/// Transform `impl_trait_for_X::method` to a direct method call.
82fn 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    // Find the function declaration corresponding to this impl.
100    let bound_fn = trait_impl.methods.get(*method_id)?;
101    if !method_generics.matches(&bound_fn.params) {
102        return None;
103    }
104
105    // Make the two levels of binding explicit: outer binder for the impl block, inner binder for
106    // the method.
107    let fn_ref: Binder<Binder<FunDeclRef>> = Binder::new(
108        BinderKind::Other,
109        trait_impl.generics.clone(),
110        bound_fn.clone(),
111    );
112    // Substitute the appropriate generics into the function call.
113    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}