1use 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 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 a_binders: u32,
99 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 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 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 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 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 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 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}