Skip to main content

flux_middle/rty/
refining.rs

1//! *Refining* is the process of generating a refined version of a rust type.
2//!
3//! Concretely, this module provides functions to go from types in [`flux_rustc_bridge::ty`] to
4//! types in [`rty`].
5
6use flux_arc_interner::{List, SliceInternable};
7use flux_common::bug;
8use flux_rustc_bridge::{ty, ty::GenericArgsExt as _};
9use itertools::Itertools;
10use rustc_abi::VariantIdx;
11use rustc_data_structures::fx::FxHashMap;
12use rustc_hir::def_id::DefId;
13use rustc_middle::ty::ParamTy;
14use rustc_span::Symbol;
15use rustc_type_ir::INNERMOST;
16
17use super::{
18    RefineArgsExt,
19    fold::{TypeFoldable, TypeFolder, TypeVisitable},
20};
21use crate::{
22    global_env::{GlobalEnv, WeakKvarInfo, WeakKvarMap},
23    queries::{QueryErr, QueryResult},
24    query_bug,
25    rty::{self, Expr, fold::TypeSuperFoldable},
26};
27
28pub fn refine_generics(generics: &ty::Generics) -> rty::Generics {
29    let params = generics
30        .params
31        .iter()
32        .map(|param| refine_generic_param_def(false, param))
33        .collect();
34
35    rty::Generics {
36        own_params: params,
37        parent: generics.parent(),
38        parent_count: generics.parent_count(),
39        has_self: generics.orig.has_self,
40    }
41}
42
43pub(crate) fn refine_generic_param_def(
44    as_type: bool,
45    param: &ty::GenericParamDef,
46) -> rty::GenericParamDef {
47    rty::GenericParamDef {
48        kind: refine_generic_param_def_kind(as_type, param.kind),
49        index: param.index,
50        name: param.name,
51        def_id: param.def_id,
52    }
53}
54
55fn refine_generic_param_def_kind(
56    as_type: bool,
57    kind: ty::GenericParamDefKind,
58) -> rty::GenericParamDefKind {
59    match kind {
60        ty::GenericParamDefKind::Lifetime => rty::GenericParamDefKind::Lifetime,
61        ty::GenericParamDefKind::Type { has_default } => {
62            if as_type {
63                rty::GenericParamDefKind::Type { has_default }
64            } else {
65                rty::GenericParamDefKind::Base { has_default }
66            }
67        }
68        ty::GenericParamDefKind::Const { has_default, .. } => {
69            rty::GenericParamDefKind::Const { has_default }
70        }
71    }
72}
73
74pub struct Refiner<'genv, 'tcx> {
75    genv: GlobalEnv<'genv, 'tcx>,
76    def_id: DefId,
77    generics: rty::Generics,
78    refine: fn(rty::BaseTy) -> rty::SubsetTyCtor,
79}
80
81impl<'genv, 'tcx> Refiner<'genv, 'tcx> {
82    pub fn new_for_item(
83        genv: GlobalEnv<'genv, 'tcx>,
84        def_id: DefId,
85        refine: fn(rty::BaseTy) -> rty::SubsetTyCtor,
86    ) -> QueryResult<Self> {
87        let generics = genv.generics_of(def_id)?;
88        Ok(Self { genv, def_id, generics, refine })
89    }
90
91    pub fn default_for_item(genv: GlobalEnv<'genv, 'tcx>, def_id: DefId) -> QueryResult<Self> {
92        Self::new_for_item(genv, def_id, refine_default)
93    }
94
95    pub fn with_holes(genv: GlobalEnv<'genv, 'tcx>, def_id: DefId) -> QueryResult<Self> {
96        Self::new_for_item(genv, def_id, |bty| {
97            let sort = bty.sort();
98            let constr = rty::SubsetTy::new(
99                bty.shift_in_escaping(1),
100                rty::Expr::nu(),
101                rty::Expr::hole(rty::HoleKind::Pred),
102            );
103            rty::Binder::bind_with_sort(constr, sort)
104        })
105    }
106
107    pub fn refine<T: Refine + ?Sized>(&self, t: &T) -> QueryResult<T::Output> {
108        t.refine(self)
109    }
110
111    fn refine_existential_predicate_generic_args(
112        &self,
113        def_id: DefId,
114        args: &ty::GenericArgs,
115    ) -> QueryResult<rty::GenericArgs> {
116        let generics = self.generics_of(def_id)?;
117        args.iter()
118            .enumerate()
119            .map(|(idx, arg)| {
120                // We need to skip the generic for Self
121                let param = generics.param_at(idx + 1, self.genv)?;
122                self.refine_generic_arg(&param, arg)
123            })
124            .try_collect()
125    }
126
127    pub fn refine_variant_def(
128        &self,
129        adt_def_id: DefId,
130        variant_idx: VariantIdx,
131    ) -> QueryResult<rty::PolyVariant> {
132        let adt_def = self.adt_def(adt_def_id)?;
133        let variant_def = adt_def.variant(variant_idx);
134        let fields = variant_def
135            .fields
136            .iter()
137            .map(|fld| {
138                let ty = self.genv.lower_type_of(fld.did)?.instantiate_identity();
139                ty.refine(self)
140            })
141            .try_collect()?;
142
143        let idx = if adt_def.sort_def().is_struct() {
144            rty::Expr::unit_struct(adt_def_id)
145        } else {
146            rty::Expr::ctor_enum(adt_def_id, variant_idx)
147        };
148        let value = rty::VariantSig::new(
149            adt_def,
150            rty::GenericArg::identity_for_item(self.genv, adt_def_id)?,
151            fields,
152            idx,
153            List::empty(),
154        );
155
156        Ok(rty::Binder::bind_with_vars(value, List::empty()))
157    }
158
159    pub fn refine_generic_args(
160        &self,
161        def_id: DefId,
162        args: &ty::GenericArgs,
163    ) -> QueryResult<rty::GenericArgs> {
164        let generics = self.generics_of(def_id)?;
165        args.iter()
166            .enumerate()
167            .map(|(idx, arg)| {
168                let param = generics.param_at(idx, self.genv)?;
169                self.refine_generic_arg(&param, arg)
170            })
171            .collect()
172    }
173
174    pub fn refine_generic_arg(
175        &self,
176        param: &rty::GenericParamDef,
177        arg: &ty::GenericArg,
178    ) -> QueryResult<rty::GenericArg> {
179        match (&param.kind, arg) {
180            (rty::GenericParamDefKind::Type { .. }, ty::GenericArg::Ty(ty)) => {
181                Ok(rty::GenericArg::Ty(ty.refine(self)?))
182            }
183            (rty::GenericParamDefKind::Base { .. }, ty::GenericArg::Ty(ty)) => {
184                let rty::TyOrBase::Base(contr) = self.refine_ty_or_base(ty)? else {
185                    return Err(QueryErr::InvalidGenericArg { def_id: param.def_id });
186                };
187                Ok(rty::GenericArg::Base(contr))
188            }
189            (rty::GenericParamDefKind::Lifetime, ty::GenericArg::Lifetime(re)) => {
190                Ok(rty::GenericArg::Lifetime(*re))
191            }
192            (rty::GenericParamDefKind::Const { .. }, ty::GenericArg::Const(ct)) => {
193                Ok(rty::GenericArg::Const(ct.clone()))
194            }
195            _ => bug!("mismatched generic arg `{arg:?}` `{param:?}`"),
196        }
197    }
198
199    fn refine_alias_term(&self, alias_term: &ty::AliasTerm) -> QueryResult<rty::AliasTerm> {
200        let args = self.refine_generic_args(alias_term.def_id(), &alias_term.args)?;
201        Ok(rty::AliasTerm::new(alias_term.kind, args))
202    }
203
204    fn refine_alias_ty(&self, alias_ty: &ty::AliasTy) -> QueryResult<rty::AliasTy> {
205        match alias_ty.kind {
206            // Only opaque types carry refinement arguments
207            ty::AliasKind::Opaque { def_id } => {
208                let args = self.refine_generic_args(def_id, &alias_ty.args)?;
209                let refine_args = rty::RefineArgs::for_item(self.genv, def_id, |param, _| {
210                    let param = param.instantiate(self.genv.tcx(), &args, &[]);
211                    Ok(rty::Expr::hole(rty::HoleKind::Expr(param.sort)))
212                })?;
213                Ok(rty::AliasTy::new(alias_ty.kind, args, refine_args))
214            }
215            ty::AliasKind::Projection { def_id } | ty::AliasKind::Free { def_id } => {
216                let args = self.refine_generic_args(def_id, &alias_ty.args)?;
217                Ok(rty::AliasTy::new(alias_ty.kind, args, List::empty()))
218            }
219        }
220    }
221
222    pub fn refine_ty_or_base(&self, ty: &ty::Ty) -> QueryResult<rty::TyOrBase> {
223        let bty = match ty.kind() {
224            ty::TyKind::Closure(did, args) => {
225                let no_panic = self.genv.no_panic(*did);
226                let closure_args = args.as_closure();
227                let upvar_tys = closure_args
228                    .upvar_tys()
229                    .iter()
230                    .map(|ty| ty.refine(self))
231                    .try_collect()?;
232                rty::BaseTy::Closure(*did, upvar_tys, args.clone(), no_panic)
233            }
234            ty::TyKind::Coroutine(did, args) => {
235                let coroutine_args = args.as_coroutine();
236                let resume_ty = coroutine_args.resume_ty().refine(self)?;
237                let upvar_tys = coroutine_args
238                    .upvar_tys()
239                    .map(|ty| ty.refine(self))
240                    .try_collect()?;
241                rty::BaseTy::Coroutine(*did, resume_ty, upvar_tys, args.clone())
242            }
243            ty::TyKind::CoroutineWitness(..) => {
244                bug!("implement when we know what this is");
245            }
246            ty::TyKind::Never => rty::BaseTy::Never,
247            ty::TyKind::Ref(r, ty, mutbl) => rty::BaseTy::Ref(*r, ty.refine(self)?, *mutbl),
248            ty::TyKind::Float(float_ty) => rty::BaseTy::Float(*float_ty),
249            ty::TyKind::Tuple(tys) => {
250                let tys = tys.iter().map(|ty| ty.refine(self)).try_collect()?;
251                rty::BaseTy::Tuple(tys)
252            }
253            ty::TyKind::Array(ty, len) => rty::BaseTy::Array(ty.refine(self)?, len.clone()),
254            ty::TyKind::Param(param_ty) => {
255                match self.param(*param_ty)?.kind {
256                    rty::GenericParamDefKind::Type { .. } => {
257                        return Ok(rty::TyOrBase::Ty(rty::Ty::param(*param_ty)));
258                    }
259                    rty::GenericParamDefKind::Base { .. } => rty::BaseTy::Param(*param_ty),
260                    rty::GenericParamDefKind::Lifetime | rty::GenericParamDefKind::Const { .. } => {
261                        bug!()
262                    }
263                }
264            }
265            ty::TyKind::Adt(adt_def, args) => {
266                let adt_def = self.genv.adt_def(adt_def.did())?;
267                let args = self.refine_generic_args(adt_def.did(), args)?;
268                rty::BaseTy::adt(adt_def, args)
269            }
270            ty::TyKind::FnDef(def_id, args) => {
271                let args = self.refine_generic_args(*def_id, args)?;
272                rty::BaseTy::fn_def(*def_id, args)
273            }
274            ty::TyKind::Alias(alias_ty) => {
275                let alias_ty = self.as_default().refine_alias_ty(alias_ty)?;
276                rty::BaseTy::Alias(alias_ty)
277            }
278            ty::TyKind::Bool => rty::BaseTy::Bool,
279            ty::TyKind::Int(int_ty) => rty::BaseTy::Int(*int_ty),
280            ty::TyKind::Uint(uint_ty) => rty::BaseTy::Uint(*uint_ty),
281            ty::TyKind::Foreign(def_id) => rty::BaseTy::Foreign(*def_id),
282            ty::TyKind::Str => rty::BaseTy::Str,
283            ty::TyKind::Slice(ty) => rty::BaseTy::Slice(ty.refine(self)?),
284            ty::TyKind::Char => rty::BaseTy::Char,
285            ty::TyKind::FnPtr(poly_fn_sig) => {
286                rty::BaseTy::FnPtr(poly_fn_sig.refine(&self.as_default())?)
287            }
288            ty::TyKind::RawPtr(ty, mu) => rty::BaseTy::RawPtr(ty.refine(&self.as_default())?, *mu),
289            ty::TyKind::Dynamic(exi_preds, r) => {
290                let exi_preds = exi_preds
291                    .iter()
292                    .map(|pred| pred.refine(self))
293                    .try_collect()?;
294                rty::BaseTy::Dynamic(exi_preds, *r)
295            }
296            ty::TyKind::Pat => rty::BaseTy::Pat,
297        };
298        Ok(rty::TyOrBase::Base((self.refine)(bty)))
299    }
300
301    fn as_default(&self) -> Self {
302        Refiner { refine: refine_default, generics: self.generics.clone(), ..*self }
303    }
304
305    fn adt_def(&self, def_id: DefId) -> QueryResult<rty::AdtDef> {
306        self.genv.adt_def(def_id)
307    }
308
309    fn generics_of(&self, def_id: DefId) -> QueryResult<rty::Generics> {
310        self.genv.generics_of(def_id)
311    }
312
313    fn param(&self, param_ty: ParamTy) -> QueryResult<rty::GenericParamDef> {
314        self.generics.param_at(param_ty.index as usize, self.genv)
315    }
316}
317
318pub trait Refine {
319    type Output;
320
321    fn refine(&self, refiner: &Refiner) -> QueryResult<Self::Output>;
322}
323
324impl Refine for ty::Ty {
325    type Output = rty::Ty;
326
327    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::Ty> {
328        Ok(refiner.refine_ty_or_base(self)?.into_ty())
329    }
330}
331
332impl<T: Refine> Refine for ty::Binder<T> {
333    type Output = rty::Binder<T::Output>;
334
335    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::Binder<T::Output>> {
336        let vars = refine_bound_variables(self.vars());
337        let inner = self.skip_binder_ref().refine(refiner)?;
338        Ok(rty::Binder::bind_with_vars(inner, vars))
339    }
340}
341
342impl Refine for ty::FnSig {
343    type Output = rty::FnSig;
344
345    // TODO(hof2)
346    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::FnSig> {
347        let inputs = self
348            .inputs()
349            .iter()
350            .map(|ty| ty.refine(refiner))
351            .try_collect()?;
352        let ret = self.output().refine(refiner)?.shift_in_escaping(1);
353        let output = rty::Binder::bind_with_vars(rty::FnOutput::new(ret, vec![]), List::empty());
354        // TODO(hof2) make a hoister to hoist all the stuff out of the inputs,
355        // the hoister will have a list of all the variables it hoisted and the
356        // single hole for the "requires"; then we "fill" the hole with a KVAR
357        // and generate a PolyFnSig with the hoisted variables
358        // see `into_bb_env` in `type_env.rs` for an example.
359        Ok(rty::FnSig::new(self.safety, self.abi, List::empty(), inputs, output, Expr::ff(), true))
360    }
361}
362
363impl Refine for ty::Clause {
364    type Output = rty::Clause;
365
366    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::Clause> {
367        Ok(rty::Clause { kind: self.kind.refine(refiner)? })
368    }
369}
370
371impl Refine for ty::TraitRef {
372    type Output = rty::TraitRef;
373
374    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::TraitRef> {
375        Ok(rty::TraitRef {
376            def_id: self.def_id,
377            args: refiner.refine_generic_args(self.def_id, &self.args)?,
378        })
379    }
380}
381
382impl Refine for ty::ClauseKind {
383    type Output = rty::ClauseKind;
384
385    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::ClauseKind> {
386        let kind = match self {
387            ty::ClauseKind::Trait(trait_pred) => {
388                let pred = rty::TraitPredicate { trait_ref: trait_pred.trait_ref.refine(refiner)? };
389                rty::ClauseKind::Trait(pred)
390            }
391            ty::ClauseKind::Projection(proj_pred) => {
392                let rty::TyOrBase::Base(term) = refiner.refine_ty_or_base(&proj_pred.term)? else {
393                    return Err(query_bug!(
394                        refiner.def_id,
395                        "sorry, we can't handle non-base associated types"
396                    ));
397                };
398                let pred = rty::ProjectionPredicate {
399                    projection_term: refiner.refine_alias_term(&proj_pred.projection_term)?,
400                    term,
401                };
402                rty::ClauseKind::Projection(pred)
403            }
404            ty::ClauseKind::RegionOutlives(pred) => {
405                let pred = rty::OutlivesPredicate(pred.0, pred.1);
406                rty::ClauseKind::RegionOutlives(pred)
407            }
408            ty::ClauseKind::TypeOutlives(pred) => {
409                let pred = rty::OutlivesPredicate(pred.0.refine(refiner)?, pred.1);
410                rty::ClauseKind::TypeOutlives(pred)
411            }
412            ty::ClauseKind::ConstArgHasType(const_, ty) => {
413                rty::ClauseKind::ConstArgHasType(const_.clone(), ty.refine(&refiner.as_default())?)
414            }
415            ty::ClauseKind::UnstableFeature(sym) => rty::ClauseKind::UnstableFeature(*sym),
416        };
417        Ok(kind)
418    }
419}
420
421impl Refine for ty::ExistentialPredicate {
422    type Output = rty::ExistentialPredicate;
423
424    fn refine(&self, refiner: &Refiner) -> QueryResult<Self::Output> {
425        let pred = match self {
426            ty::ExistentialPredicate::Trait(trait_ref) => {
427                rty::ExistentialPredicate::Trait(rty::ExistentialTraitRef {
428                    def_id: trait_ref.def_id,
429                    args: refiner.refine_existential_predicate_generic_args(
430                        trait_ref.def_id,
431                        &trait_ref.args,
432                    )?,
433                })
434            }
435            ty::ExistentialPredicate::Projection(projection) => {
436                let rty::TyOrBase::Base(term) = refiner.refine_ty_or_base(&projection.term)? else {
437                    return Err(query_bug!(
438                        refiner.def_id,
439                        "sorry, we can't handle non-base associated types"
440                    ));
441                };
442                rty::ExistentialPredicate::Projection(rty::ExistentialProjection {
443                    def_id: projection.def_id,
444                    args: refiner.refine_existential_predicate_generic_args(
445                        projection.def_id,
446                        &projection.args,
447                    )?,
448                    term,
449                })
450            }
451            ty::ExistentialPredicate::AutoTrait(def_id) => {
452                rty::ExistentialPredicate::AutoTrait(*def_id)
453            }
454        };
455        Ok(pred)
456    }
457}
458
459impl Refine for ty::GenericPredicates {
460    type Output = rty::GenericPredicates;
461
462    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::GenericPredicates> {
463        Ok(rty::GenericPredicates {
464            parent: self.parent,
465            predicates: refiner.refine(&self.predicates)?,
466        })
467    }
468}
469
470impl<T> Refine for List<T>
471where
472    T: SliceInternable,
473    T: Refine<Output: SliceInternable>,
474{
475    type Output = rty::List<T::Output>;
476
477    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::List<T::Output>> {
478        refiner.refine(&self[..])
479    }
480}
481
482impl<T> Refine for [T]
483where
484    T: Refine<Output: SliceInternable>,
485{
486    type Output = rty::List<T::Output>;
487
488    fn refine(&self, refiner: &Refiner) -> QueryResult<rty::List<T::Output>> {
489        self.iter().map(|t| refiner.refine(t)).try_collect()
490    }
491}
492
493fn refine_default(bty: rty::BaseTy) -> rty::SubsetTyCtor {
494    let sort = bty.sort();
495    let constr = rty::SubsetTy::trivial(bty.shift_in_escaping(1), rty::Expr::nu());
496    rty::Binder::bind_with_sort(constr, sort)
497}
498
499pub fn refine_bound_variables(vars: &[ty::BoundVariableKind]) -> List<rty::BoundVariableKind> {
500    vars.iter()
501        .map(|kind| {
502            match kind {
503                ty::BoundVariableKind::Region(kind) => rty::BoundVariableKind::Region(*kind),
504            }
505        })
506        .collect()
507}
508
509impl rty::PolyFnSig {
510    pub fn add_weak_kvars(self, genv: GlobalEnv, def_id: DefId) -> QueryResult<Self> {
511        let refinement_generics = genv.refinement_generics_of(def_id)?;
512        let early_param_sorts: FxHashMap<Symbol, rty::Sort> = refinement_generics
513            .0
514            .own_params
515            .iter()
516            .map(|param| (param.name, param.sort.clone()))
517            .collect();
518        let early_vars = self
519            .early_params()
520            .into_iter()
521            .filter_map(|param| {
522                let sort = early_param_sorts.get(&param.name).unwrap().clone();
523                if !sort.is_param() && !sort.is_loc() {
524                    Some((rty::Var::EarlyParam(param), sort))
525                } else {
526                    None
527                }
528            })
529            .collect_vec();
530        let late_vars = make_vars_and_sorts_from_bound_vars(self.vars());
531        Ok(self.map(|fn_sig| {
532            let mut params = late_vars.into_iter().chain(early_vars).collect_vec();
533            let mut wkvar_inserter = WeakKVarInserter {
534                wkvar_map: WeakKvarMap::default(),
535                def_id,
536                kvid: rty::KVid::from(0_usize),
537                existential_params: Vec::new(),
538                params: params.clone(),
539            };
540            let requires_wkvar = make_weak_kvar(
541                &mut wkvar_inserter.wkvar_map,
542                def_id,
543                &mut wkvar_inserter.kvid,
544                Vec::new(),
545                params.clone(),
546            );
547            let inputs = fn_sig
548                .inputs
549                .iter()
550                .map(|input| wkvar_inserter.fold_ty(input))
551                .collect();
552            shift_in_vars(&mut params);
553            let output_binder_params = make_vars_and_sorts_from_bound_vars(fn_sig.output.vars());
554            params.extend(output_binder_params);
555            wkvar_inserter.params = params.clone();
556            let ensures = if !fn_sig.output.vars().is_empty() {
557                let ensures_wkvar = make_weak_kvar(
558                    &mut wkvar_inserter.wkvar_map,
559                    def_id,
560                    &mut wkvar_inserter.kvid,
561                    make_vars_and_sorts_from_bound_vars(fn_sig.output.vars()),
562                    params.clone(),
563                );
564                fn_sig
565                    .output
566                    .skip_binder_ref()
567                    .ensures
568                    .iter()
569                    .cloned()
570                    .chain(std::iter::once(rty::Ensures::Pred(rty::Expr::wkvar(ensures_wkvar))))
571                    .collect()
572            } else {
573                fn_sig.output.skip_binder_ref().ensures.clone()
574            };
575            let output = fn_sig
576                .output
577                .map(|output| rty::FnOutput { ret: wkvar_inserter.fold_ty(&output.ret), ensures });
578            genv.feed_weak_kvars(def_id, wkvar_inserter.wkvar_map);
579
580            rty::FnSig {
581                abi: fn_sig.abi,
582                safety: fn_sig.safety,
583                inputs,
584                // NOTE(CK): Not sure whether we can avoid the clone.
585                requires: fn_sig
586                    .requires
587                    .iter()
588                    .cloned()
589                    .chain(std::iter::once(rty::Expr::wkvar(requires_wkvar)))
590                    .collect(),
591                output,
592                lifted: fn_sig.lifted,
593                no_panic: fn_sig.no_panic,
594            }
595        }))
596    }
597}
598
599struct WeakKVarInserter {
600    wkvar_map: WeakKvarMap,
601    def_id: DefId,
602    kvid: rty::KVid,
603    existential_params: Vec<Vec<(rty::Var, rty::Sort)>>,
604    params: Vec<(rty::Var, rty::Sort)>,
605}
606
607impl TypeFolder for WeakKVarInserter {
608    fn fold_ty(&mut self, ty: &rty::Ty) -> rty::Ty {
609        use rty::{Expr, Ty, TyKind::*};
610        match ty.kind() {
611            // This is the only recursive case where we need to update the params
612            // since we're going under a binder.
613            //
614            // We handle the shifting in and out explicitly rather than using
615            // the enter_binder and exit_binder methods because we immediately
616            // use the bound vars to make a weak kvar.
617            Exists(bound_ty) => {
618                for v in &mut self.existential_params {
619                    shift_in_vars(v);
620                }
621                shift_in_vars(&mut self.params);
622                let exist_params = make_vars_and_sorts_from_bound_vars(bound_ty.vars());
623                // Take all of the current existential params + the current params,
624                // AFTER shifting in.
625                let params = self
626                    .existential_params
627                    .iter()
628                    .flatten()
629                    .chain(self.params.iter())
630                    .cloned()
631                    .collect();
632                // We pass the params immediately under this binder as the self args.
633                //
634                // The purpose of self args is to ensure that we don't have duplication
635                // of suggestions.
636                //
637                // Suppose after we add weak kvars we have the type
638                //
639                //     fn ({exists v0. Vec<i32>[v0] | $wk1[v0]()}) requires $wk0[]()
640                //
641                // If we are looking to instantiate a weak kvar to the
642                // expression `2 > 1` (for some reason), we can validly put it
643                // in both $wk0 and $wk1. But the self arg ensures that we don't
644                // put it in $wk1, since it requires the expression contain one
645                // of its self args (in this case, just `v0`).
646                let wkvar = make_weak_kvar(
647                    &mut self.wkvar_map,
648                    self.def_id,
649                    &mut self.kvid,
650                    exist_params.clone(),
651                    params,
652                );
653                // Now we add the params for future weak kvars.
654                self.existential_params.push(exist_params);
655                let new_ty = bound_ty.skip_binder_ref().super_fold_with(self);
656                self.existential_params.pop();
657                for v in &mut self.existential_params {
658                    shift_out_vars(v);
659                }
660                shift_out_vars(&mut self.params);
661                Ty::exists(rty::Binder::bind_with_vars(
662                    Ty::constr(Expr::wkvar(wkvar), new_ty),
663                    bound_ty.vars().clone(),
664                ))
665            }
666            _ => ty.super_fold_with(self),
667        }
668    }
669
670    fn fold_bty(&mut self, bty: &rty::BaseTy) -> rty::BaseTy {
671        use rty::{BaseTy, Expr, GenericArg};
672        match bty {
673            BaseTy::Adt(adt_def, args) => {
674                let new_args = args
675                    .iter()
676                    .map(|arg| {
677                        match arg {
678                            GenericArg::Base(subset_ty) => {
679                                for v in &mut self.existential_params {
680                                    shift_in_vars(v);
681                                }
682                                shift_in_vars(&mut self.params);
683                                let exist_params =
684                                    make_vars_and_sorts_from_bound_vars(subset_ty.vars());
685                                // Take all of the current existential params + the current params,
686                                // AFTER shifting in.
687                                let params = self
688                                    .existential_params
689                                    .iter()
690                                    .flatten()
691                                    .chain(self.params.iter())
692                                    .cloned()
693                                    .collect();
694                                // We pass the params immediately under this binder as the self args.
695                                // see the TyKind::Exists case.
696                                let wkvar = make_weak_kvar(
697                                    &mut self.wkvar_map,
698                                    self.def_id,
699                                    &mut self.kvid,
700                                    exist_params.clone(),
701                                    params,
702                                );
703                                // Now we add the params for future weak kvars.
704                                self.existential_params.push(exist_params);
705                                let new_ty = subset_ty.skip_binder_ref().super_fold_with(self);
706                                let new_ty_with_wkvar = new_ty.strengthen(Expr::wkvar(wkvar));
707                                self.existential_params.pop();
708                                for v in &mut self.existential_params {
709                                    shift_out_vars(v);
710                                }
711                                shift_out_vars(&mut self.params);
712                                GenericArg::Base(rty::Binder::bind_with_vars(
713                                    new_ty_with_wkvar,
714                                    subset_ty.vars().clone(),
715                                ))
716                            }
717                            _ => arg.fold_with(self),
718                        }
719                    })
720                    .collect();
721                BaseTy::Adt(adt_def.clone(), new_args)
722            }
723            // For these specific btys, we will recur and add wkvars
724            BaseTy::Ref(..) | BaseTy::Tuple(..) | BaseTy::Array(..) | BaseTy::Slice(..) => {
725                bty.super_fold_with(self)
726            }
727            // By default we will not recur on the bty to add wkvars
728            _ => bty.clone(),
729        }
730    }
731
732    fn fold_expr(&mut self, expr: &Expr) -> Expr {
733        expr.clone()
734    }
735
736    fn fold_sort(&mut self, sort: &rty::Sort) -> rty::Sort {
737        sort.clone()
738    }
739}
740
741/// NOTE(CK):
742///   * Skips params (we don't presently handle polymorphism, though even if we did,
743///     I'm not sure that we need to pass params to the weak kvars).
744///   * Skips locs because we can't encode those.
745///   * Skips unit + unit adts because they otherwise get encoded as a 0 tuple
746///     to fixpoint because we use them in the args to a weak kvar, which
747///     we don't want to do.
748fn make_vars_and_sorts_from_bound_vars<'a, I, II>(vars: I) -> Vec<(rty::Var, rty::Sort)>
749where
750    I: IntoIterator<IntoIter = II>,
751    II: DoubleEndedIterator<Item = &'a rty::BoundVariableKind>,
752{
753    vars.into_iter()
754        .enumerate()
755        .filter_map(|(i, var_kind)| {
756            if let rty::BoundVariableKind::Refine(sort, _, reft_kind) = var_kind
757                && !sort.is_param()
758                && !sort.is_loc()
759                && !sort.is_unit()
760                && sort.is_unit_adt().is_none()
761            {
762                let bound_reft = rty::BoundReft { var: rty::BoundVar::from(i), kind: *reft_kind };
763                Some((rty::Var::Bound(INNERMOST, bound_reft), sort.clone()))
764            } else {
765                None
766            }
767        })
768        .collect_vec()
769}
770
771// TODO: Use a Vec<Vec<_>> solution, per Nico.
772// This is a sort of annoying rearchitecture, but nothing impossible.
773fn shift_in_vars(vars: &mut [(rty::Var, rty::Sort)]) {
774    for (var, _) in vars.iter_mut() {
775        *var = var.shift_in(1);
776    }
777}
778
779fn shift_out_vars(vars: &mut [(rty::Var, rty::Sort)]) {
780    for (var, _) in vars.iter_mut() {
781        *var = var.shift_out(1);
782    }
783}
784
785// TODO: Don't make a weak kvar if the self_args is empty if there's a weak kvar
786//        that's been created before it with a superset of its params.
787fn make_weak_kvar(
788    wkvar_map: &mut WeakKvarMap,
789    def_id: DefId,
790    kvid: &mut rty::KVid,
791    self_args: Vec<(rty::Var, rty::Sort)>,
792    params: Vec<(rty::Var, rty::Sort)>,
793) -> rty::WKVar {
794    let num_self_args = self_args.len();
795    let (args, sorts): (Vec<rty::Var>, Vec<rty::Sort>) =
796        self_args.into_iter().chain(params).unzip();
797    let arg_exprs = args.into_iter().map(rty::Expr::var).collect();
798    // We don't have any solutions because these weak kvars are being generated
799    // (solutions only come from user annotations).
800    wkvar_map.insert(kvid.as_u32(), WeakKvarInfo { solutions: vec![], sorts });
801    let ret = rty::WKVar {
802        wkvid: rty::WKVid::new(def_id, *kvid),
803        self_args: num_self_args,
804        args: arg_exprs,
805    };
806    *kvid += 1;
807    ret
808}