Skip to main content

flux_fhir_analysis/conv/
struct_compat.rs

1//! Check whether two refinemnt types/signatures are structurally compatible.
2//!
3//! Used to check if a user spec is compatible with the underlying rust type. The code also
4//! infer types annotated with `_` in the surface syntax.
5
6use std::{fmt, iter};
7
8use flux_common::bug;
9use flux_errors::Errors;
10use flux_middle::{
11    def_id::MaybeExternId,
12    fhir,
13    global_env::GlobalEnv,
14    queries::QueryResult,
15    rty::{
16        self,
17        fold::{TypeFoldable, TypeFolder, TypeSuperFoldable},
18        refining::{Refine as _, Refiner},
19    },
20};
21use flux_rustc_bridge::ty::{self, FieldIdx, VariantIdx};
22use rustc_ast::Mutability;
23use rustc_data_structures::unord::UnordMap;
24use rustc_type_ir::{DebruijnIndex, INNERMOST, InferConst};
25
26pub(crate) fn type_alias(
27    genv: GlobalEnv,
28    alias: &fhir::TyAlias,
29    alias_ty: &rty::TyCtor,
30    def_id: MaybeExternId,
31) -> QueryResult<rty::TyCtor> {
32    let rust_ty = genv.lower_type_of(def_id.resolved_id())?.skip_binder();
33    let expected = rust_ty.refine(&Refiner::default_for_item(genv, def_id.resolved_id())?)?;
34    let mut zipper = Zipper::new(genv, def_id);
35
36    if zipper
37        .enter_a_binder(alias_ty, |zipper, ty| zipper.zip_ty(ty, &expected))
38        .is_err()
39    {
40        zipper
41            .errors
42            .emit(errors::IncompatibleRefinement::type_alias(genv, def_id, alias));
43    }
44
45    zipper.errors.to_result()?;
46
47    Ok(zipper.holes.replace_holes(alias_ty))
48}
49
50pub(crate) fn fn_sig(
51    genv: GlobalEnv,
52    decl: &fhir::FnDecl,
53    fn_sig: &rty::PolyFnSig,
54    def_id: MaybeExternId,
55) -> QueryResult<rty::PolyFnSig> {
56    let rust_fn_sig = genv.lower_fn_sig(def_id.resolved_id())?.skip_binder();
57
58    let expected = Refiner::default_for_item(genv, def_id.resolved_id())?.refine(&rust_fn_sig)?;
59
60    let mut zipper = Zipper::new(genv, def_id);
61    if let Err(err) = zipper.zip_poly_fn_sig(fn_sig, &expected) {
62        zipper.emit_fn_sig_err(err, decl);
63    }
64
65    zipper.errors.to_result()?;
66
67    Ok(zipper.holes.replace_holes(fn_sig))
68}
69
70pub(crate) fn variants(
71    genv: GlobalEnv,
72    variants: &[rty::PolyVariant],
73    adt_def_id: MaybeExternId,
74) -> QueryResult<Vec<rty::PolyVariant>> {
75    let refiner = Refiner::default_for_item(genv, adt_def_id.resolved_id())?;
76    let mut zipper = Zipper::new(genv, adt_def_id);
77    // TODO check same number of variants
78    for (i, variant) in variants.iter().enumerate() {
79        let variant_idx = VariantIdx::from_usize(i);
80        let expected = refiner.refine_variant_def(adt_def_id.resolved_id(), variant_idx)?;
81        zipper.zip_variant(variant, &expected, variant_idx);
82    }
83
84    zipper.errors.to_result()?;
85
86    Ok(variants
87        .iter()
88        .map(|v| zipper.holes.replace_holes(v))
89        .collect())
90}
91
92struct Zipper<'genv, 'tcx> {
93    genv: GlobalEnv<'genv, 'tcx>,
94    owner_id: MaybeExternId,
95    locs: UnordMap<rty::Loc, rty::Ty>,
96    holes: Holes,
97    /// Number of binders we've entered in `a`
98    a_binders: u32,
99    /// Each element in the vector correspond to a binder in `b`. For some binders we map it to
100    /// a corresponding binder in `a`. We assume that expressions filling holes will only contain
101    /// variables pointing to some of these mapped binders.
102    b_binder_to_a_binder: Vec<Option<u32>>,
103    errors: Errors<'genv>,
104}
105
106#[derive(Default)]
107struct Holes {
108    sorts: UnordMap<rty::SortVid, rty::Sort>,
109    subset_tys: UnordMap<rty::TyVid, rty::SubsetTy>,
110    types: UnordMap<rty::TyVid, rty::Ty>,
111    regions: UnordMap<rty::RegionVid, rty::Region>,
112    consts: UnordMap<rty::ConstVid, rty::Const>,
113}
114
115impl TypeFolder for &Holes {
116    fn fold_sort(&mut self, sort: &rty::Sort) -> rty::Sort {
117        if let rty::Sort::Infer(vid) = sort {
118            self.sorts
119                .get(vid)
120                .cloned()
121                .unwrap_or_else(|| bug!("unfilled sort hole {vid:?}"))
122        } else {
123            sort.super_fold_with(self)
124        }
125    }
126
127    fn fold_ty(&mut self, ty: &rty::Ty) -> rty::Ty {
128        if let rty::TyKind::Infer(vid) = ty.kind() {
129            self.types
130                .get(vid)
131                .cloned()
132                .unwrap_or_else(|| bug!("unfilled type hole {vid:?}"))
133        } else {
134            ty.super_fold_with(self)
135        }
136    }
137
138    fn fold_subset_ty(&mut self, constr: &rty::SubsetTy) -> rty::SubsetTy {
139        if let rty::BaseTy::Infer(vid) = &constr.bty {
140            self.subset_tys
141                .get(vid)
142                .cloned()
143                .unwrap_or_else(|| bug!("unfilled type hole {vid:?}"))
144        } else {
145            constr.super_fold_with(self)
146        }
147    }
148
149    fn fold_region(&mut self, r: &rty::Region) -> rty::Region {
150        if let rty::Region::ReVar(vid) = r {
151            self.regions
152                .get(vid)
153                .copied()
154                .unwrap_or_else(|| bug!("unfilled region hole {vid:?}"))
155        } else {
156            *r
157        }
158    }
159
160    fn fold_const(&mut self, ct: &rty::Const) -> rty::Const {
161        if let rty::ConstKind::Infer(InferConst::Var(cid)) = ct.kind {
162            self.consts
163                .get(&cid)
164                .cloned()
165                .unwrap_or_else(|| bug!("unfilled const hole {cid:?}"))
166        } else {
167            ct.super_fold_with(self)
168        }
169    }
170}
171
172impl Holes {
173    fn replace_holes<T: TypeFoldable>(&self, t: &T) -> T {
174        let mut this = self;
175        t.fold_with(&mut this)
176    }
177}
178
179impl<'genv, 'tcx> Zipper<'genv, 'tcx> {
180    fn new(genv: GlobalEnv<'genv, 'tcx>, owner_id: MaybeExternId) -> Self {
181        Self {
182            genv,
183            owner_id,
184            locs: UnordMap::default(),
185            holes: Default::default(),
186            a_binders: 0,
187            b_binder_to_a_binder: vec![],
188            errors: Errors::new(genv.sess()),
189        }
190    }
191
192    fn is_async_fn(&self) -> bool {
193        self.genv
194            .tcx()
195            .asyncness(self.owner_id.resolved_id())
196            .is_async()
197    }
198
199    fn zip_poly_fn_sig(&mut self, a: &rty::PolyFnSig, b: &rty::PolyFnSig) -> Result<(), FnSigErr> {
200        self.enter_binders(a, b, |this, a, b| this.zip_fn_sig(a, b))
201    }
202
203    fn zip_variant(&mut self, a: &rty::PolyVariant, b: &rty::PolyVariant, variant_idx: VariantIdx) {
204        self.enter_binders(a, b, |this, a, b| {
205            // The args are always `GenericArgs::identity_for_item` inside the `EarlyBinder`
206            debug_assert_eq!(a.args, b.args);
207
208            if a.fields.len() != b.fields.len() {
209                this.errors.emit(errors::FieldCountMismatch::new(
210                    this.genv,
211                    a.fields.len(),
212                    this.owner_id,
213                    variant_idx,
214                ));
215                return;
216            }
217            for (i, (ty_a, ty_b)) in iter::zip(&a.fields, &b.fields).enumerate() {
218                let field_idx = FieldIdx::from_usize(i);
219                if this.zip_ty(ty_a, ty_b).is_err() {
220                    this.errors.emit(errors::IncompatibleRefinement::field(
221                        this.genv,
222                        this.owner_id,
223                        variant_idx,
224                        field_idx,
225                    ));
226                }
227            }
228        });
229    }
230
231    fn zip_fn_sig(&mut self, a: &rty::FnSig, b: &rty::FnSig) -> Result<(), FnSigErr> {
232        if a.inputs().len() != b.inputs().len() {
233            Err(FnSigErr::ArgCountMismatch)?;
234        }
235        for (i, (ty_a, ty_b)) in iter::zip(a.inputs(), b.inputs()).enumerate() {
236            self.zip_ty(ty_a, ty_b).map_err(|_| FnSigErr::FnInput(i))?;
237        }
238        self.enter_binders(&a.output, &b.output, |this, output_a, output_b| {
239            this.zip_output(output_a, output_b)
240        })
241    }
242
243    fn zip_output(&mut self, a: &rty::FnOutput, b: &rty::FnOutput) -> Result<(), FnSigErr> {
244        self.zip_ty(&a.ret, &b.ret).map_err(FnSigErr::FnOutput)?;
245
246        for (i, ensures) in a.ensures.iter().enumerate() {
247            if let rty::Ensures::Type(path, ty_a) = ensures {
248                let loc = path.to_loc().unwrap();
249                let ty_b = self.locs.get(&loc).unwrap().shift_in_escaping(1);
250                self.zip_ty(ty_a, &ty_b)
251                    .map_err(|_| FnSigErr::Ensures { i, expected: ty_b })?;
252            }
253        }
254        Ok(())
255    }
256
257    fn zip_ty(&mut self, a: &rty::Ty, b: &rty::Ty) -> Result<(), Mismatch> {
258        match (a.kind(), b.kind()) {
259            (rty::TyKind::Infer(vid), _) => {
260                assert_ne!(vid.as_u32(), 0);
261                let b = self.adjust_bvars(b);
262                self.holes.types.insert(*vid, b);
263                Ok(())
264            }
265            (rty::TyKind::Exists(ctor_a), _) => {
266                self.enter_a_binder(ctor_a, |this, ty_a| this.zip_ty(ty_a, b))
267            }
268            (_, rty::TyKind::Exists(ctor_b)) => {
269                self.enter_b_binder(ctor_b, |this, ty_b| this.zip_ty(a, ty_b))
270            }
271            (rty::TyKind::Constr(_, ty_a), _) => self.zip_ty(ty_a, b),
272            (_, rty::TyKind::Constr(_, ty_b)) => self.zip_ty(a, ty_b),
273            (rty::TyKind::Indexed(bty_a, _), rty::TyKind::Indexed(bty_b, _)) => {
274                self.zip_bty(bty_a, bty_b)
275            }
276            (rty::TyKind::StrgRef(re_a, path, ty_a), rty::Ref!(re_b, ty_b, Mutability::Mut)) => {
277                let loc = path.to_loc().unwrap();
278                self.locs.insert(loc, ty_b.clone());
279
280                self.zip_region(re_a, re_b);
281                self.zip_ty(ty_a, ty_b)
282            }
283            (rty::TyKind::Param(pty_a), rty::TyKind::Param(pty_b)) => {
284                assert_eq_or_incompatible(pty_a, pty_b)
285            }
286            (
287                rty::TyKind::Ptr(_, _)
288                | rty::TyKind::Discr(..)
289                | rty::TyKind::Downcast(_, _, _, _, _)
290                | rty::TyKind::Blocked(_)
291                | rty::TyKind::Uninit,
292                _,
293            ) => {
294                bug!("unexpected type {a:?}");
295            }
296            _ => Err(Mismatch::new(a, b)),
297        }
298    }
299
300    fn zip_bty(&mut self, a: &rty::BaseTy, b: &rty::BaseTy) -> Result<(), Mismatch> {
301        match (a, b) {
302            (rty::BaseTy::Int(ity_a), rty::BaseTy::Int(ity_b)) => {
303                assert_eq_or_incompatible(ity_a, ity_b)
304            }
305            (rty::BaseTy::Uint(uity_a), rty::BaseTy::Uint(uity_b)) => {
306                assert_eq_or_incompatible(uity_a, uity_b)
307            }
308            (rty::BaseTy::Bool, rty::BaseTy::Bool) => Ok(()),
309            (rty::BaseTy::Str, rty::BaseTy::Str) => Ok(()),
310            (rty::BaseTy::Char, rty::BaseTy::Char) => Ok(()),
311            (rty::BaseTy::Float(fty_a), rty::BaseTy::Float(fty_b)) => {
312                assert_eq_or_incompatible(fty_a, fty_b)
313            }
314            (rty::BaseTy::Slice(ty_a), rty::BaseTy::Slice(ty_b)) => self.zip_ty(ty_a, ty_b),
315            (rty::BaseTy::Adt(adt_def_a, args_a), rty::BaseTy::Adt(adt_def_b, args_b)) => {
316                assert_eq_or_incompatible(adt_def_a.did(), adt_def_b.did())?;
317                assert_eq_or_incompatible(args_a.len(), args_b.len())?;
318                for (arg_a, arg_b) in iter::zip(args_a, args_b) {
319                    self.zip_generic_arg(arg_a, arg_b)?;
320                }
321                Ok(())
322            }
323            (rty::BaseTy::RawPtr(ty_a, mutbl_a), rty::BaseTy::RawPtr(ty_b, mutbl_b)) => {
324                assert_eq_or_incompatible(mutbl_a, mutbl_b)?;
325                self.zip_ty(ty_a, ty_b)
326            }
327            (rty::BaseTy::Ref(re_a, ty_a, mutbl_a), rty::BaseTy::Ref(re_b, ty_b, mutbl_b)) => {
328                assert_eq_or_incompatible(mutbl_a, mutbl_b)?;
329                self.zip_region(re_a, re_b);
330                self.zip_ty(ty_a, ty_b)
331            }
332            (rty::BaseTy::FnPtr(poly_sig_a), rty::BaseTy::FnPtr(poly_sig_b)) => {
333                // Check that safety and abi of a fn-ptr in a spec (see `desugar_bare_fn`)
334                // matches the rust one, to not refine an `unsafe fn` pointer with a safe one.
335                let (sig_a, sig_b) = (poly_sig_a.skip_binder_ref(), poly_sig_b.skip_binder_ref());
336                assert_eq_or_incompatible(sig_a.safety, sig_b.safety)?;
337                assert_eq_or_incompatible(sig_a.abi, sig_b.abi)?;
338                self.zip_poly_fn_sig(poly_sig_a, poly_sig_b)
339                    .map_err(|_| Mismatch::new(poly_sig_a, poly_sig_b))
340            }
341            (rty::BaseTy::Tuple(tys_a), rty::BaseTy::Tuple(tys_b)) => {
342                assert_eq_or_incompatible(tys_a.len(), tys_b.len())?;
343                for (ty_a, ty_b) in iter::zip(tys_a, tys_b) {
344                    self.zip_ty(ty_a, ty_b)?;
345                }
346                Ok(())
347            }
348            (rty::BaseTy::Alias(aty_a), rty::BaseTy::Alias(aty_b)) => {
349                assert_eq_or_incompatible(aty_a.kind, aty_b.kind)?;
350                assert_eq_or_incompatible(aty_a.args.len(), aty_b.args.len())?;
351                for (arg_a, arg_b) in iter::zip(&aty_a.args, &aty_b.args) {
352                    self.zip_generic_arg(arg_a, arg_b)?;
353                }
354                Ok(())
355            }
356            (rty::BaseTy::Array(ty_a, len_a), rty::BaseTy::Array(ty_b, len_b)) => {
357                self.zip_const(len_a, len_b)?;
358                self.zip_ty(ty_a, ty_b)
359            }
360            (rty::BaseTy::Never, rty::BaseTy::Never) => Ok(()),
361            (rty::BaseTy::Param(pty_a), rty::BaseTy::Param(pty_b)) => {
362                assert_eq_or_incompatible(pty_a, pty_b)
363            }
364            (rty::BaseTy::Dynamic(preds_a, re_a), rty::BaseTy::Dynamic(preds_b, re_b)) => {
365                assert_eq_or_incompatible(preds_a.len(), preds_b.len())?;
366                for (pred_a, pred_b) in iter::zip(preds_a, preds_b) {
367                    self.zip_poly_existential_pred(pred_a, pred_b)?;
368                }
369                self.zip_region(re_a, re_b);
370                Ok(())
371            }
372            (rty::BaseTy::Foreign(def_id_a), rty::BaseTy::Foreign(def_id_b)) => {
373                assert_eq_or_incompatible(def_id_a, def_id_b)
374            }
375            (rty::BaseTy::Closure(..) | rty::BaseTy::Coroutine(..), _) => {
376                bug!("unexpected type `{a:?}`");
377            }
378            _ => Err(Mismatch::new(a, b)),
379        }
380    }
381
382    fn zip_generic_arg(
383        &mut self,
384        a: &rty::GenericArg,
385        b: &rty::GenericArg,
386    ) -> Result<(), Mismatch> {
387        match (a, b) {
388            (rty::GenericArg::Ty(ty_a), rty::GenericArg::Ty(ty_b)) => self.zip_ty(ty_a, ty_b),
389            (rty::GenericArg::Base(ctor_a), rty::GenericArg::Base(ctor_b)) => {
390                self.zip_sorts(&ctor_a.sort(), &ctor_b.sort());
391                self.enter_binders(ctor_a, ctor_b, |this, sty_a, sty_b| {
392                    this.zip_subset_ty(sty_a, sty_b)
393                })
394            }
395            (rty::GenericArg::Lifetime(re_a), rty::GenericArg::Lifetime(re_b)) => {
396                self.zip_region(re_a, re_b);
397                Ok(())
398            }
399            (rty::GenericArg::Const(ct_a), rty::GenericArg::Const(ct_b)) => {
400                self.zip_const(ct_a, ct_b)
401            }
402            _ => Err(Mismatch::new(a, b)),
403        }
404    }
405
406    fn zip_sorts(&mut self, a: &rty::Sort, b: &rty::Sort) {
407        if let rty::Sort::Infer(vid) = a {
408            assert_ne!(vid.as_u32(), 0);
409            self.holes.sorts.insert(*vid, b.clone());
410        }
411    }
412
413    fn zip_subset_ty(&mut self, a: &rty::SubsetTy, b: &rty::SubsetTy) -> Result<(), Mismatch> {
414        if let rty::BaseTy::Infer(vid) = a.bty {
415            assert_ne!(vid.as_u32(), 0);
416            let b = self.adjust_bvars(b);
417            self.holes.subset_tys.insert(vid, b);
418            Ok(())
419        } else {
420            self.zip_bty(&a.bty, &b.bty)
421        }
422    }
423
424    fn zip_const(&mut self, a: &rty::Const, b: &ty::Const) -> Result<(), Mismatch> {
425        match (&a.kind, &b.kind) {
426            (rty::ConstKind::Infer(ty::InferConst::Var(cid)), _) => {
427                self.holes.consts.insert(*cid, b.clone());
428                Ok(())
429            }
430            (rty::ConstKind::Param(param_const_a), ty::ConstKind::Param(param_const_b)) => {
431                assert_eq_or_incompatible(param_const_a, param_const_b)
432            }
433            (rty::ConstKind::Value(ty_a, val_a), ty::ConstKind::Value(ty_b, val_b)) => {
434                assert_eq_or_incompatible(ty_a, ty_b)?;
435                assert_eq_or_incompatible(val_a, val_b)
436            }
437            (rty::ConstKind::Alias(c1), ty::ConstKind::Alias(c2)) => {
438                assert_eq_or_incompatible(c1, c2)
439            }
440            _ => Err(Mismatch::new(a, b)),
441        }
442    }
443
444    fn zip_region(&mut self, a: &rty::Region, b: &ty::Region) {
445        if let rty::Region::ReVar(vid) = a {
446            let re = self.adjust_bvars(b);
447            self.holes.regions.insert(*vid, re);
448        }
449    }
450
451    fn zip_poly_existential_pred(
452        &mut self,
453        a: &rty::Binder<rty::ExistentialPredicate>,
454        b: &rty::Binder<rty::ExistentialPredicate>,
455    ) -> Result<(), Mismatch> {
456        self.enter_binders(a, b, |this, a, b| {
457            match (a, b) {
458                (
459                    rty::ExistentialPredicate::Trait(trait_ref_a),
460                    rty::ExistentialPredicate::Trait(trait_ref_b),
461                ) => {
462                    assert_eq_or_incompatible(trait_ref_a.def_id, trait_ref_b.def_id)?;
463                    assert_eq_or_incompatible(trait_ref_a.args.len(), trait_ref_b.args.len())?;
464                    for (arg_a, arg_b) in iter::zip(&trait_ref_a.args, &trait_ref_b.args) {
465                        this.zip_generic_arg(arg_a, arg_b)?;
466                    }
467                    Ok(())
468                }
469                (
470                    rty::ExistentialPredicate::Projection(projection_a),
471                    rty::ExistentialPredicate::Projection(projection_b),
472                ) => {
473                    assert_eq_or_incompatible(projection_a.def_id, projection_b.def_id)?;
474                    assert_eq_or_incompatible(projection_a.args.len(), projection_b.args.len())?;
475                    for (arg_a, arg_b) in iter::zip(&projection_a.args, &projection_b.args) {
476                        this.zip_generic_arg(arg_a, arg_b)?;
477                    }
478                    this.enter_binders(&projection_a.term, &projection_b.term, |this, a, b| {
479                        this.zip_bty(&a.bty, &b.bty)
480                    })
481                }
482                (
483                    rty::ExistentialPredicate::AutoTrait(def_id_a),
484                    rty::ExistentialPredicate::AutoTrait(def_id_b),
485                ) => assert_eq_or_incompatible(def_id_a, def_id_b),
486                _ => Err(Mismatch::new(a, b)),
487            }
488        })
489    }
490
491    /// Enter a binder in both `a` and `b` creating a mapping between the two.
492    fn enter_binders<T, R>(
493        &mut self,
494        a: &rty::Binder<T>,
495        b: &rty::Binder<T>,
496        f: impl FnOnce(&mut Self, &T, &T) -> R,
497    ) -> R {
498        self.b_binder_to_a_binder.push(Some(self.a_binders));
499        self.a_binders += 1;
500        let r = f(self, a.skip_binder_ref(), b.skip_binder_ref());
501        self.a_binders -= 1;
502        self.b_binder_to_a_binder.pop();
503        r
504    }
505
506    /// Enter a binder in `a` without a corresponding mapping in `b`
507    fn enter_a_binder<T, R>(
508        &mut self,
509        t: &rty::Binder<T>,
510        f: impl FnOnce(&mut Self, &T) -> R,
511    ) -> R {
512        self.a_binders += 1;
513        let r = f(self, t.skip_binder_ref());
514        self.a_binders -= 1;
515        r
516    }
517
518    /// Enter a binder in `b` without a corresponding mapping in `a`
519    fn enter_b_binder<T, R>(
520        &mut self,
521        t: &rty::Binder<T>,
522        f: impl FnOnce(&mut Self, &T) -> R,
523    ) -> R {
524        self.b_binder_to_a_binder.push(None);
525        let r = f(self, t.skip_binder_ref());
526        self.b_binder_to_a_binder.pop();
527        r
528    }
529
530    fn adjust_bvars<T: TypeFoldable + Clone + std::fmt::Debug>(&self, t: &T) -> T {
531        struct Adjuster<'a, 'genv, 'tcx> {
532            current_index: DebruijnIndex,
533            zipper: &'a Zipper<'genv, 'tcx>,
534        }
535
536        impl Adjuster<'_, '_, '_> {
537            fn adjust(&self, debruijn: DebruijnIndex) -> DebruijnIndex {
538                let b_binders = self.zipper.b_binder_to_a_binder.len();
539                let mapped_binder = self.zipper.b_binder_to_a_binder
540                    [b_binders - debruijn.as_usize() - 1]
541                    .unwrap_or_else(|| {
542                        bug!("bound var without corresponding binder: `{debruijn:?}`")
543                    });
544                DebruijnIndex::from_u32(self.zipper.a_binders - mapped_binder - 1)
545                    .shifted_in(self.current_index.as_u32())
546            }
547        }
548
549        impl TypeFolder for Adjuster<'_, '_, '_> {
550            fn enter_binder(&mut self, _: &rty::BoundVariableKinds) {
551                self.current_index.shift_in(1);
552            }
553
554            fn exit_binder(&mut self) {
555                self.current_index.shift_out(1);
556            }
557
558            fn fold_region(&mut self, re: &rty::Region) -> rty::Region {
559                if let rty::ReBound(debruijn, br) = *re
560                    && debruijn >= self.current_index
561                {
562                    rty::ReBound(self.adjust(debruijn), br)
563                } else {
564                    *re
565                }
566            }
567
568            fn fold_expr(&mut self, expr: &rty::Expr) -> rty::Expr {
569                if let rty::ExprKind::Var(rty::Var::Bound(debruijn, breft)) = expr.kind()
570                    && *debruijn >= self.current_index
571                {
572                    rty::Expr::bvar(self.adjust(*debruijn), breft.var, breft.kind)
573                } else {
574                    expr.super_fold_with(self)
575                }
576            }
577        }
578        t.fold_with(&mut Adjuster { current_index: INNERMOST, zipper: self })
579    }
580
581    fn emit_fn_sig_err(&mut self, err: FnSigErr, decl: &fhir::FnDecl) {
582        match err {
583            FnSigErr::ArgCountMismatch => {
584                self.errors.emit(errors::IncompatibleParamCount::new(
585                    self.genv,
586                    decl,
587                    self.owner_id,
588                ));
589            }
590            FnSigErr::FnInput(i) => {
591                self.errors.emit(errors::IncompatibleRefinement::fn_input(
592                    self.genv,
593                    self.owner_id,
594                    decl,
595                    i,
596                ));
597            }
598            FnSigErr::FnOutput(_) => {
599                self.errors.emit(errors::IncompatibleRefinement::fn_output(
600                    self.genv,
601                    self.owner_id,
602                    decl,
603                    self.is_async_fn(),
604                ));
605            }
606            FnSigErr::Ensures { i, expected } => {
607                self.errors.emit(errors::IncompatibleRefinement::ensures(
608                    self.genv,
609                    self.owner_id,
610                    decl,
611                    &expected,
612                    i,
613                ));
614            }
615        }
616    }
617}
618
619fn assert_eq_or_incompatible<T: Eq + fmt::Debug>(a: T, b: T) -> Result<(), Mismatch> {
620    if a != b {
621        return Err(Mismatch::new(a, b));
622    }
623    Ok(())
624}
625
626#[expect(dead_code, reason = "we use the the String for debugging")]
627struct Mismatch(String);
628
629impl Mismatch {
630    fn new<T: fmt::Debug>(a: T, b: T) -> Self {
631        Self(format!("{a:?} != {b:?}"))
632    }
633}
634
635enum FnSigErr {
636    ArgCountMismatch,
637    FnInput(usize),
638    #[expect(dead_code, reason = "we use the struct for debugging")]
639    FnOutput(Mismatch),
640    Ensures {
641        i: usize,
642        expected: rty::Ty,
643    },
644}
645
646mod errors {
647    use flux_common::span_bug;
648    use flux_errors::E0999;
649    use flux_macros::Diagnostic;
650    use flux_middle::{def_id::MaybeExternId, fhir, global_env::GlobalEnv, rty};
651    use flux_rustc_bridge::{
652        ToRustc,
653        ty::{FieldIdx, VariantIdx},
654    };
655    use rustc_span::{DUMMY_SP, Span};
656
657    #[derive(Diagnostic)]
658    #[diag("{$def_descr} has an incompatible refinement annotation", code = E0999)]
659    #[note("a refinement annotation must match the unrefined definition structurally")]
660    pub(super) struct IncompatibleRefinement<'tcx> {
661        #[primary_span]
662        #[label("expected a refinement of `{$expected_ty}`")]
663        span: Span,
664        #[label("unrefined {$def_descr} found here")]
665        expected_span: Option<Span>,
666        expected_ty: rustc_middle::ty::Ty<'tcx>,
667        def_descr: &'static str,
668        #[help("mark the flux signature as `async fn`")]
669        async_hint: Option<()>,
670    }
671
672    impl<'tcx> IncompatibleRefinement<'tcx> {
673        pub(super) fn type_alias(
674            genv: GlobalEnv<'_, 'tcx>,
675            def_id: MaybeExternId,
676            type_alias: &fhir::TyAlias,
677        ) -> Self {
678            let tcx = genv.tcx();
679            Self {
680                span: type_alias.ty.span,
681                def_descr: tcx.def_descr(def_id.resolved_id()),
682                expected_span: Some(tcx.def_span(def_id)),
683                expected_ty: tcx.type_of(def_id).skip_binder(),
684                async_hint: None,
685            }
686        }
687
688        pub(super) fn fn_input(
689            genv: GlobalEnv<'_, 'tcx>,
690            fn_id: MaybeExternId,
691            decl: &fhir::FnDecl,
692            pos: usize,
693        ) -> Self {
694            let expected_span = match fn_id {
695                MaybeExternId::Local(local_id) => {
696                    genv.tcx()
697                        .hir_node_by_def_id(local_id)
698                        .fn_decl()
699                        .and_then(|fn_decl| fn_decl.inputs.get(pos))
700                        .map(|input| input.span)
701                }
702                MaybeExternId::Extern(_, extern_id) => Some(genv.tcx().def_span(extern_id)),
703            };
704
705            let expected_ty = genv
706                .tcx()
707                .fn_sig(fn_id.resolved_id())
708                .skip_binder()
709                .inputs()
710                .map_bound(|inputs| inputs[pos])
711                .skip_binder();
712
713            Self {
714                span: decl.inputs[pos].span,
715                def_descr: genv.tcx().def_descr(fn_id.resolved_id()),
716                expected_span,
717                expected_ty,
718                async_hint: None,
719            }
720        }
721
722        pub(super) fn fn_output(
723            genv: GlobalEnv<'_, 'tcx>,
724            fn_id: MaybeExternId,
725            decl: &fhir::FnDecl,
726            is_async: bool,
727        ) -> Self {
728            let expected_span = match fn_id {
729                MaybeExternId::Local(local_id) => {
730                    genv.tcx()
731                        .hir_node_by_def_id(local_id)
732                        .fn_decl()
733                        .map(|fn_decl| fn_decl.output.span())
734                }
735                MaybeExternId::Extern(_, extern_id) => Some(genv.tcx().def_span(extern_id)),
736            };
737
738            let expected_ty = genv
739                .tcx()
740                .fn_sig(fn_id.resolved_id())
741                .skip_binder()
742                .output()
743                .skip_binder();
744            let spec_span = decl.output.ret.span;
745
746            let async_hint = if is_async { Some(()) } else { None };
747
748            Self {
749                span: spec_span,
750                def_descr: genv.tcx().def_descr(fn_id.resolved_id()),
751                expected_span,
752                expected_ty,
753                async_hint,
754            }
755        }
756
757        pub(super) fn ensures(
758            genv: GlobalEnv<'_, 'tcx>,
759            fn_id: MaybeExternId,
760            decl: &fhir::FnDecl,
761            expected: &rty::Ty,
762            i: usize,
763        ) -> Self {
764            let fhir::Ensures::Type(_, ty) = &decl.output.ensures[i] else {
765                span_bug!(decl.span, "expected `fhir::Ensures::Type`");
766            };
767            let tcx = genv.tcx();
768            Self {
769                span: ty.span,
770                def_descr: tcx.def_descr(fn_id.resolved_id()),
771                expected_span: None,
772                expected_ty: expected.to_rustc(tcx),
773                async_hint: None,
774            }
775        }
776
777        pub(super) fn field(
778            genv: GlobalEnv<'_, 'tcx>,
779            adt_id: MaybeExternId,
780            variant_idx: VariantIdx,
781            field_idx: FieldIdx,
782        ) -> Self {
783            let tcx = genv.tcx();
784            let adt_def = tcx.adt_def(adt_id);
785            let field_def = &adt_def.variant(variant_idx).fields[field_idx];
786
787            let item = genv.fhir_expect_item(adt_id.local_id()).unwrap();
788            let span = match &item.kind {
789                fhir::ItemKind::Enum(enum_def) => {
790                    enum_def.variants[variant_idx.as_usize()].fields[field_idx.as_usize()]
791                        .ty
792                        .span
793                }
794                fhir::ItemKind::Struct(struct_def)
795                    if let fhir::StructKind::Transparent { fields } = &struct_def.kind =>
796                {
797                    fields[field_idx.as_usize()].ty.span
798                }
799                _ => DUMMY_SP,
800            };
801
802            Self {
803                span,
804                def_descr: tcx.def_descr(field_def.did),
805                expected_span: Some(tcx.def_span(field_def.did)),
806                expected_ty: tcx.type_of(field_def.did).skip_binder(),
807                async_hint: None,
808            }
809        }
810    }
811
812    #[derive(Diagnostic)]
813    #[diag("{$def_descr} has an incompatible refinement annotation", code = E0999)]
814    pub(super) struct IncompatibleParamCount {
815        #[primary_span]
816        #[label(
817            "refined signature has {$found} {$found ->
818                [one] parameter
819                *[other] parameters
820            }"
821        )]
822        span: Span,
823        found: usize,
824        #[label(
825            "unrefined signature has {$expected} {$expected ->
826                [one] parameter
827                *[other] parameters
828            }"
829        )]
830        expected_span: Span,
831        expected: usize,
832        def_descr: &'static str,
833    }
834
835    impl IncompatibleParamCount {
836        pub(super) fn new(genv: GlobalEnv, decl: &fhir::FnDecl, def_id: MaybeExternId) -> Self {
837            let def_descr = genv.tcx().def_descr(def_id.resolved_id());
838
839            let span = if !decl.inputs.is_empty() {
840                decl.inputs[decl.inputs.len() - 1]
841                    .span
842                    .with_lo(decl.inputs[0].span.lo())
843            } else {
844                decl.span
845            };
846
847            let expected_span = if let Some(local_id) = def_id.as_local()
848                && let expected_decl = genv.tcx().hir_node_by_def_id(local_id).fn_decl().unwrap()
849                && !expected_decl.inputs.is_empty()
850            {
851                expected_decl.inputs[expected_decl.inputs.len() - 1]
852                    .span
853                    .with_lo(expected_decl.inputs[0].span.lo())
854            } else {
855                genv.tcx().def_span(def_id)
856            };
857
858            let expected = genv
859                .tcx()
860                .fn_sig(def_id)
861                .skip_binder()
862                .skip_binder()
863                .inputs()
864                .len();
865
866            Self { span, found: decl.inputs.len(), expected_span, expected, def_descr }
867        }
868    }
869
870    #[derive(Diagnostic)]
871    #[diag("variant has an incompatible refinement annotation", code = E0999)]
872    pub(super) struct FieldCountMismatch {
873        #[primary_span]
874        #[label(
875            "expected {$expected_fields} {$expected_fields ->
876                [one] field
877                *[other] fields
878            }, found {$fields}"
879        )]
880        span: Span,
881        fields: usize,
882        #[label("unrefined variant defined here")]
883        expected_span: Span,
884        expected_fields: usize,
885    }
886
887    impl FieldCountMismatch {
888        pub(super) fn new(
889            genv: GlobalEnv,
890            found: usize,
891            adt_def_id: MaybeExternId,
892            variant_idx: VariantIdx,
893        ) -> Self {
894            let adt_def = genv.tcx().adt_def(adt_def_id);
895            let expected_variant = adt_def.variant(variant_idx);
896
897            // Get the span of the variant if this is an enum. Structs cannot have produce a field
898            // count mismatch.
899            let span = if let Ok(fhir::Node::Item(item)) = genv.fhir_node(adt_def_id.local_id())
900                && let fhir::ItemKind::Enum(enum_def) = &item.kind
901                && let Some(variant) = enum_def.variants.get(variant_idx.as_usize())
902            {
903                variant.span
904            } else {
905                DUMMY_SP
906            };
907
908            Self {
909                span,
910                fields: found,
911                expected_span: genv.tcx().def_span(expected_variant.def_id),
912                expected_fields: expected_variant.fields.len(),
913            }
914        }
915    }
916}