Skip to main content

flux_refineck/
type_env.rs

1mod place_ty;
2
3use std::{iter, ops::ControlFlow};
4
5use flux_common::{
6    bug,
7    dbg::{SpanTrace, debug_assert_eq3},
8    tracked_span_bug, tracked_span_dbg_assert_eq,
9};
10use flux_infer::{
11    fixpoint_encoding::KVarEncoding,
12    infer::{ConstrReason, InferCtxt, InferCtxtAt, InferCtxtRoot, InferResult},
13    refine_tree::Scope,
14};
15use flux_macros::DebugAsJson;
16use flux_middle::{
17    PlaceExt as _,
18    global_env::GlobalEnv,
19    pretty::{PrettyCx, PrettyNested},
20    queries::QueryResult,
21    rty::{
22        BaseTy, Binder, BoundReftKind, BoundVariableKind, Ctor, Expr, ExprKind, FnOutput, FnSig,
23        GenericArg, HoleKind, INNERMOST, Lambda, List, Loc, Mutability, Path, PolyFnSig, PtrKind,
24        Region, SortCtor, SubsetTy, SubsetTyCtor, Ty, TyKind, VariantIdx,
25        canonicalize::{Hoister, LocalHoister},
26        fold::{FallibleTypeFolder, TypeFoldable, TypeVisitable, TypeVisitor},
27        region_matching::{rty_match_regions, ty_match_regions},
28    },
29};
30use flux_rustc_bridge::{
31    self,
32    mir::{BasicBlock, Body, Local, LocalDecl, LocalDecls, Place, PlaceElem},
33    ty,
34};
35use itertools::{Itertools, izip};
36use rustc_data_structures::unord::UnordMap;
37use rustc_index::{IndexSlice, IndexVec};
38use rustc_middle::{mir::RETURN_PLACE, ty::TyCtxt};
39use rustc_span::{Span, Symbol};
40use rustc_type_ir::BoundVar;
41use serde::Serialize;
42
43use self::place_ty::{LocKind, PlacesTree};
44use super::rty::Sort;
45
46#[derive(Clone, Default)]
47pub struct TypeEnv<'a> {
48    bindings: PlacesTree,
49    local_decls: &'a LocalDecls,
50}
51
52pub struct BasicBlockEnvShape {
53    scope: Scope,
54    bindings: PlacesTree,
55}
56
57pub struct BasicBlockEnv {
58    data: Binder<BasicBlockEnvData>,
59    scope: Scope,
60}
61
62#[derive(Debug)]
63struct BasicBlockEnvData {
64    constrs: List<Expr>,
65    bindings: PlacesTree,
66}
67
68impl<'a> TypeEnv<'a> {
69    pub fn new(infcx: &mut InferCtxt, body: &'a Body, fn_sig: &FnSig) -> TypeEnv<'a> {
70        let mut env = TypeEnv { bindings: PlacesTree::default(), local_decls: &body.local_decls };
71
72        for requires in fn_sig.requires() {
73            infcx.assume_pred(requires);
74        }
75
76        for (local, ty) in body.args_iter().zip(fn_sig.inputs()) {
77            let ty = infcx.unpack(ty);
78            infcx.assume_invariants(&ty);
79            env.alloc_with_ty(local, ty);
80        }
81
82        for local in body.vars_and_temps_iter() {
83            env.alloc(local);
84        }
85
86        env.alloc(RETURN_PLACE);
87        env
88    }
89
90    pub fn empty() -> TypeEnv<'a> {
91        TypeEnv { bindings: PlacesTree::default(), local_decls: IndexSlice::empty() }
92    }
93
94    fn alloc_with_ty(&mut self, local: Local, ty: Ty) {
95        let ty = ty_match_regions(&ty, &self.local_decls[local].ty);
96        self.bindings.insert(local.into(), LocKind::Local, ty);
97    }
98
99    fn alloc(&mut self, local: Local) {
100        self.bindings
101            .insert(local.into(), LocKind::Local, Ty::uninit());
102    }
103
104    pub(crate) fn into_infer(self, scope: Scope) -> BasicBlockEnvShape {
105        BasicBlockEnvShape::new(scope, self)
106    }
107
108    pub(crate) fn lookup_rust_ty(&self, genv: GlobalEnv, place: &Place) -> QueryResult<ty::Ty> {
109        Ok(place.ty(genv, self.local_decls)?.ty)
110    }
111
112    pub(crate) fn lookup_place(
113        &mut self,
114        infcx: &mut InferCtxtAt,
115        place: &Place,
116    ) -> InferResult<Ty> {
117        let span = infcx.span;
118        let result = self.bindings.lookup_unfolding(infcx, place, span)?;
119        Ok(result.ty)
120    }
121
122    pub(crate) fn get(&self, path: &Path) -> Ty {
123        self.bindings.get(path)
124    }
125
126    pub fn update_path(&mut self, path: &Path, new_ty: Ty, span: Span) {
127        self.bindings.lookup(path, span).update(new_ty);
128    }
129
130    /// When checking a borrow in the right hand side of an assignment `x = &'?n p`, we use the
131    /// annotated region `'?n` in the type of the result. This region will only be used temporarily
132    /// and then replaced by the region in the type of `x` after the assignment. See [`TypeEnv::assign`]
133    pub(crate) fn borrow(
134        &mut self,
135        infcx: &mut InferCtxtAt,
136        re: Region,
137        mutbl: Mutability,
138        place: &Place,
139    ) -> InferResult<Ty> {
140        let span = infcx.span;
141        let result = self.bindings.lookup_unfolding(infcx, place, span)?;
142        if result.is_strg && mutbl == Mutability::Mut {
143            Ok(Ty::ptr(PtrKind::Mut(re), result.path()))
144        } else {
145            // FIXME(nilehmann) we should block the place here. That would require a notion
146            // of shared vs mutable block types because sometimes blocked places from a shared
147            // reference never get unblocked and we should still allow reads through them.
148            Ok(Ty::mk_ref(re, result.ty, mutbl))
149        }
150    }
151
152    // FIXME(nilehmann) this is only used in a single place and we have it because [`TypeEnv`]
153    // doesn't expose a lookup without unfolding
154    pub(crate) fn ptr_to_ref_at_place(
155        &mut self,
156        infcx: &mut InferCtxtAt,
157        place: &Place,
158    ) -> InferResult {
159        let lookup = self.bindings.lookup(place, infcx.span);
160        let TyKind::Ptr(PtrKind::Mut(re), path) = lookup.ty.kind() else {
161            tracked_span_bug!("ptr_to_borrow called on non mutable pointer type")
162        };
163
164        let ref_ty =
165            self.ptr_to_ref(infcx, ConstrReason::Other, *re, path, PtrToRefBound::Infer)?;
166
167        self.bindings.lookup(place, infcx.span).update(ref_ty);
168
169        Ok(())
170    }
171
172    /// Convert a (strong) pointer to a mutable reference.
173    ///
174    /// This roughly implements the following inference rule:
175    /// ```text
176    ///                   t₁ <: t₂
177    /// -------------------------------------------------
178    /// Γ₁,ℓ:t1,Γ₂ ; ptr(mut, ℓ) => Γ₁,ℓ:†t₂,Γ₂ ; &mut t2
179    /// ```
180    /// That's it, we first get the current type `t₁` at location `ℓ` and check it is a subtype
181    /// of `t₂`. Then, we update the type of `ℓ` to `t₂` and block the place.
182    ///
183    /// The bound `t₂` can be either inferred ([`PtrToRefBound::Infer`]), explicitly provided
184    /// ([`PtrToRefBound::Ty`]), or made equal to `t₁` ([`PtrToRefBound::Identity`]).
185    ///
186    /// As an example, consider the environment `x: i32[a]` and the pointer `ptr(mut, x)`.
187    /// Converting the pointer to a mutable reference with an inferred bound produces the following
188    /// derivation (roughly):
189    ///
190    /// ```text
191    ///                    i32[a] <: i32{v: $k(v)}
192    /// ----------------------------------------------------------------
193    /// x: i32[a] ; ptr(mut, x) => x:†i32{v: $k(v)} ; &mut i32{v: $k(v)}
194    /// ```
195    pub(crate) fn ptr_to_ref(
196        &mut self,
197        infcx: &mut InferCtxtAt,
198        reason: ConstrReason,
199        re: Region,
200        path: &Path,
201        bound: PtrToRefBound,
202    ) -> InferResult<Ty> {
203        // ℓ: t1
204        let t1 = self.bindings.lookup(path, infcx.span).fold(infcx)?;
205
206        // t1 <: t2
207        let t2 = match bound {
208            PtrToRefBound::Ty(t2) => {
209                let t2 = rty_match_regions(&t2, &t1);
210                infcx.subtyping_with_env(self, &t1, &t2, reason)?;
211                t2
212            }
213            PtrToRefBound::Infer => {
214                let t2 = t1.with_holes().replace_holes(|sorts, kind| {
215                    debug_assert_eq!(kind, HoleKind::Pred);
216                    infcx.fresh_kvar(sorts, KVarEncoding::Conj)
217                });
218                infcx.subtyping_with_env(self, &t1, &t2, reason)?;
219                t2
220            }
221            PtrToRefBound::Identity => t1.clone(),
222        };
223
224        // ℓ: †t2
225        self.bindings
226            .lookup(path, infcx.span)
227            .block_with(t2.clone());
228
229        Ok(Ty::mk_ref(re, t2, Mutability::Mut))
230    }
231
232    pub(crate) fn fold_local_ptrs(&mut self, infcx: &mut InferCtxtAt) -> InferResult {
233        for (loc, bound, ty) in self.bindings.local_ptrs() {
234            infcx.subtyping(&ty, &bound, ConstrReason::FoldLocal)?;
235            self.bindings.remove_local(&loc);
236        }
237        Ok(())
238    }
239
240    /// Updates the type of `place` to `new_ty`. This may involve a *strong update* if we have
241    /// ownership of `place` or a *weak update* if it's behind a reference (which fires a subtyping
242    /// constraint)
243    ///
244    /// When strong updating, the process involves recovering the original regions (lifetimes) used
245    /// in the (unrefined) Rust type of `place` and then substituting these regions in `new_ty`. For
246    /// instance, if we are assigning a value of type `S<&'?10 i32{v: v > 0}>` to a variable `x`,
247    /// and the (unrefined) Rust type of `x` is `S<&'?5 i32>`, before the assignment, we identify a
248    /// substitution that maps the region `'?10` to `'?5`. After applying this substitution, the
249    /// type of the place `x` is updated accordingly. This ensures that the lifetimes in the
250    /// assigned type are consistent with those expected by the place's original type definition.
251    pub(crate) fn assign(
252        &mut self,
253        infcx: &mut InferCtxtAt,
254        place: &Place,
255        new_ty: Ty,
256    ) -> InferResult {
257        let rustc_ty = place.ty(infcx.genv, self.local_decls)?.ty;
258        let new_ty = ty_match_regions(&new_ty, &rustc_ty);
259        let span = infcx.span;
260        let result = self.bindings.lookup_unfolding(infcx, place, span)?;
261        if result.is_strg {
262            result.update(new_ty);
263        } else if !place.behind_raw_ptr(infcx.genv, self.local_decls)? {
264            infcx.subtyping(&new_ty, &result.ty, ConstrReason::Assign)?;
265        }
266        Ok(())
267    }
268
269    pub(crate) fn move_place(&mut self, infcx: &mut InferCtxtAt, place: &Place) -> InferResult<Ty> {
270        let span = infcx.span;
271        let result = self.bindings.lookup_unfolding(infcx, place, span)?;
272        if result.is_strg {
273            let uninit = Ty::uninit();
274            Ok(result.update(uninit))
275        } else {
276            // ignore the 'move' and trust rustc managed the move correctly
277            // https://github.com/flux-rs/flux/issues/725#issuecomment-2295065634
278            Ok(result.ty)
279        }
280    }
281
282    pub(crate) fn unpack(&mut self, infcx: &mut InferCtxt) {
283        self.bindings
284            .fmap_mut(|_loc, ty| infcx.hoister(true).hoist(ty));
285    }
286
287    pub(crate) fn unblock(&mut self, infcx: &mut InferCtxt, place: &Place) {
288        self.bindings.unblock(infcx, place);
289    }
290
291    pub(crate) fn check_goto(
292        self,
293        infcx: &mut InferCtxtAt,
294        bb_env: &BasicBlockEnv,
295        target: BasicBlock,
296    ) -> InferResult {
297        infcx.ensure_resolved_evars(|infcx| {
298            let bb_env = bb_env
299                .data
300                .replace_bound_refts_with(|sort, mode, _| infcx.fresh_infer_var(sort, mode));
301
302            // Check constraints
303            for constr in &bb_env.constrs {
304                infcx.check_pred(constr, ConstrReason::Goto(target));
305            }
306
307            // Check subtyping
308            let bb_env = bb_env.bindings.flatten();
309            for (path, _, ty2) in bb_env {
310                let ty1 = self.bindings.get(&path);
311                infcx.subtyping(&ty1.unblocked(), &ty2.unblocked(), ConstrReason::Goto(target))?;
312            }
313            Ok(())
314        })
315    }
316
317    pub(crate) fn fold(&mut self, infcx: &mut InferCtxtAt, place: &Place) -> InferResult {
318        let span = infcx.span;
319        self.bindings.lookup(place, span).fold(infcx)?;
320        Ok(())
321    }
322
323    pub(crate) fn unfold_local_ptr(
324        &mut self,
325        infcx: &mut InferCtxt,
326        bound: &Ty,
327    ) -> InferResult<Loc> {
328        let name = infcx.define_unknown_var(&Sort::Loc);
329        let loc = Loc::from(name);
330        let ty = infcx.unpack(bound);
331        self.bindings
332            .insert(loc, LocKind::LocalPtr(bound.clone()), ty);
333        Ok(loc)
334    }
335
336    /// ```text
337    /// -----------------------------------
338    /// Γ ; &strg <ℓ: t> => Γ,ℓ: t ; ptr(ℓ)
339    /// ```
340    pub(crate) fn unfold_strg_ref(
341        &mut self,
342        infcx: &mut InferCtxt,
343        path: &Path,
344        ty: &Ty,
345    ) -> InferResult<Loc> {
346        if let Some(loc) = path.to_loc() {
347            let ty = infcx.unpack(ty);
348            self.bindings.insert(loc, LocKind::Universal, ty);
349            Ok(loc)
350        } else {
351            bug!("unfold_strg_ref: unexpected path {path:?}")
352        }
353    }
354
355    pub(crate) fn unfold(
356        &mut self,
357        infcx: &mut InferCtxt,
358        place: &Place,
359        span: Span,
360    ) -> InferResult {
361        self.bindings.unfold(infcx, place, span)
362    }
363
364    pub(crate) fn downcast(
365        &mut self,
366        infcx: &mut InferCtxtAt,
367        place: &Place,
368        variant_idx: VariantIdx,
369    ) -> InferResult {
370        let mut down_place = place.clone();
371        let span = infcx.span;
372        down_place
373            .projection
374            .push(PlaceElem::Downcast(None, variant_idx));
375        self.bindings.unfold(infcx, &down_place, span)?;
376        Ok(())
377    }
378
379    pub fn fully_resolve_evars(&mut self, infcx: &InferCtxt) {
380        self.bindings
381            .fmap_mut(|_loc, ty| infcx.fully_resolve_evars(ty));
382    }
383}
384
385pub(crate) enum PtrToRefBound {
386    Ty(Ty),
387    Infer,
388    Identity,
389}
390
391impl flux_infer::infer::LocEnv for TypeEnv<'_> {
392    fn ptr_to_ref(
393        &mut self,
394        infcx: &mut InferCtxtAt,
395        reason: ConstrReason,
396        re: Region,
397        path: &Path,
398        bound: Ty,
399    ) -> InferResult<Ty> {
400        self.ptr_to_ref(infcx, reason, re, path, PtrToRefBound::Ty(bound))
401    }
402
403    fn get(&self, path: &Path) -> Ty {
404        self.get(path)
405    }
406
407    fn unfold_strg_ref(&mut self, infcx: &mut InferCtxt, path: &Path, ty: &Ty) -> InferResult<Loc> {
408        self.unfold_strg_ref(infcx, path, ty)
409    }
410
411    fn unfold_local_ptr(&mut self, infcx: &mut InferCtxt, bound: &Ty) -> InferResult<Loc> {
412        self.unfold_local_ptr(infcx, bound)
413    }
414
415    fn fold_local_ptrs(&mut self, infcx: &mut InferCtxtAt) -> InferResult {
416        self.fold_local_ptrs(infcx)
417    }
418
419    fn update_path(&mut self, path: &Path, new_ty: Ty, span: Span) {
420        self.update_path(path, new_ty, span);
421    }
422}
423
424impl BasicBlockEnvShape {
425    pub fn enter<'a>(&self, local_decls: &'a LocalDecls) -> TypeEnv<'a> {
426        TypeEnv { bindings: self.bindings.clone(), local_decls }
427    }
428
429    fn new(scope: Scope, env: TypeEnv) -> BasicBlockEnvShape {
430        let mut bindings = env.bindings;
431        bindings.fmap_mut(|_loc, ty| BasicBlockEnvShape::pack_ty(&scope, ty));
432        BasicBlockEnvShape { scope, bindings }
433    }
434
435    fn pack_ty(scope: &Scope, ty: &Ty) -> Ty {
436        match ty.kind() {
437            TyKind::Indexed(bty, idxs) => {
438                let bty = BasicBlockEnvShape::pack_bty(scope, bty);
439                if scope.has_free_vars(idxs) {
440                    Ty::exists_with_constr(bty, Expr::hole(HoleKind::Pred))
441                } else {
442                    Ty::indexed(bty, idxs.clone())
443                }
444            }
445            TyKind::Downcast(adt, args, ty, variant, fields) => {
446                debug_assert!(!scope.has_free_vars(args));
447                debug_assert!(!scope.has_free_vars(ty));
448                let fields = fields.iter().map(|ty| Self::pack_ty(scope, ty)).collect();
449                Ty::downcast(adt.clone(), args.clone(), ty.clone(), *variant, fields)
450            }
451            TyKind::Blocked(ty) => Ty::blocked(BasicBlockEnvShape::pack_ty(scope, ty)),
452            // FIXME(nilehmann) [`TyKind::Exists`] could also contain free variables.
453            TyKind::Exists(_)
454            | TyKind::Discr(..)
455            | TyKind::Ptr(..)
456            | TyKind::Uninit
457            | TyKind::Param(_)
458            | TyKind::Constr(_, _) => ty.clone(),
459            TyKind::Infer(_) => bug!("unexpected hole whecn checking function body"),
460            TyKind::StrgRef(..) => bug!("unexpected strong reference when checking function body"),
461        }
462    }
463
464    fn pack_bty(scope: &Scope, bty: &BaseTy) -> BaseTy {
465        match bty {
466            BaseTy::Adt(adt_def, args) => {
467                let args = List::from_vec(
468                    args.iter()
469                        .map(|arg| Self::pack_generic_arg(scope, arg))
470                        .collect(),
471                );
472                BaseTy::adt(adt_def.clone(), args)
473            }
474            BaseTy::FnDef(def_id, args) => {
475                let args = List::from_vec(
476                    args.iter()
477                        .map(|arg| Self::pack_generic_arg(scope, arg))
478                        .collect(),
479                );
480                BaseTy::fn_def(*def_id, args)
481            }
482            BaseTy::Tuple(tys) => {
483                let tys = tys
484                    .iter()
485                    .map(|ty| BasicBlockEnvShape::pack_ty(scope, ty))
486                    .collect();
487                BaseTy::Tuple(tys)
488            }
489            BaseTy::Slice(ty) => BaseTy::Slice(Self::pack_ty(scope, ty)),
490            BaseTy::Ref(r, ty, mutbl) => BaseTy::Ref(*r, Self::pack_ty(scope, ty), *mutbl),
491            BaseTy::Array(ty, c) => BaseTy::Array(Self::pack_ty(scope, ty), c.clone()),
492            BaseTy::Int(_)
493            | BaseTy::Param(_)
494            | BaseTy::Uint(_)
495            | BaseTy::Bool
496            | BaseTy::Float(_)
497            | BaseTy::Str
498            | BaseTy::RawPtr(_, _)
499            | BaseTy::RawPtrMetadata(_)
500            | BaseTy::Char
501            | BaseTy::Never
502            | BaseTy::Closure(..)
503            | BaseTy::Dynamic(..)
504            | BaseTy::Alias(..)
505            | BaseTy::FnPtr(..)
506            | BaseTy::Foreign(..)
507            | BaseTy::Coroutine(..) => {
508                if scope.has_free_vars(bty) {
509                    tracked_span_bug!("unexpected type with free vars")
510                } else {
511                    bty.clone()
512                }
513            }
514            BaseTy::Pat => {
515                todo!()
516            }
517            BaseTy::Infer(..) => {
518                tracked_span_bug!("unexpected infer type")
519            }
520        }
521    }
522
523    fn pack_generic_arg(scope: &Scope, arg: &GenericArg) -> GenericArg {
524        match arg {
525            GenericArg::Ty(ty) => GenericArg::Ty(Self::pack_ty(scope, ty)),
526            GenericArg::Base(ctor) => GenericArg::Base(Self::pack_subset_ty_ctor(scope, ctor)),
527            GenericArg::Lifetime(re) => GenericArg::Lifetime(*re),
528            GenericArg::Const(c) => GenericArg::Const(c.clone()),
529        }
530    }
531
532    fn pack_subset_ty_ctor(scope: &Scope, ctor: &SubsetTyCtor) -> SubsetTyCtor {
533        let sty = ctor.as_ref().skip_binder();
534        debug_assert!(sty.idx.is_nu());
535        let bty = Self::pack_bty(scope, &sty.bty);
536        let pred = if scope.has_free_vars(&sty.pred) {
537            Expr::hole(HoleKind::Pred)
538        } else {
539            sty.pred.clone()
540        };
541        let sort = bty.sort();
542        Binder::bind_with_sort(SubsetTy::new(bty, Expr::nu(), pred), sort)
543    }
544
545    fn update(&mut self, path: &Path, ty: Ty, span: Span) {
546        self.bindings.lookup(path, span).update(ty);
547    }
548
549    /// join(self, genv, other) consumes the bindings in other, to "update"
550    /// `self` in place, and returns `true` if there was an actual change
551    /// or `false` indicating no change (i.e., a fixpoint was reached).
552    pub(crate) fn join(&mut self, other: TypeEnv, span: Span) -> bool {
553        let paths = self.bindings.paths();
554
555        // Join types
556        let mut modified = false;
557        for path in &paths {
558            let ty1 = self.bindings.get(path);
559            let ty2 = other.bindings.get(path);
560            let ty = if ty1 == ty2 { ty1.clone() } else { self.join_ty(&ty1, &ty2) };
561            modified |= ty1 != ty;
562            self.update(path, ty, span);
563        }
564
565        modified
566    }
567
568    fn join_ty(&self, ty1: &Ty, ty2: &Ty) -> Ty {
569        match (ty1.kind(), ty2.kind()) {
570            (TyKind::Blocked(ty1), _) => Ty::blocked(self.join_ty(ty1, &ty2.unblocked())),
571            (_, TyKind::Blocked(ty2)) => Ty::blocked(self.join_ty(&ty1.unblocked(), ty2)),
572            (TyKind::Uninit, _) | (_, TyKind::Uninit) => Ty::uninit(),
573            (TyKind::Exists(ty1), _) => self.join_ty(ty1.as_ref().skip_binder(), ty2),
574            (_, TyKind::Exists(ty2)) => self.join_ty(ty1, ty2.as_ref().skip_binder()),
575            (TyKind::Constr(_, ty1), _) => self.join_ty(ty1, ty2),
576            (_, TyKind::Constr(_, ty2)) => self.join_ty(ty1, ty2),
577            (TyKind::Indexed(bty1, idx1), TyKind::Indexed(bty2, idx2)) => {
578                let bty = self.join_bty(bty1, bty2);
579                let mut sorts = vec![];
580                let idx = self.join_idx(idx1, idx2, &bty.sort(), &mut sorts);
581                if sorts.is_empty() {
582                    Ty::indexed(bty, idx)
583                } else {
584                    let ty = Ty::constr(Expr::hole(HoleKind::Pred), Ty::indexed(bty, idx));
585                    Ty::exists(Binder::bind_with_sorts(ty, &sorts))
586                }
587            }
588            (TyKind::Ptr(rk1, path1), TyKind::Ptr(rk2, path2)) => {
589                debug_assert_eq!(rk1, rk2);
590                debug_assert_eq!(path1, path2);
591                Ty::ptr(*rk1, path1.clone())
592            }
593            (TyKind::Param(param_ty1), TyKind::Param(param_ty2)) => {
594                debug_assert_eq!(param_ty1, param_ty2);
595                Ty::param(*param_ty1)
596            }
597            (
598                TyKind::Downcast(adt1, args1, ty1, variant1, fields1),
599                TyKind::Downcast(adt2, args2, ty2, variant2, fields2),
600            ) => {
601                debug_assert_eq!(adt1, adt2);
602                debug_assert_eq!(args1, args2);
603                debug_assert!(ty1 == ty2 && !self.scope.has_free_vars(ty2));
604                debug_assert_eq!(variant1, variant2);
605                debug_assert_eq!(fields1.len(), fields2.len());
606                let fields = iter::zip(fields1, fields2)
607                    .map(|(ty1, ty2)| self.join_ty(ty1, ty2))
608                    .collect();
609                Ty::downcast(adt1.clone(), args1.clone(), ty1.clone(), *variant1, fields)
610            }
611            _ => tracked_span_bug!("unexpected types: `{ty1:?}` - `{ty2:?}`"),
612        }
613    }
614
615    fn join_idx(&self, e1: &Expr, e2: &Expr, sort: &Sort, bound_sorts: &mut Vec<Sort>) -> Expr {
616        match (e1.kind(), e2.kind(), sort) {
617            (ExprKind::Tuple(es1), ExprKind::Tuple(es2), Sort::Tuple(sorts)) => {
618                debug_assert_eq3!(es1.len(), es2.len(), sorts.len());
619                Expr::tuple(
620                    izip!(es1, es2, sorts)
621                        .map(|(e1, e2, sort)| self.join_idx(e1, e2, sort, bound_sorts))
622                        .collect(),
623                )
624            }
625            (
626                ExprKind::Ctor(Ctor::Struct(_), flds1),
627                ExprKind::Ctor(Ctor::Struct(_), flds2),
628                Sort::App(SortCtor::Adt(sort_def), args),
629            ) => {
630                let sorts = sort_def.struct_variant().field_sorts(args);
631                debug_assert_eq3!(flds1.len(), flds2.len(), sorts.len());
632
633                Expr::ctor_struct(
634                    sort_def.did(),
635                    izip!(flds1, flds2, &sorts)
636                        .map(|(f1, f2, sort)| self.join_idx(f1, f2, sort, bound_sorts))
637                        .collect(),
638                )
639            }
640            _ => {
641                let has_free_vars2 = self.scope.has_free_vars(e2);
642                let has_escaping_vars1 = e1.has_escaping_bvars();
643                let has_escaping_vars2 = e2.has_escaping_bvars();
644                if !has_free_vars2 && !has_escaping_vars1 && !has_escaping_vars2 && e1 == e2 {
645                    e1.clone()
646                } else if sort.is_pred() {
647                    // FIXME(nilehmann) we shouldn't special case predicates here. Instead, we
648                    // should differentiate between generics and indices.
649                    let fsort = sort.expect_func().expect_mono();
650                    Expr::abs(Lambda::bind_with_fsort(Expr::hole(HoleKind::Pred), fsort))
651                } else {
652                    bound_sorts.push(sort.clone());
653                    Expr::bvar(
654                        INNERMOST,
655                        BoundVar::from_usize(bound_sorts.len() - 1),
656                        BoundReftKind::Anon,
657                    )
658                }
659            }
660        }
661    }
662
663    fn join_bty(&self, bty1: &BaseTy, bty2: &BaseTy) -> BaseTy {
664        match (bty1, bty2) {
665            (BaseTy::Adt(def1, args1), BaseTy::Adt(def2, args2)) => {
666                tracked_span_dbg_assert_eq!(def1.did(), def2.did());
667                let args = iter::zip(args1, args2)
668                    .map(|(arg1, arg2)| self.join_generic_arg(arg1, arg2))
669                    .collect();
670                BaseTy::adt(def1.clone(), List::from_vec(args))
671            }
672            (BaseTy::Tuple(fields1), BaseTy::Tuple(fields2)) => {
673                let fields = iter::zip(fields1, fields2)
674                    .map(|(ty1, ty2)| self.join_ty(ty1, ty2))
675                    .collect();
676                BaseTy::Tuple(fields)
677            }
678            (BaseTy::Alias(alias_ty1), BaseTy::Alias(alias_ty2)) => {
679                tracked_span_dbg_assert_eq!(alias_ty1, alias_ty2);
680                BaseTy::Alias(alias_ty1.clone())
681            }
682            (BaseTy::Ref(r1, ty1, mutbl1), BaseTy::Ref(r2, ty2, mutbl2)) => {
683                tracked_span_dbg_assert_eq!(r1, r2);
684                tracked_span_dbg_assert_eq!(mutbl1, mutbl2);
685                BaseTy::Ref(*r1, self.join_ty(ty1, ty2), *mutbl1)
686            }
687            (BaseTy::Array(ty1, len1), BaseTy::Array(ty2, len2)) => {
688                tracked_span_dbg_assert_eq!(len1, len2);
689                BaseTy::Array(self.join_ty(ty1, ty2), len1.clone())
690            }
691            (BaseTy::Slice(ty1), BaseTy::Slice(ty2)) => BaseTy::Slice(self.join_ty(ty1, ty2)),
692            (BaseTy::FnPtr(sig1), BaseTy::FnPtr(sig2)) if sig1 != sig2 => {
693                // Generalize to a signature with holes, which are replaced by kvars in the basic
694                // block env, e.g., `fn(i32{v: $k0(v)}) -> i32{v: $k1(v)}`. Fn subtyping (against
695                // both signatures) is checked when jumping to the join point.
696                BaseTy::FnPtr(generalize_fn_sig(sig1))
697            }
698            _ => {
699                tracked_span_dbg_assert_eq!(bty1, bty2);
700                bty1.clone()
701            }
702        }
703    }
704
705    fn join_generic_arg(&self, arg1: &GenericArg, arg2: &GenericArg) -> GenericArg {
706        match (arg1, arg2) {
707            (GenericArg::Ty(ty1), GenericArg::Ty(ty2)) => GenericArg::Ty(self.join_ty(ty1, ty2)),
708            (GenericArg::Base(ctor1), GenericArg::Base(ctor2)) => {
709                let sty1 = ctor1.as_ref().skip_binder();
710                let sty2 = ctor2.as_ref().skip_binder();
711                debug_assert!(sty1.idx.is_nu());
712                debug_assert!(sty2.idx.is_nu());
713
714                let bty = self.join_bty(&sty1.bty, &sty2.bty);
715                let pred = if self.scope.has_free_vars(&sty2.pred) || sty1.pred != sty2.pred {
716                    Expr::hole(HoleKind::Pred)
717                } else {
718                    sty1.pred.clone()
719                };
720                let sort = bty.sort();
721                let ctor = Binder::bind_with_sort(SubsetTy::new(bty, Expr::nu(), pred), sort);
722                GenericArg::Base(ctor)
723            }
724            (GenericArg::Lifetime(re1), GenericArg::Lifetime(_re2)) => {
725                // TODO(nilehmann) loop_abstract_refinement.rs is triggering this assertion to fail
726                // wee should fix it.
727                // debug_assert_eq!(re1, _re2);
728                GenericArg::Lifetime(*re1)
729            }
730            (GenericArg::Const(c1), GenericArg::Const(c2)) => {
731                debug_assert_eq!(c1, c2);
732                GenericArg::Const(c1.clone())
733            }
734            _ => tracked_span_bug!("unexpected generic args: `{arg1:?}` - `{arg2:?}`"),
735        }
736    }
737
738    pub fn into_bb_env(self, infcx: &mut InferCtxtRoot, body: &Body) -> BasicBlockEnv {
739        let mut delegate = LocalHoister::default();
740        let mut hoister = Hoister::with_delegate(&mut delegate).transparent();
741
742        let mut bindings = self.bindings;
743        bindings.fmap_mut(|loc, ty| {
744            let name = if let Loc::Local(local) = loc {
745                body.local_names.get(local).copied()
746            } else {
747                None
748            };
749            hoister.delegate.name = name;
750            hoister.hoist(ty)
751        });
752
753        BasicBlockEnv {
754            // We are relying on all the types in `bindings` not having escaping bvars, otherwise
755            // we would have to shift them in since we are creating a new binder.
756            data: delegate.bind(|vars, preds| {
757                // Replace all holes with a single fresh kvar on all parameters
758                let mut constrs = preds
759                    .into_iter()
760                    .filter(|pred| !matches!(pred.kind(), ExprKind::Hole(HoleKind::Pred)))
761                    .collect_vec();
762                let kvar = infcx.fresh_kvar_in_scope(
763                    std::slice::from_ref(&vars),
764                    &self.scope,
765                    KVarEncoding::Conj,
766                );
767                constrs.push(kvar);
768
769                // Replace remaining holes by fresh kvars
770                let mut kvar_gen = |binders: &[_], kind| {
771                    debug_assert_eq!(kind, HoleKind::Pred);
772                    let binders = std::iter::once(vars.clone())
773                        .chain(binders.iter().cloned())
774                        .collect_vec();
775                    infcx.fresh_kvar_in_scope(&binders, &self.scope, KVarEncoding::Conj)
776                };
777                bindings.fmap_mut(|_, binding| binding.replace_holes(&mut kvar_gen));
778
779                BasicBlockEnvData { constrs: constrs.into(), bindings }
780            }),
781            scope: self.scope,
782        }
783    }
784}
785
786/// Generalizes a fn pointer signature to one with the same (rust) shape where all refinements are
787/// holes, e.g., `for<n> fn(i32[n]) -> i32[n + 1]` becomes `fn({v. i32[v] | *}) -> {v. i32[v] | *}`.
788/// The refinement params of the signature are dropped, and so are the `requires`, `ensures` and
789/// `no_panic` (which could mention them). The region vars are kept: they come before the
790/// refinement params in the binder, so dropping the latter doesn't shift them.
791///
792/// This is the same template that refining the rust signature with [`Refiner::with_holes`] would
793/// produce (see the `Refine` impl for `ty::FnSig`), but it doesn't need a refiner.
794///
795/// [`Refiner::with_holes`]: flux_middle::rty::refining::Refiner::with_holes
796fn generalize_fn_sig(sig: &PolyFnSig) -> PolyFnSig {
797    let vars = sig
798        .vars()
799        .iter()
800        .filter(|var| matches!(var, BoundVariableKind::Region(_)))
801        .cloned()
802        .collect();
803    let fn_sig = sig.skip_binder_ref();
804    let ret = fn_sig.output().skip_binder().ret.with_holes();
805    let fn_sig = FnSig::new(
806        fn_sig.safety,
807        fn_sig.abi,
808        List::empty(),
809        fn_sig.inputs.with_holes(),
810        Binder::dummy(FnOutput::new(ret, vec![])),
811        Expr::ff(),
812        fn_sig.lifted,
813    );
814    Binder::bind_with_vars(fn_sig, vars)
815}
816
817impl TypeVisitable for BasicBlockEnvData {
818    fn visit_with<V: TypeVisitor>(&self, _visitor: &mut V) -> ControlFlow<V::BreakTy> {
819        unimplemented!()
820    }
821}
822
823impl TypeFoldable for BasicBlockEnvData {
824    fn try_fold_with<F: FallibleTypeFolder>(
825        &self,
826        folder: &mut F,
827    ) -> std::result::Result<Self, F::Error> {
828        Ok(BasicBlockEnvData {
829            constrs: self.constrs.try_fold_with(folder)?,
830            bindings: self.bindings.try_fold_with(folder)?,
831        })
832    }
833}
834
835impl BasicBlockEnv {
836    pub(crate) fn enter<'a>(
837        &self,
838        infcx: &mut InferCtxt,
839        local_decls: &'a LocalDecls,
840    ) -> TypeEnv<'a> {
841        let data = self.data.replace_bound_refts_with(|sort, _, kind| {
842            Expr::fvar(infcx.define_bound_reft_var(sort, kind))
843        });
844        for constr in &data.constrs {
845            infcx.assume_pred(constr);
846        }
847        TypeEnv { bindings: data.bindings, local_decls }
848    }
849
850    pub(crate) fn scope(&self) -> &Scope {
851        &self.scope
852    }
853}
854
855mod pretty {
856    use std::fmt;
857
858    use flux_middle::pretty::*;
859
860    use super::*;
861
862    impl Pretty for TypeEnv<'_> {
863        fn fmt(&self, cx: &PrettyCx, f: &mut fmt::Formatter<'_>) -> fmt::Result {
864            w!(cx, f, "{:?}", &self.bindings)
865        }
866
867        fn default_cx(tcx: TyCtxt) -> PrettyCx {
868            PlacesTree::default_cx(tcx)
869        }
870    }
871
872    impl Pretty for BasicBlockEnvShape {
873        fn fmt(&self, cx: &PrettyCx, f: &mut fmt::Formatter<'_>) -> fmt::Result {
874            w!(cx, f, "{:?} {:?}", &self.scope, &self.bindings)
875        }
876
877        fn default_cx(tcx: TyCtxt) -> PrettyCx {
878            PlacesTree::default_cx(tcx)
879        }
880    }
881
882    impl Pretty for BasicBlockEnv {
883        fn fmt(&self, cx: &PrettyCx, f: &mut fmt::Formatter<'_>) -> fmt::Result {
884            w!(cx, f, "{:?} ", &self.scope)?;
885
886            let vars = self.data.vars();
887            cx.with_bound_vars(vars, || {
888                if !vars.is_empty() {
889                    cx.fmt_bound_vars(true, "for<", vars, "> ", f)?;
890                }
891                let data = self.data.as_ref().skip_binder();
892                if !data.constrs.is_empty() {
893                    w!(
894                        cx,
895                        f,
896                        "{:?} ⇒ ",
897                        join!(", ", data.constrs.iter().filter(|pred| !pred.is_trivially_true()))
898                    )?;
899                }
900                w!(cx, f, "{:?}", &data.bindings)
901            })
902        }
903
904        fn default_cx(tcx: TyCtxt) -> PrettyCx {
905            PlacesTree::default_cx(tcx)
906        }
907    }
908
909    impl_debug_with_default_cx! {
910        TypeEnv<'_> => "type_env",
911        BasicBlockEnvShape => "basic_block_env_shape",
912        BasicBlockEnv => "basic_block_env"
913    }
914}
915
916/// A very explicit representation of [`TypeEnv`] for debugging/tracing/serialization ONLY.
917#[derive(Serialize, DebugAsJson)]
918pub struct TypeEnvTrace(Vec<TypeEnvBind>);
919
920#[derive(Serialize)]
921struct TypeEnvBind {
922    local: LocInfo,
923    name: Option<String>,
924    kind: String,
925    ty: String,
926    span: Option<SpanTrace>,
927}
928
929#[derive(Serialize)]
930enum LocInfo {
931    Local(String),
932    Var(String),
933}
934
935fn loc_info(loc: &Loc) -> LocInfo {
936    match loc {
937        Loc::Local(local) => LocInfo::Local(format!("{local:?}")),
938        Loc::Var(var) => LocInfo::Var(format!("{var:?}")),
939    }
940}
941
942fn loc_name(local_names: &UnordMap<Local, Symbol>, loc: &Loc) -> Option<String> {
943    if let Loc::Local(local) = loc {
944        let name = local_names.get(local)?;
945        return Some(format!("{name}"));
946    }
947    None
948}
949
950fn loc_span(
951    genv: GlobalEnv,
952    local_decls: &IndexVec<Local, LocalDecl>,
953    loc: &Loc,
954) -> Option<SpanTrace> {
955    if let Loc::Local(local) = loc {
956        return local_decls
957            .get(*local)
958            .map(|local_decl| SpanTrace::new(genv.tcx(), local_decl.source_info.span));
959    }
960    None
961}
962
963impl TypeEnvTrace {
964    pub fn new(
965        genv: GlobalEnv,
966        local_names: &UnordMap<Local, Symbol>,
967        local_decls: &IndexVec<Local, LocalDecl>,
968        cx: PrettyCx,
969        env: &TypeEnv,
970    ) -> Self {
971        let mut bindings = vec![];
972        env.bindings
973            .iter()
974            .filter(|(_, binding)| !binding.ty.is_uninit())
975            .sorted_by(|(loc1, _), (loc2, _)| loc1.cmp(loc2))
976            .for_each(|(loc, binding)| {
977                let name = loc_name(local_names, loc);
978                let local = loc_info(loc);
979                let kind = format!("{:?}", binding.kind);
980                let ty = binding.ty.nested_string(&cx);
981                let span = loc_span(genv, local_decls, loc);
982                bindings.push(TypeEnvBind { name, local, kind, ty, span });
983            });
984
985        TypeEnvTrace(bindings)
986    }
987}