Skip to main content

charon_lib/ast/type_level/
substitute.rs

1use crate::ast::*;
2use derive_generic_visitor::*;
3use std::borrow::Cow;
4use std::convert::Infallible;
5use std::fmt::Debug;
6use std::iter::Iterator;
7
8/// Visitor for type-level variables. Used to visit the variables contained in a value, as seen
9/// from the outside of the value. This means that any variable bound inside the value will be
10/// skipped, and all the seen De Bruijn indices will count from the outside of the value. The
11/// returned value, if any, will be put in place of the variable.
12pub trait VarsVisitor {
13    fn visit_erased_region(&mut self) -> Option<Region> {
14        None
15    }
16    fn visit_region_var(&mut self, _v: RegionDbVar) -> Option<Region> {
17        None
18    }
19    fn visit_type_var(&mut self, _v: TypeDbVar) -> Option<Ty> {
20        None
21    }
22    fn visit_const_generic_var(&mut self, _v: ConstGenericDbVar) -> Option<ConstantExprKind> {
23        None
24    }
25    fn visit_clause_var(&mut self, _v: ClauseDbVar) -> Option<TraitRefKind> {
26        None
27    }
28    fn visit_self_clause(&mut self) -> Option<TraitRefKind> {
29        None
30    }
31    fn visit_metadata_value(&mut self, _value: &MetadataValue) {}
32}
33
34/// Visitor for the [TyVisitable::substitute] function.
35/// This substitutes variables bound at the level where we start to substitute (level 0).
36#[derive(Visitor)]
37pub(crate) struct SubstVisitor<'a> {
38    generics: &'a GenericArgs,
39    self_ref: Option<&'a TraitRefKind>,
40    /// Whether to substitute explicit variables only (types, regions, const generics).
41    explicits_only: bool,
42    had_error: bool,
43}
44impl<'a> SubstVisitor<'a> {
45    pub(crate) fn new(
46        generics: &'a GenericArgs,
47        self_ref: Option<&'a TraitRefKind>,
48        explicits_only: bool,
49    ) -> Self {
50        Self {
51            generics,
52            self_ref,
53            explicits_only,
54            had_error: false,
55        }
56    }
57
58    pub fn visit<T: TyVisitable>(mut self, mut x: T) -> Result<T, GenericsMismatch> {
59        if x.type_info().is_closed() {
60            return Ok(x);
61        }
62        x.visit_vars(&mut self);
63        if self.had_error {
64            Err(GenericsMismatch)
65        } else {
66            Ok(x)
67        }
68    }
69
70    /// Returns the value for this variable, if any.
71    fn process_var<Id, T>(
72        &mut self,
73        var: DeBruijnVar<Id>,
74        get: impl Fn(Id) -> Option<&'a T>,
75    ) -> Option<T>
76    where
77        Id: Copy,
78        T: Clone + TyVisitable,
79        DeBruijnVar<Id>: Into<T>,
80    {
81        match var {
82            DeBruijnVar::Bound(dbid, varid) => {
83                Some(if let Some(dbid) = dbid.sub(DeBruijnId::one()) {
84                    // This is bound outside the binder we're substituting for.
85                    DeBruijnVar::Bound(dbid, varid).into()
86                } else {
87                    match get(varid) {
88                        Some(v) => v.clone(),
89                        None => {
90                            self.had_error = true;
91                            return None;
92                        }
93                    }
94                })
95            }
96            DeBruijnVar::Free(..) => None,
97        }
98    }
99}
100impl VarsVisitor for SubstVisitor<'_> {
101    fn visit_region_var(&mut self, v: RegionDbVar) -> Option<Region> {
102        self.process_var(v, |id| self.generics.regions.get(id))
103    }
104    fn visit_type_var(&mut self, v: TypeDbVar) -> Option<Ty> {
105        self.process_var(v, |id| self.generics.types.get(id))
106    }
107    fn visit_const_generic_var(&mut self, v: ConstGenericDbVar) -> Option<ConstantExprKind> {
108        self.process_var(v, |id| {
109            self.generics.const_generics.get(id).map(|c| c.kind())
110        })
111    }
112    fn visit_clause_var(&mut self, v: ClauseDbVar) -> Option<TraitRefKind> {
113        if self.explicits_only {
114            None
115        } else {
116            self.process_var(v, |id| Some(&self.generics.trait_refs.get(id)?.kind))
117        }
118    }
119    fn visit_self_clause(&mut self) -> Option<TraitRefKind> {
120        Some(self.self_ref.cloned().expect(
121            "used `substitute` on an item coming from a trait; \
122            use `substitute_with_self` or `substitute_inner_binder` instead.",
123        ))
124    }
125    fn visit_metadata_value(&mut self, _value: &MetadataValue) {
126        self.had_error = true;
127    }
128}
129
130#[derive(Debug)]
131pub struct GenericsMismatch;
132
133/// Types that are involved at the type-level and may be substituted around.
134pub trait TyVisitable: Sized + AstVisitable {
135    /// Compute various bits of information about the contents of this value. See methods on
136    /// [`TypeInfo`].
137    fn type_info(&self) -> TypeInfo {
138        TypeInfo::compute(self)
139    }
140
141    /// Visit the variables contained in `self`, as seen from the outside of `self`. This means
142    /// that any variable bound inside `self` will be skipped, and all the seen De Bruijn indices
143    /// will count from the outside of `self`.
144    fn visit_vars(&mut self, v: &mut impl VarsVisitor) {
145        #[derive(Visitor)]
146        struct Wrap<'v, V> {
147            v: &'v mut V,
148            depth: DeBruijnId,
149        }
150        impl<V> VisitorWithBinderDepth for Wrap<'_, V> {
151            fn binder_depth_mut(&mut self) -> &mut DeBruijnId {
152                &mut self.depth
153            }
154        }
155        impl<V: VarsVisitor> VisitAstMut for Wrap<'_, V> {
156            fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
157                VisitWithBinderDepth::new(self).visit(x)
158            }
159
160            fn exit_region(&mut self, r: &mut Region) {
161                match r {
162                    Region::Var(var)
163                        if let Some(var) = var.move_out_from_depth(self.depth)
164                            && let Some(new_r) = self.v.visit_region_var(var) =>
165                    {
166                        *r = new_r.move_under_binders(self.depth);
167                    }
168                    Region::Erased | Region::Body(..)
169                        if let Some(new_r) = self.v.visit_erased_region() =>
170                    {
171                        *r = new_r.move_under_binders(self.depth);
172                    }
173                    _ => (),
174                }
175            }
176            fn exit_ty(&mut self, ty: &mut Ty) {
177                if let TyKind::TypeVar(var) = ty.kind()
178                    && let Some(var) = var.move_out_from_depth(self.depth)
179                    && let Some(new_ty) = self.v.visit_type_var(var)
180                {
181                    *ty = new_ty.move_under_binders(self.depth);
182                }
183            }
184            fn exit_constant_expr_kind(&mut self, kind: &mut ConstantExprKind) {
185                if let ConstantExprKind::Var(var) = kind
186                    && let Some(var) = var.move_out_from_depth(self.depth)
187                    && let Some(new_cg) = self.v.visit_const_generic_var(var)
188                {
189                    *kind = new_cg.move_under_binders(self.depth);
190                }
191            }
192            fn exit_trait_ref_kind(&mut self, kind: &mut TraitRefKind) {
193                match kind {
194                    TraitRefKind::SelfId => {
195                        if let Some(new_kind) = self.v.visit_self_clause() {
196                            *kind = new_kind.move_under_binders(self.depth);
197                        }
198                    }
199                    TraitRefKind::Clause(var) => {
200                        if let Some(var) = var.move_out_from_depth(self.depth)
201                            && let Some(new_kind) = self.v.visit_clause_var(var)
202                        {
203                            *kind = new_kind.move_under_binders(self.depth);
204                        }
205                    }
206                    _ => {}
207                }
208            }
209            fn enter_metadata_value(&mut self, value: &mut MetadataValue) {
210                self.v.visit_metadata_value(value);
211            }
212        }
213        Wrap {
214            v,
215            depth: DeBruijnId::zero(),
216        }
217        .visit(self);
218    }
219
220    /// Substitute the generic variables inside `self` by replacing them with the provided values.
221    /// Note: if `self` is an item that comes from a `TraitDecl`, you must use
222    /// `substitute_with_self` or `substitute_inner_binder`, otherwise you'll get panics.
223    fn substitute(self, generics: &GenericArgs) -> Self {
224        SubstVisitor::new(generics, None, false)
225            .visit(self)
226            .unwrap()
227    }
228    /// Substitute the generic variables inside `self` by replacing them with the provided values.
229    /// This is appropriate when substituting an inner binder.
230    fn substitute_inner_binder(self, generics: &GenericArgs) -> Self {
231        self.substitute_with_self(generics, &TraitRefKind::SelfId)
232    }
233    /// Substitute only the type, region and const generic args.
234    fn substitute_explicits(self, generics: &GenericArgs) -> Self {
235        SubstVisitor::new(generics, None, true).visit(self).unwrap()
236    }
237    /// Substitute the generic variables as well as the `TraitRefKind::SelfId` trait ref.
238    fn substitute_with_self(self, generics: &GenericArgs, self_ref: &TraitRefKind) -> Self {
239        self.try_substitute_with_self(generics, self_ref).unwrap()
240    }
241    /// Substitute the generic variables as well as the `TraitRefKind::SelfId` trait ref.
242    fn substitute_with_tref(self, tref: &TraitRef) -> Self {
243        let pred = tref.trait_decl_ref.clone().erase();
244        self.substitute_with_self(&pred.generics, &tref.kind)
245    }
246    /// Substitute the generic variables as well as the `TraitRefKind::SelfId` trait ref.
247    fn try_substitute_with_tref(self, tref: &TraitRef) -> Result<Self, GenericsMismatch> {
248        let pred = tref.trait_decl_ref.clone().erase();
249        self.try_substitute_with_self(&pred.generics, &tref.kind)
250    }
251
252    fn try_substitute(self, generics: &GenericArgs) -> Result<Self, GenericsMismatch> {
253        SubstVisitor::new(generics, None, false).visit(self)
254    }
255    fn try_substitute_with_self(
256        self,
257        generics: &GenericArgs,
258        self_ref: &TraitRefKind,
259    ) -> Result<Self, GenericsMismatch> {
260        SubstVisitor::new(generics, Some(self_ref), false).visit(self)
261    }
262
263    /// Move under one binder.
264    fn move_under_binder(self) -> Self {
265        self.move_under_binders(DeBruijnId::one())
266    }
267
268    /// Move under `depth` binders.
269    fn move_under_binders(mut self, depth: DeBruijnId) -> Self {
270        if !depth.is_zero() {
271            let Continue(()) = self.visit_db_id::<Infallible>(|id| {
272                *id = id.plus(depth);
273                Continue(())
274            });
275        }
276        self
277    }
278
279    /// Move from under one binder.
280    fn move_from_under_binder(self) -> Option<Self> {
281        self.move_from_under_binders(DeBruijnId::one())
282    }
283
284    /// Move the value out of `depth` binders. Returns `None` if it contains a variable bound in
285    /// one of these `depth` binders.
286    fn move_from_under_binders(mut self, depth: DeBruijnId) -> Option<Self> {
287        match self.type_info().max_de_bruijn_id() {
288            None => return Some(self),
289            Some(max) if max < depth => return None,
290            Some(_) => {}
291        }
292        self.visit_db_id::<()>(|id| match id.sub(depth) {
293            Some(sub) => {
294                *id = sub;
295                Continue(())
296            }
297            None => Break(()),
298        })
299        .is_continue()
300        .then_some(self)
301    }
302
303    /// Visit the de Bruijn ids contained in `self`, as seen from the outside of `self`. This means
304    /// that any variable bound inside `self` will be skipped, and all the seen indices will count
305    /// from the outside of self.
306    fn visit_db_id<B>(
307        &mut self,
308        f: impl FnMut(&mut DeBruijnId) -> ControlFlow<B>,
309    ) -> ControlFlow<B> {
310        if self.type_info().max_de_bruijn_id().is_none() {
311            return Continue(());
312        }
313
314        struct Wrap<F> {
315            f: F,
316            depth: DeBruijnId,
317        }
318        impl<B, F> Visitor for Wrap<F>
319        where
320            F: FnMut(&mut DeBruijnId) -> ControlFlow<B>,
321        {
322            type Break = B;
323        }
324        impl<F> VisitorWithBinderDepth for Wrap<F> {
325            fn binder_depth_mut(&mut self) -> &mut DeBruijnId {
326                &mut self.depth
327            }
328        }
329        impl<B, F> VisitAstMut for Wrap<F>
330        where
331            F: FnMut(&mut DeBruijnId) -> ControlFlow<B>,
332        {
333            fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
334                VisitWithBinderDepth::new(self).visit(x)
335            }
336
337            fn visit_with_cached_type_info<T: AstVisitable>(
338                &mut self,
339                value: &mut WithCachedTypeInfo<T>,
340            ) -> ControlFlow<Self::Break> {
341                if value
342                    .type_info()
343                    .max_de_bruijn_id()
344                    .is_none_or(|max| max < self.depth)
345                {
346                    Continue(())
347                } else {
348                    self.visit_inner(value)
349                }
350            }
351
352            fn visit_de_bruijn_id(&mut self, x: &mut DeBruijnId) -> ControlFlow<Self::Break> {
353                if let Some(mut shifted) = x.sub(self.depth) {
354                    (self.f)(&mut shifted)?;
355                    *x = shifted.plus(self.depth)
356                }
357                Continue(())
358            }
359        }
360        Wrap {
361            f,
362            depth: DeBruijnId::zero(),
363        }
364        .visit(self)
365    }
366
367    /// Collect the regions contained in `self`.
368    fn collect_regions(&self) -> impl Iterator<Item = Region> {
369        let mut regions = SeqHashSet::new();
370        self.dyn_visit(|region: &Region| {
371            regions.insert(*region);
372        });
373        regions.into_iter()
374    }
375
376    /// Replace all the erased regions by the output of the provided function. Binders levels are
377    /// handled automatically.
378    fn replace_erased_regions(mut self, f: impl FnMut() -> Region) -> Self {
379        if !self.type_info().has_erased_or_body_regions() {
380            return self;
381        }
382
383        #[derive(Visitor)]
384        struct RefreshErasedRegions<F>(F);
385        impl<F: FnMut() -> Region> VarsVisitor for RefreshErasedRegions<F> {
386            fn visit_erased_region(&mut self) -> Option<Region> {
387                Some((self.0)())
388            }
389        }
390        self.visit_vars(&mut RefreshErasedRegions(f));
391        self
392    }
393}
394
395impl<T: AstVisitable> TyVisitable for T {}
396
397/// A value of type `T` applied to some `GenericArgs`, except we havent applied them yet to avoid a
398/// deep clone.
399#[derive(Debug, Clone)]
400pub struct Substituted<'a, T> {
401    pub val: &'a T,
402    pub generics: Cow<'a, GenericArgs>,
403    pub trait_self: Option<&'a TraitRefKind>,
404}
405
406impl<'a, T> Substituted<'a, T> {
407    pub fn new(val: &'a T, generics: &'a GenericArgs) -> Self {
408        Self {
409            val,
410            generics: Cow::Borrowed(generics),
411            trait_self: None,
412        }
413    }
414    pub fn new_for_trait(
415        val: &'a T,
416        generics: &'a GenericArgs,
417        trait_self: &'a TraitRefKind,
418    ) -> Self {
419        Self {
420            val,
421            generics: Cow::Borrowed(generics),
422            trait_self: Some(trait_self),
423        }
424    }
425    pub fn new_for_trait_ref(val: &'a T, tref: &'a TraitRef) -> Self {
426        Self {
427            val,
428            generics: Cow::Owned(*tref.trait_decl_ref.clone().erase().generics),
429            trait_self: Some(&tref.kind),
430        }
431    }
432
433    pub fn rebind<U>(&self, val: &'a U) -> Substituted<'a, U> {
434        Substituted {
435            val,
436            generics: self.generics.clone(),
437            trait_self: self.trait_self,
438        }
439    }
440
441    pub fn substitute(&self) -> T
442    where
443        T: TyVisitable + Clone,
444    {
445        self.try_substitute().unwrap()
446    }
447    pub fn try_substitute(&self) -> Result<T, GenericsMismatch>
448    where
449        T: TyVisitable + Clone,
450    {
451        match self.trait_self {
452            None => self.val.clone().try_substitute(&self.generics),
453            Some(trait_self) => self
454                .val
455                .clone()
456                .try_substitute_with_self(&self.generics, trait_self),
457        }
458    }
459
460    pub fn iter<Item: 'a>(&self) -> impl Iterator<Item = Substituted<'a, Item>>
461    where
462        &'a T: IntoIterator<Item = &'a Item>,
463    {
464        self.val.into_iter().map(move |x| self.rebind(x))
465    }
466}
467
468/// A value of type `T` bound by the generic parameters of item
469/// `item`. Used when dealing with multiple items at a time, to
470/// ensure we don't mix up generics.
471///
472/// To get the value, use `under_binder_of` or `subst_for`.
473#[derive(Debug, Copy, Clone)]
474pub struct ItemBinder<ItemId, T> {
475    pub item_id: ItemId,
476    val: T,
477}
478
479impl<ItemId, T> ItemBinder<ItemId, T>
480where
481    ItemId: Debug + Copy + PartialEq,
482{
483    pub fn new(item_id: ItemId, val: T) -> Self {
484        Self { item_id, val }
485    }
486
487    pub fn as_ref(&self) -> ItemBinder<ItemId, &T> {
488        ItemBinder {
489            item_id: self.item_id,
490            val: &self.val,
491        }
492    }
493
494    pub fn map_bound<U>(self, f: impl FnOnce(T) -> U) -> ItemBinder<ItemId, U> {
495        ItemBinder {
496            item_id: self.item_id,
497            val: f(self.val),
498        }
499    }
500
501    fn assert_item_id(&self, item_id: ItemId) {
502        assert_eq!(
503            self.item_id, item_id,
504            "Trying to use item bound for {:?} as if it belonged to {:?}",
505            self.item_id, item_id
506        );
507    }
508
509    /// Assert that the value is bound for item `item_id`, and returns it. This is used when we
510    /// plan to store the returned value inside that item.
511    pub fn under_binder_of(self, item_id: ItemId) -> T {
512        self.assert_item_id(item_id);
513        self.val
514    }
515
516    /// Given generic args for `item_id`, assert that the value is bound for `item_id` and
517    /// substitute it with the provided generic arguments. Because the arguments are bound in the
518    /// context of another item, so it the resulting substituted value.
519    pub fn substitute<OtherItem: Debug + Copy + PartialEq>(
520        self,
521        args: ItemBinder<OtherItem, &GenericArgs>,
522    ) -> ItemBinder<OtherItem, T>
523    where
524        ItemId: Into<ItemId>,
525        T: TyVisitable,
526    {
527        args.map_bound(|args| self.val.substitute(args))
528    }
529}
530
531/// Dummy item identifier that represents the current item when not ambiguous.
532#[derive(Debug, Copy, Clone, PartialEq, Eq)]
533pub struct CurrentItem;
534
535impl<T> ItemBinder<CurrentItem, T> {
536    pub fn under_current_binder(self) -> T {
537        self.val
538    }
539}