Skip to main content

charon_lib/transform/add_missing_info/
add_missing_alias_clauses.rs

1//! Rust doesn't require bounds on type aliases to be well-formed. When a type alias mentions
2//! `<T as Trait>::Assoc` without a corresponding `T: Trait` clause, translation leaves an unknown
3//! trait ref. This pass tries to add these missing clauses.
4
5use crate::ast::*;
6use crate::transform::{TransformCtx, ctx::TransformPass};
7use rustc_hash::FxHashMap as HashMap;
8
9#[derive(Visitor)]
10struct ClauseExtractor<'a> {
11    params: &'a mut GenericParams,
12    span: Span,
13    binder_stack: BindingStack<GenericParams>,
14    extracted_clauses: HashMap<PolyTraitDeclRef, TraitClauseId>,
15}
16
17impl<'a> ClauseExtractor<'a> {
18    fn new(params: &'a mut GenericParams, span: Span) -> Self {
19        Self {
20            binder_stack: BindingStack::new(params.clone()),
21            params,
22            span,
23            extracted_clauses: HashMap::default(),
24        }
25    }
26
27    /// Move a trait ref out of the binders to make it a trait clause. Collects all the region
28    /// binders on the way to here into a single binder to make a HRTB.
29    fn extract_trait_clause(&self, mut trait_: PolyTraitDeclRef) -> Option<PolyTraitDeclRef> {
30        // Iterate over the binders on the way to this trait ref, skipping the first binder (the
31        // item binder).
32        let mut scope_regions = Vec::new();
33        for (dbid, params) in self.binder_stack.iter_enumerated().rev().skip(1) {
34            for (old_id, region) in params.regions.iter_enumerated() {
35                let new_id = trait_.regions.push_with(|index| {
36                    let mut region = region.clone();
37                    region.index = index;
38                    region
39                });
40                scope_regions.push((dbid, old_id, new_id));
41            }
42        }
43
44        if !scope_regions.is_empty() {
45            // Make all the region variables point at the outer binder.
46            #[derive(Visitor)]
47            struct MoveRegionsToHrtb {
48                binder_depth: DeBruijnId,
49                scope_regions: Vec<(DeBruijnId, RegionId, RegionId)>,
50            }
51
52            impl VisitorWithBinderDepth for MoveRegionsToHrtb {
53                fn binder_depth_mut(&mut self) -> &mut DeBruijnId {
54                    &mut self.binder_depth
55                }
56            }
57
58            impl VisitAstMut for MoveRegionsToHrtb {
59                fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
60                    VisitWithBinderDepth::new(self).visit(x)
61                }
62
63                fn enter_region(&mut self, region: &mut Region) {
64                    let Region::Var(var) = region else {
65                        return;
66                    };
67                    let DeBruijnVar::Bound(dbid, old_id) = *var else {
68                        return;
69                    };
70                    let Some(outer_depth) = dbid.sub(self.binder_depth.incr()) else {
71                        return;
72                    };
73                    let Some((_, _, new_id)) = self
74                        .scope_regions
75                        .iter()
76                        .find(|(dbid, id, _)| *dbid == outer_depth && *id == old_id)
77                    else {
78                        return;
79                    };
80                    *var = DeBruijnVar::bound(self.binder_depth, *new_id);
81                }
82            }
83
84            MoveRegionsToHrtb {
85                binder_depth: DeBruijnId::zero(),
86                scope_regions,
87            }
88            .visit(&mut trait_.skip_binder);
89        }
90
91        trait_.move_from_under_binders(self.binder_stack.depth())
92    }
93}
94
95impl VisitorWithBinderStack for ClauseExtractor<'_> {
96    fn binder_stack_mut(&mut self) -> &mut BindingStack<GenericParams> {
97        &mut self.binder_stack
98    }
99}
100
101impl VisitAstMut for ClauseExtractor<'_> {
102    fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
103        VisitWithBinderStack::new(self).visit(x)
104    }
105
106    fn exit_trait_ref_contents(&mut self, tref: &mut TraitRefContents) {
107        if matches!(tref.kind, TraitRefKind::Unknown(_))
108            && let Some(trait_) = self.extract_trait_clause(tref.trait_decl_ref.clone())
109        {
110            let clause_id = if let Some(clause_id) = self.extracted_clauses.get(&trait_) {
111                *clause_id
112            } else {
113                let clause_id = self.params.trait_clauses.push_with(|clause_id| TraitParam {
114                    clause_id,
115                    span: Some(self.span),
116                    origin: PredicateOrigin::WhereClauseOnType,
117                    trait_: trait_.clone(),
118                });
119                self.extracted_clauses.insert(trait_, clause_id);
120                clause_id
121            };
122            tref.kind =
123                TraitRefKind::Clause(DeBruijnVar::bound(self.binder_stack.depth(), clause_id));
124        }
125    }
126}
127
128pub struct Transform;
129impl TransformPass for Transform {
130    fn transform_ctx(&self, ctx: &mut TransformCtx) {
131        for tdecl in &mut ctx.translated.type_decls {
132            if matches!(tdecl.kind, TypeDeclKind::Alias(_)) {
133                let mut extractor = ClauseExtractor::new(&mut tdecl.generics, tdecl.item_meta.span);
134                extractor.visit(&mut tdecl.kind);
135                extractor.visit(&mut tdecl.layout);
136            }
137        }
138    }
139}