Skip to main content

charon_lib/ast/
type_level.rs

1use derive_generic_visitor::*;
2use itertools::Itertools;
3use serde_state::{DeserializeState, SerializeState};
4use std::{collections::HashSet, mem};
5
6use crate::ast::*;
7
8pub mod regions;
9pub mod substitute;
10pub mod trait_proofs;
11pub mod type_info;
12pub mod types;
13pub mod vars;
14
15pub use regions::*;
16pub use substitute::*;
17pub use trait_proofs::*;
18pub use type_info::*;
19pub use types::*;
20pub use vars::*;
21
22/// A set of generic arguments.
23#[derive(Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
24#[derive(SerializeState, DeserializeState, Drive, DriveMut, DriveTwo)]
25pub struct GenericArgs {
26    pub regions: IndexVec<RegionId, Region>,
27    pub types: IndexVec<TypeVarId, Ty>,
28    pub const_generics: IndexVec<ConstGenericVarId, ConstantExpr>,
29    pub trait_refs: IndexVec<TraitClauseId, TraitRef>,
30}
31
32/// A quantified trait predicate, e.g. `for<'a> Type<'a>: Trait<'a, Args>`.
33pub type PolyTraitDeclRef = RegionBinder<TraitDeclRef>;
34
35/// .0 outlives .1
36#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
37#[derive(SerializeState, DeserializeState, Drive, DriveMut, DriveTwo)]
38pub struct OutlivesPred<T, U>(pub T, pub U);
39
40pub type RegionOutlives = OutlivesPred<Region, Region>;
41pub type TypeOutlives = OutlivesPred<Ty, Region>;
42
43/// A constraint over a trait associated type.
44///
45/// Example:
46/// ```text
47/// T : Foo<S = String>
48///         ^^^^^^^^^^
49/// ```
50#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
51#[derive(SerializeState, DeserializeState, Drive, DriveMut, DriveTwo)]
52pub struct TraitTypeConstraint {
53    pub trait_ref: TraitRef,
54    pub type_id: AssocTypeId,
55    pub ty: Ty,
56}
57
58pub type BoxedArgs = Box<GenericArgs>;
59
60/// Generic parameters for a declaration, including predicates.
61#[derive(Default, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
62#[derive(SerializeState, DeserializeState, Drive, DriveMut, DriveTwo)]
63pub struct GenericParams {
64    #[serde_state(stateless)]
65    pub regions: IndexVec<RegionId, RegionParam>,
66    #[serde_state(stateless)]
67    pub types: IndexVec<TypeVarId, TypeParam>,
68    pub const_generics: IndexVec<ConstGenericVarId, ConstGenericParam>,
69    // TODO: rename to match [GenericArgs]?
70    pub trait_clauses: IndexVec<TraitClauseId, TraitParam>,
71    /// The first region in the pair outlives the second region
72    pub regions_outlive: Vec<RegionBinder<RegionOutlives>>,
73    /// The type outlives the region
74    pub types_outlive: Vec<RegionBinder<TypeOutlives>>,
75    /// Constraints over trait associated types
76    pub trait_type_constraints: IndexVec<TraitTypeConstraintId, RegionBinder<TraitTypeConstraint>>,
77}
78
79#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
80#[derive(SerializeState, DeserializeState, Drive, DriveMut, DriveTwo)]
81#[cfg_attr(feature = "charon_on_charon", charon::variants_prefix("BK"))]
82pub enum BinderKind {
83    /// The parameters of a generic associated type.
84    TraitType(TraitDeclId, AssocTypeId),
85    /// The parameters of a trait method. Used in the `methods` lists in trait decls and trait
86    /// impls.
87    TraitMethod(TraitDeclId, TraitMethodId),
88    /// The parameters bound in a non-trait `impl` block. Used in the `Name`s of inherent methods.
89    InherentImplBlock,
90    /// Binder used for `dyn Trait` existential predicates.
91    Dyn,
92    /// Some other use of a binder outside the main Charon ast.
93    Other,
94}
95
96/// A value of type `T` bound by generic parameters. Used in any context where we're adding generic
97/// parameters that aren't on the top-level item, e.g. `for<'a>` clauses (uses `RegionBinder` for
98/// now), trait methods, GATs (TODO).
99#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
100#[derive(SerializeState, DeserializeState, Drive, DriveMut, DriveTwo)]
101pub struct Binder<T> {
102    #[cfg_attr(feature = "charon_on_charon", charon::rename("binder_params"))]
103    pub params: GenericParams,
104    /// Named this way to highlight accesses to the inner value that might be handling parameters
105    /// incorrectly. Prefer using helper methods.
106    #[cfg_attr(feature = "charon_on_charon", charon::rename("binder_value"))]
107    pub skip_binder: T,
108    /// The kind of binder this is.
109    #[cfg_attr(feature = "charon_on_charon", charon::opaque)]
110    pub kind: BinderKind,
111}
112
113/// A value of type `T` bound by regions. We should use `binder` instead but this causes name clash
114/// issues in the derived ocaml visitors.
115#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
116#[derive(SerializeState, DeserializeState, Drive, DriveMut, DriveTwo)]
117pub struct RegionBinder<T> {
118    #[cfg_attr(feature = "charon_on_charon", charon::rename("binder_regions"))]
119    #[serde_state(stateless)]
120    pub regions: IndexVec<RegionId, RegionParam>,
121    /// Named this way to highlight accesses to the inner value that might be handling parameters
122    /// incorrectly. Prefer using helper methods.
123    #[cfg_attr(feature = "charon_on_charon", charon::rename("binder_value"))]
124    pub skip_binder: T,
125}
126
127impl GenericArgs {
128    pub fn len(&self) -> usize {
129        let GenericArgs {
130            regions,
131            types,
132            const_generics,
133            trait_refs,
134        } = self;
135        regions.len() + types.len() + const_generics.len() + trait_refs.len()
136    }
137
138    pub fn is_empty(&self) -> bool {
139        self.len() == 0
140    }
141    /// Whether this has any explicit arguments (types, regions or const generics).
142    pub fn has_explicits(&self) -> bool {
143        !self.regions.is_empty() || !self.types.is_empty() || !self.const_generics.is_empty()
144    }
145    /// Whether this has any implicit arguments (trait refs).
146    pub fn has_implicits(&self) -> bool {
147        !self.trait_refs.is_empty()
148    }
149
150    pub fn empty() -> Self {
151        GenericArgs {
152            regions: Default::default(),
153            types: Default::default(),
154            const_generics: Default::default(),
155            trait_refs: Default::default(),
156        }
157    }
158
159    pub fn new(
160        regions: IndexVec<RegionId, Region>,
161        types: IndexVec<TypeVarId, Ty>,
162        const_generics: IndexVec<ConstGenericVarId, ConstantExpr>,
163        trait_refs: IndexVec<TraitClauseId, TraitRef>,
164    ) -> Self {
165        Self {
166            regions,
167            types,
168            const_generics,
169            trait_refs,
170        }
171    }
172    pub fn new_types(types: IndexVec<TypeVarId, Ty>) -> Self {
173        Self {
174            types,
175            ..Self::empty()
176        }
177    }
178    pub fn new_lifetimes(regions: IndexVec<RegionId, Region>) -> Self {
179        Self {
180            regions,
181            ..Self::empty()
182        }
183    }
184
185    /// Check whether this matches the given `GenericParams`.
186    /// TODO: check more things, e.g. that the trait refs use the correct trait and generics.
187    pub fn matches(&self, params: &GenericParams) -> bool {
188        params.regions.len() == self.regions.len()
189            && params.types.len() == self.types.len()
190            && params.const_generics.len() == self.const_generics.len()
191            && params.trait_clauses.len() == self.trait_refs.len()
192    }
193
194    /// Return the same generics, but where we pop the first type arguments.
195    /// This is useful for trait references (for pretty printing for instance),
196    /// because the first type argument is the type for which the trait is
197    /// implemented.
198    pub fn pop_first_type_arg(&self) -> (Ty, Self) {
199        let mut generics = self.clone();
200        let mut it = mem::take(&mut generics.types).into_iter();
201        let ty = it.next().unwrap();
202        generics.types = it.collect();
203        (ty, generics)
204    }
205
206    /// Concatenate this set of arguments with another one. Use with care, you must manage the
207    /// order of arguments correctly.
208    pub fn concat(mut self, other: &Self) -> Self {
209        let Self {
210            regions,
211            types,
212            const_generics,
213            trait_refs,
214        } = other;
215        self.regions.clone_extend_from_other(regions);
216        self.types.clone_extend_from_other(types);
217        self.const_generics.clone_extend_from_other(const_generics);
218        self.trait_refs.clone_extend_from_other(trait_refs);
219        self
220    }
221}
222
223impl GenericParams {
224    pub fn empty() -> Self {
225        Self::default()
226    }
227
228    pub fn is_empty(&self) -> bool {
229        self.len() == 0
230    }
231    /// Whether this has any explicit arguments (types, regions or const generics).
232    pub fn has_explicits(&self) -> bool {
233        !self.regions.is_empty() || !self.types.is_empty() || !self.const_generics.is_empty()
234    }
235    /// Whether this has any implicit arguments (trait clauses, outlives relations, associated type
236    /// equality constraints).
237    pub fn has_predicates(&self) -> bool {
238        !self.trait_clauses.is_empty()
239            || !self.types_outlive.is_empty()
240            || !self.regions_outlive.is_empty()
241            || !self.trait_type_constraints.is_empty()
242    }
243
244    /// Run some sanity checks.
245    pub fn check_consistency(&self) {
246        // Sanity check: check the clause ids are consistent.
247        assert!(
248            self.trait_clauses
249                .iter()
250                .enumerate()
251                .all(|(i, c)| c.clause_id.index() == i)
252        );
253
254        // Sanity check: region names are pairwise distinct (this caused trouble when generating
255        // names for the backward functions in Aeneas): at some point, Rustc introduced names equal
256        // to `Some("'_")` for the anonymous regions, instead of using `None` (we now check in
257        // [translate_region_name] and ignore names equal to "'_").
258        let mut s = HashSet::new();
259        for r in &self.regions {
260            if let Some(name) = &r.name {
261                assert!(
262                    !s.contains(name),
263                    "Name \"{}\" reused for two different lifetimes",
264                    name
265                );
266                s.insert(name);
267            }
268        }
269    }
270
271    pub fn len(&self) -> usize {
272        let GenericParams {
273            regions,
274            types,
275            const_generics,
276            trait_clauses,
277            regions_outlive,
278            types_outlive,
279            trait_type_constraints,
280        } = self;
281        regions.len()
282            + types.len()
283            + const_generics.len()
284            + trait_clauses.len()
285            + regions_outlive.len()
286            + types_outlive.len()
287            + trait_type_constraints.len()
288    }
289
290    /// Construct a set of generic arguments in the scope of `self` that matches `self` and feeds
291    /// each required parameter with itself. E.g. given parameters for `<T, U> where U:
292    /// PartialEq<T>`, the arguments would be `<T, U>[TraitClause0]`.
293    pub fn identity_args(&self) -> GenericArgs {
294        self.identity_args_at_depth(DeBruijnId::zero())
295    }
296
297    /// Like `identity_args` but uses variables bound at the given depth.
298    pub fn identity_args_at_depth(&self, depth: DeBruijnId) -> GenericArgs {
299        GenericArgs {
300            regions: self
301                .regions
302                .map_ref_indexed(|id, _| Region::Var(DeBruijnVar::bound(depth, id))),
303            types: self
304                .types
305                .map_ref_indexed(|id, _| TyKind::TypeVar(DeBruijnVar::bound(depth, id)).into_ty()),
306            const_generics: self.const_generics.map_ref_indexed(|id, c| {
307                ConstantExpr::new(
308                    ConstantExprKind::Var(DeBruijnVar::bound(depth, id)),
309                    c.ty.clone(),
310                )
311            }),
312            trait_refs: self
313                .trait_clauses
314                .map_ref(|clause| clause.identity_tref_at_depth(depth)),
315        }
316    }
317
318    /// Take the predicates from the another `GenericParams`. This assumes the clause ids etc are
319    /// already consistent.
320    pub fn take_predicates_from(&mut self, other: GenericParams) {
321        assert!(!other.has_explicits());
322        let num_clauses = self.trait_clauses.len();
323        let GenericParams {
324            regions: _,
325            types: _,
326            const_generics: _,
327            trait_clauses,
328            regions_outlive,
329            types_outlive,
330            trait_type_constraints,
331        } = other;
332        self.trait_clauses
333            .extend(trait_clauses.into_iter().update(|clause| {
334                clause.clause_id += num_clauses;
335            }));
336        self.regions_outlive.extend(regions_outlive);
337        self.types_outlive.extend(types_outlive);
338        self.trait_type_constraints.extend(trait_type_constraints);
339    }
340
341    /// Take the predicates from the another `GenericParams`. This assumes that the two
342    /// `GenericParams` are independent, hence will shift clause ids if `other` has any
343    /// trait refs that reference its own clauses.
344    pub fn merge_predicates_from(&mut self, mut other: GenericParams) {
345        // Drop the explicits params.
346        other.types.clear();
347        other.regions.clear();
348        other.const_generics.clear();
349        // The contents of `other` may refer to its own trait clauses, so we must shift clause ids.
350        struct ShiftClausesVisitor(usize);
351        impl VarsVisitor for ShiftClausesVisitor {
352            fn visit_clause_var(&mut self, v: ClauseDbVar) -> Option<TraitRefKind> {
353                if let DeBruijnVar::Bound(DeBruijnId::ZERO, clause_id) = v {
354                    // Replace clause 0 and decrement the others.
355                    Some(TraitRefKind::Clause(DeBruijnVar::Bound(
356                        DeBruijnId::ZERO,
357                        clause_id + self.0,
358                    )))
359                } else {
360                    None
361                }
362            }
363        }
364        let num_clauses = self.trait_clauses.len();
365        other.visit_vars(&mut ShiftClausesVisitor(num_clauses));
366        self.take_predicates_from(other);
367    }
368}
369
370impl<T> Binder<T> {
371    /// Wrap the value in an empty binder, shifting variables appropriately.
372    pub fn empty(kind: BinderKind, x: T) -> Self
373    where
374        T: TyVisitable,
375    {
376        Binder {
377            params: Default::default(),
378            skip_binder: x.move_under_binder(),
379            kind,
380        }
381    }
382    pub fn new(kind: BinderKind, params: GenericParams, skip_binder: T) -> Self {
383        Self {
384            params,
385            skip_binder,
386            kind,
387        }
388    }
389
390    /// Whether this binder binds any variables.
391    pub fn binds_anything(&self) -> bool {
392        !self.params.is_empty()
393    }
394
395    /// Retreive the contents of this binder if the binder binds no variables. This is the invers
396    /// of `Binder::empty`.
397    pub fn get_if_binds_nothing(&self) -> Option<T>
398    where
399        T: TyVisitable + Clone,
400    {
401        self.params
402            .is_empty()
403            .then(|| self.skip_binder.clone().move_from_under_binder().unwrap())
404    }
405
406    pub fn map<U>(self, f: impl FnOnce(T) -> U) -> Binder<U> {
407        Binder {
408            params: self.params,
409            skip_binder: f(self.skip_binder),
410            kind: self.kind,
411        }
412    }
413
414    pub fn map_ref<U>(&self, f: impl FnOnce(&T) -> U) -> Binder<U> {
415        Binder {
416            params: self.params.clone(),
417            skip_binder: f(&self.skip_binder),
418            kind: self.kind.clone(),
419        }
420    }
421
422    /// Substitute the provided arguments for the variables bound in this binder and return the
423    /// substituted inner value.
424    pub fn apply(self, args: &GenericArgs) -> T
425    where
426        T: TyVisitable,
427    {
428        self.skip_binder.substitute(args)
429    }
430
431    /// Like `apply`, but also keep the parameters: predicates mention them and therefore need to
432    /// be substituted before use too.
433    pub fn apply_keep_params(self, args: &GenericArgs) -> (GenericParams, T)
434    where
435        T: TyVisitable,
436    {
437        (
438            self.params.substitute(args),
439            self.skip_binder.substitute(args),
440        )
441    }
442}
443
444impl<T: AstVisitable> Binder<Binder<T>> {
445    /// Flatten two levels of binders into a single one.
446    pub fn flatten(self) -> Binder<T> {
447        #[derive(Visitor)]
448        struct FlattenVisitor<'a> {
449            shift_by: &'a GenericParams,
450            binder_depth: DeBruijnId,
451        }
452        impl VisitorWithBinderDepth for FlattenVisitor<'_> {
453            fn binder_depth_mut(&mut self) -> &mut DeBruijnId {
454                &mut self.binder_depth
455            }
456        }
457        impl VisitAstMut for FlattenVisitor<'_> {
458            fn visit<T: AstVisitable>(&mut self, x: &mut T) -> ControlFlow<Self::Break> {
459                VisitWithBinderDepth::new(self).visit(x)
460            }
461
462            fn enter_de_bruijn_id(&mut self, db_id: &mut DeBruijnId) {
463                if *db_id > self.binder_depth {
464                    // We started visiting at the inner binder, so in this branch we're either
465                    // mentioning the outer binder or a binder further beyond. Either way we
466                    // decrease the depth; variables that point to the outer binder don't have to
467                    // be shifted.
468                    *db_id = db_id.decr();
469                }
470            }
471            fn enter_region(&mut self, x: &mut Region) {
472                if let Region::Var(var) = x
473                    && let Some(id) = var.bound_at_depth_mut(self.binder_depth)
474                {
475                    *id += self.shift_by.regions.len();
476                }
477            }
478            fn enter_ty_kind(&mut self, x: &mut TyKind) {
479                if let TyKind::TypeVar(var) = x
480                    && let Some(id) = var.bound_at_depth_mut(self.binder_depth)
481                {
482                    *id += self.shift_by.types.len();
483                }
484            }
485            fn enter_constant_expr_kind(&mut self, kind: &mut ConstantExprKind) {
486                if let ConstantExprKind::Var(var) = kind
487                    && let Some(id) = var.bound_at_depth_mut(self.binder_depth)
488                {
489                    *id += self.shift_by.const_generics.len();
490                }
491            }
492            fn enter_trait_ref_kind(&mut self, x: &mut TraitRefKind) {
493                if let TraitRefKind::Clause(var) = x
494                    && let Some(id) = var.bound_at_depth_mut(self.binder_depth)
495                {
496                    *id += self.shift_by.trait_clauses.len();
497                }
498            }
499        }
500
501        // We will concatenate both sets of params.
502        let mut outer_params = self.params;
503
504        // The inner value needs to change:
505        // - at binder level 0 we shift all variable ids to match the concatenated params;
506        // - at binder level > 0 we decrease binding level because there's one fewer binder.
507        let mut bound_value = self.skip_binder.skip_binder;
508        let _ = bound_value.drive_mut(&mut FlattenVisitor {
509            shift_by: &outer_params,
510            binder_depth: Default::default(),
511        });
512
513        // The inner params must also be updated, as they can refer to themselves and the outer
514        // one.
515        let mut inner_params = self.skip_binder.params;
516        let _ = inner_params.drive_mut(&mut FlattenVisitor {
517            shift_by: &outer_params,
518            binder_depth: Default::default(),
519        });
520        inner_params
521            .regions
522            .iter_mut()
523            .for_each(|v| v.index += outer_params.regions.len());
524        inner_params
525            .types
526            .iter_mut()
527            .for_each(|v| v.index += outer_params.types.len());
528        inner_params
529            .const_generics
530            .iter_mut()
531            .for_each(|v| v.index += outer_params.const_generics.len());
532        inner_params
533            .trait_clauses
534            .iter_mut()
535            .for_each(|v| v.clause_id += outer_params.trait_clauses.len());
536
537        let GenericParams {
538            regions,
539            types,
540            const_generics,
541            trait_clauses,
542            regions_outlive,
543            types_outlive,
544            trait_type_constraints,
545        } = &inner_params;
546        outer_params.regions.clone_extend_from_other(regions);
547        outer_params.types.clone_extend_from_other(types);
548        outer_params
549            .const_generics
550            .clone_extend_from_other(const_generics);
551        outer_params
552            .trait_clauses
553            .clone_extend_from_other(trait_clauses);
554        outer_params
555            .regions_outlive
556            .extend_from_slice(regions_outlive);
557        outer_params.types_outlive.extend_from_slice(types_outlive);
558        outer_params
559            .trait_type_constraints
560            .clone_extend_from_other(trait_type_constraints);
561
562        Binder {
563            params: outer_params,
564            skip_binder: bound_value,
565            kind: BinderKind::Other,
566        }
567    }
568}
569
570impl<T> RegionBinder<T> {
571    /// Wrap the value in an empty region binder, shifting variables appropriately.
572    pub fn empty(x: T) -> Self
573    where
574        T: TyVisitable,
575    {
576        RegionBinder {
577            regions: Default::default(),
578            skip_binder: x.move_under_binder(),
579        }
580    }
581
582    pub fn map<U>(self, f: impl FnOnce(T) -> U) -> RegionBinder<U> {
583        RegionBinder {
584            regions: self.regions,
585            skip_binder: f(self.skip_binder),
586        }
587    }
588
589    pub fn map_ref<U>(&self, f: impl FnOnce(&T) -> U) -> RegionBinder<U> {
590        RegionBinder {
591            regions: self.regions.clone(),
592            skip_binder: f(&self.skip_binder),
593        }
594    }
595
596    /// Substitute the bound variables with the given lifetimes.
597    pub fn apply(self, regions: IndexVec<RegionId, Region>) -> T
598    where
599        T: TyVisitable,
600    {
601        assert_eq!(regions.len(), self.regions.len());
602        let args = GenericArgs {
603            regions,
604            ..GenericArgs::empty()
605        };
606        self.skip_binder.substitute_inner_binder(&args)
607    }
608
609    /// Substitute the bound variables with erased lifetimes.
610    pub fn erase(self) -> T
611    where
612        T: TyVisitable,
613    {
614        let regions = self.regions.map_ref_indexed(|_, _| Region::Erased);
615        self.apply(regions)
616    }
617}
618
619pub trait HasIdxVecOf<Id: Idx>: std::ops::Index<Id, Output: Sized> {
620    fn get_idx_vec(&self) -> &IndexVec<Id, Self::Output>;
621    fn get_idx_vec_mut(&mut self) -> &mut IndexVec<Id, Self::Output>;
622}
623
624/// Delegate `Index` implementations to subfields.
625macro_rules! mk_index_impls {
626    ($ty:ident.$field:ident[$idx:ty]: $output:ty) => {
627        impl std::ops::Index<$idx> for $ty {
628            type Output = $output;
629            fn index(&self, index: $idx) -> &Self::Output {
630                &self.$field[index]
631            }
632        }
633        impl std::ops::IndexMut<$idx> for $ty {
634            fn index_mut(&mut self, index: $idx) -> &mut Self::Output {
635                &mut self.$field[index]
636            }
637        }
638        impl HasIdxVecOf<$idx> for $ty {
639            fn get_idx_vec(&self) -> &IndexVec<$idx, Self::Output> {
640                &self.$field
641            }
642            fn get_idx_vec_mut(&mut self) -> &mut IndexVec<$idx, Self::Output> {
643                &mut self.$field
644            }
645        }
646    };
647}
648mk_index_impls!(GenericArgs.regions[RegionId]: Region);
649mk_index_impls!(GenericArgs.types[TypeVarId]: Ty);
650mk_index_impls!(GenericArgs.const_generics[ConstGenericVarId]: ConstantExpr);
651mk_index_impls!(GenericArgs.trait_refs[TraitClauseId]: TraitRef);
652mk_index_impls!(GenericParams.regions[RegionId]: RegionParam);
653mk_index_impls!(GenericParams.types[TypeVarId]: TypeParam);
654mk_index_impls!(GenericParams.const_generics[ConstGenericVarId]: ConstGenericParam);
655mk_index_impls!(GenericParams.trait_clauses[TraitClauseId]: TraitParam);