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#[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
32pub type PolyTraitDeclRef = RegionBinder<TraitDeclRef>;
34
35#[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#[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#[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 pub trait_clauses: IndexVec<TraitClauseId, TraitParam>,
71 pub regions_outlive: Vec<RegionBinder<RegionOutlives>>,
73 pub types_outlive: Vec<RegionBinder<TypeOutlives>>,
75 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 TraitType(TraitDeclId, AssocTypeId),
85 TraitMethod(TraitDeclId, TraitMethodId),
88 InherentImplBlock,
90 Dyn,
92 Other,
94}
95
96#[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 #[cfg_attr(feature = "charon_on_charon", charon::rename("binder_value"))]
107 pub skip_binder: T,
108 #[cfg_attr(feature = "charon_on_charon", charon::opaque)]
110 pub kind: BinderKind,
111}
112
113#[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 #[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 pub fn has_explicits(&self) -> bool {
143 !self.regions.is_empty() || !self.types.is_empty() || !self.const_generics.is_empty()
144 }
145 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 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 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 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 pub fn has_explicits(&self) -> bool {
233 !self.regions.is_empty() || !self.types.is_empty() || !self.const_generics.is_empty()
234 }
235 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 pub fn check_consistency(&self) {
246 assert!(
248 self.trait_clauses
249 .iter()
250 .enumerate()
251 .all(|(i, c)| c.clause_id.index() == i)
252 );
253
254 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 pub fn identity_args(&self) -> GenericArgs {
294 self.identity_args_at_depth(DeBruijnId::zero())
295 }
296
297 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 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 pub fn merge_predicates_from(&mut self, mut other: GenericParams) {
345 other.types.clear();
347 other.regions.clear();
348 other.const_generics.clear();
349 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 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 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 pub fn binds_anything(&self) -> bool {
392 !self.params.is_empty()
393 }
394
395 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 pub fn apply(self, args: &GenericArgs) -> T
425 where
426 T: TyVisitable,
427 {
428 self.skip_binder.substitute(args)
429 }
430
431 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 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 *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 let mut outer_params = self.params;
503
504 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 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 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 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 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
624macro_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);