1use flux_common::tracked_span_bug;
2use flux_middle::{
3 big_int::BigInt,
4 rty::{self, Binder, EarlyReftParam, InternalFuncKind, List, SpecFuncKind},
5};
6use flux_rustc_bridge::lowering::Lower;
7use itertools::Itertools;
8use rustc_hir::{def::DefKind, def_id::DefId};
9use rustc_type_ir::BoundVar;
10
11use super::{ConstKey, FixpointCtxt, fixpoint};
12use crate::fixpoint_encoding::FixpointSolution;
13
14impl<'genv, 'tcx, Tag> FixpointCtxt<'genv, 'tcx, Tag>
15where
16 Tag: std::hash::Hash + Eq + Copy,
17{
18 pub(crate) fn fixpoint_to_solution(
19 &mut self,
20 sol: &FixpointSolution,
21 ) -> rty::Binder<rty::Expr> {
22 let mut vars = vec![];
23 let mut sorts = vec![];
24 for (var, sort) in &sol.0 {
25 let fixpoint::Var::Local(local_var) = var else {
26 tracked_span_bug!("encountered non-local variable in binder: {var:?}");
27 };
28 vars.push(*local_var);
29 sorts.push(
30 self.fixpoint_to_sort(sort)
31 .unwrap_or_else(|_| tracked_span_bug!("failed to parse sort: {sort:?}")),
32 );
33 }
34 self.ecx.local_var_env.push_layer(vars);
35 let expr = self
36 .fixpoint_to_expr(&sol.1)
37 .unwrap_or_else(|err| tracked_span_bug!("failed to convert expr: {err:?}"));
38 self.ecx.local_var_env.pop_layer();
39 rty::Binder::bind_with_sorts(expr, &sorts)
40 }
41
42 fn fixpoint_to_sort_ctor(
43 &self,
44 ctor: &fixpoint::SortCtor,
45 ) -> Result<rty::SortCtor, FixpointParseError> {
46 match ctor {
47 fixpoint::SortCtor::Set => Ok(rty::SortCtor::Set),
48 fixpoint::SortCtor::Map => Ok(rty::SortCtor::Map),
49 fixpoint::SortCtor::Data(fixpoint::DataSort::Tuple(_)) => {
50 panic!("oh no! tuple!") }
52 fixpoint::SortCtor::Data(fixpoint::DataSort::User(opaque_id)) => {
53 let def_id = self.scx.opaque_sorts[opaque_id.as_usize()];
54 Ok(rty::SortCtor::User(def_id))
55 }
56 fixpoint::SortCtor::Data(fixpoint::DataSort::Adt(adt_id)) => {
57 let def_id = self.scx.adt_sorts[adt_id.as_usize()];
58 let Ok(adt_sort_def) = self.genv.adt_sort_def_of(def_id) else {
59 return Err(FixpointParseError::UnknownAdt(def_id));
60 };
61 Ok(rty::SortCtor::Adt(adt_sort_def))
62 }
63 }
64 }
65
66 pub(crate) fn fixpoint_to_sort(
67 &self,
68 fsort: &fixpoint::Sort,
69 ) -> Result<rty::Sort, FixpointParseError> {
70 match fsort {
71 fixpoint::Sort::Int => Ok(rty::Sort::Int),
72 fixpoint::Sort::Real => Ok(rty::Sort::Real),
73 fixpoint::Sort::Bool => Ok(rty::Sort::Bool),
74 fixpoint::Sort::Str => Ok(rty::Sort::Str),
75 fixpoint::Sort::Func(sorts) => {
76 let sort1 = self.fixpoint_to_sort(&sorts[0])?;
77 let sort2 = self.fixpoint_to_sort(&sorts[1])?;
78 let fsort = rty::FuncSort::new(vec![sort1], sort2);
79 let poly_sort = rty::PolyFuncSort::new(List::empty(), fsort);
80 Ok(rty::Sort::Func(poly_sort))
81 }
82 fixpoint::Sort::App(ctor, args) => {
83 let ctor = self.fixpoint_to_sort_ctor(ctor)?;
84 let args = args
85 .iter()
86 .map(|fsort| self.fixpoint_to_sort(fsort))
87 .try_collect()?;
88 Ok(rty::Sort::App(ctor, args))
89 }
90 fixpoint::Sort::BitVec(fsort) if let fixpoint::Sort::BvSize(size) = **fsort => {
91 Ok(rty::Sort::BitVec(rty::BvSize::Fixed(size)))
92 }
93 _ => unimplemented!("fixpoint_to_sort: {fsort:?}"),
94 }
95 }
96
97 fn is_curried_primop_app(
98 &mut self,
99 fhead: &fixpoint::Expr,
100 fargs: &[fixpoint::Expr],
101 op_args: &mut Vec<fixpoint::Expr>,
102 ) -> Option<rty::BinOp> {
103 match fhead {
104 fixpoint::Expr::Var(fixpoint::Var::Global(global_var, _))
105 | fixpoint::Expr::Var(fixpoint::Var::Const(global_var, _)) => {
106 if let Some(ConstKey::PrimOp(bin_op)) =
107 self.ecx.const_env.const_map_rev.get(global_var)
108 {
109 op_args.reverse();
110 Some(bin_op.clone())
111 } else {
112 None
113 }
114 }
115 fixpoint::Expr::App(fhead_inner, _, fargs_inner, _) => {
116 if fargs.len() == 1 {
117 op_args.push(fargs[0].clone());
118 }
119 self.is_curried_primop_app(fhead_inner, fargs_inner, op_args)
120 }
121 _ => None,
122 }
123 }
124
125 #[allow(dead_code)]
126 pub(crate) fn fixpoint_to_expr(
127 &mut self,
128 fexpr: &fixpoint::Expr,
129 ) -> Result<rty::Expr, FixpointParseError> {
130 match fexpr {
131 fixpoint::Expr::Constant(constant) => {
132 let c = match constant {
133 fixpoint::Constant::Numeral(num) => rty::Constant::Int(BigInt::from(*num)),
134 fixpoint::Constant::Real(dec) => rty::Constant::Real(rty::Real(dec.0)),
135 fixpoint::Constant::Boolean(b) => rty::Constant::Bool(*b),
136 fixpoint::Constant::String(s) => rty::Constant::Str(s.0),
137 fixpoint::Constant::BitVec(bv, size) => rty::Constant::BitVec(*bv, *size),
138 };
139 Ok(rty::Expr::constant(c))
140 }
141 fixpoint::Expr::Var(fvar) => {
142 match fvar {
143 fixpoint::Var::Underscore => {
144 unreachable!("Underscore should not appear in exprs")
145 }
146 fixpoint::Var::Global(global_var, _) | fixpoint::Var::Const(global_var, _) => {
147 if let Some(const_key) = self.ecx.const_env.const_map_rev.get(global_var) {
148 match const_key {
149 ConstKey::RustConst(def_id) => Ok(rty::Expr::const_def_id(*def_id)),
150 ConstKey::Alias(_flux_id, _args) => {
151 unreachable!("Should be special-cased as the head of an app")
152 }
153 ConstKey::Lambda(lambda) => Ok(rty::Expr::abs(lambda.clone())),
154 ConstKey::PrimOp(bin_op) => {
155 Ok(rty::Expr::internal_func(InternalFuncKind::Rel(
156 bin_op.clone(),
157 )))
158 }
159 ConstKey::Cast(_sort, _sort1) => {
160 unreachable!(
161 "Should be specially handled as the head of a function app."
162 )
163 }
164 ConstKey::WKVar(_, _) => {
165 unreachable!("Weak kvars are not global vars");
166 }
167 }
168 } else {
169 Err(FixpointParseError::NoGlobalVar(*global_var))
170 }
171 }
172 fixpoint::Var::Local(fname) => {
173 if let Some(expr) = self.ecx.local_var_env.reverse_map.get(fname) {
174 return Ok(expr.clone());
175 }
176
177 for (depth, layer) in self.ecx.local_var_env.layers.iter().rev().enumerate()
178 {
179 for (idx, var) in layer.iter().enumerate() {
180 if fname == var {
181 return Ok(rty::Expr::bvar(
182 rty::DebruijnIndex::from_usize(depth),
183 BoundVar::from_usize(idx),
184 rty::BoundReftKind::Anon,
185 ));
186 }
187 }
188 }
189
190 Ok(rty::Expr::fvar(rty::Name::from_u32(fname.as_u32())))
199 }
201 fixpoint::Var::DataCtor(adt_id, variant_idx) => {
202 let def_id = self.scx.adt_sorts[adt_id.as_usize()];
203 Ok(rty::Expr::ctor_enum(def_id, *variant_idx))
204 }
205 fixpoint::Var::TupleCtor { .. }
206 | fixpoint::Var::TupleProj { .. }
207 | fixpoint::Var::DataProj { .. }
208 | fixpoint::Var::UIFRel(_) => {
209 unreachable!(
210 "Trying to convert an atomic var, but reached a var that should only occur as the head of an app (and be special-cased in conversion as a result)"
211 )
212 }
213 fixpoint::Var::Param(EarlyReftParam { index, name }) => {
214 Ok(rty::Expr::early_param(*index, *name))
215 }
216 fixpoint::Var::ConstGeneric(const_generic) => {
217 Ok(rty::Expr::const_generic(*const_generic))
218 }
219 fixpoint::Var::WKVar(..) => {
220 unreachable!(
221 "Weak kvar ids should be converted as part of fixpoint::Expr::WKVar"
222 );
223 }
224 }
225 }
226 fixpoint::Expr::App(fhead, _sort_args, fargs, _out_sort) => {
227 let mut op_args = vec![];
228 if let Some(bin_op) = self.is_curried_primop_app(fhead, fargs, &mut op_args) {
229 if op_args.len() != 2 {
230 return Err(FixpointParseError::PrimOpArityMismatch(fargs.len()));
231 } else {
232 let e1 = self.fixpoint_to_expr(&op_args[0])?;
233 let e2 = self.fixpoint_to_expr(&op_args[1])?;
234 return Ok(rty::Expr::prim_val(bin_op, e1, e2));
235 }
236 }
237 match &**fhead {
238 fixpoint::Expr::Var(fixpoint::Var::TupleProj { arity, field }) => {
239 if fargs.len() == 1 {
240 let earg = self.fixpoint_to_expr(&fargs[0])?;
241 Ok(rty::Expr::field_proj(
242 earg,
243 rty::FieldProj::Tuple { arity: *arity, field: *field },
244 ))
245 } else {
246 Err(FixpointParseError::ProjArityMismatch(fargs.len()))
247 }
248 }
249 fixpoint::Expr::Var(fixpoint::Var::DataProj { adt_id, field }) => {
250 if fargs.len() == 1 {
251 let earg = self.fixpoint_to_expr(&fargs[0])?;
252 Ok(rty::Expr::field_proj(
253 earg,
254 rty::FieldProj::Adt {
255 def_id: self.scx.adt_sorts[adt_id.as_usize()],
256 field: *field,
257 },
258 ))
259 } else {
260 Err(FixpointParseError::ProjArityMismatch(fargs.len()))
261 }
262 }
263 fixpoint::Expr::Var(fixpoint::Var::TupleCtor { arity }) => {
264 if fargs.len() == *arity {
265 let eargs = fargs
266 .iter()
267 .map(|farg| self.fixpoint_to_expr(farg))
268 .try_collect()?;
269 Ok(rty::Expr::tuple(eargs))
270 } else {
271 Err(FixpointParseError::TupleCtorArityMismatch(*arity, fargs.len()))
272 }
273 }
274 fixpoint::Expr::Var(fixpoint::Var::DataCtor(adt_id, field)) => {
275 let eargs = fargs
276 .iter()
277 .map(|farg| self.fixpoint_to_expr(farg))
278 .try_collect()?;
279 let def_id = self.scx.adt_sorts[adt_id.as_usize()];
280 match self.genv.tcx().def_kind(def_id) {
281 DefKind::Struct => Ok(rty::Expr::ctor_struct(def_id, eargs)),
282 DefKind::Enum => {
283 let ctor = rty::Ctor::Enum(def_id, *field);
284 Ok(rty::Expr::ctor(ctor, eargs))
285 }
286 _ => Err(FixpointParseError::InvalidDefKindForCtor),
287 }
288 }
289 fixpoint::Expr::Var(fixpoint::Var::UIFRel(fbinrel)) => {
290 if fargs.len() == 2 {
291 let e1 = self.fixpoint_to_expr(&fargs[0])?;
292 let e2 = self.fixpoint_to_expr(&fargs[1])?;
293 let binrel = match fbinrel {
294 fixpoint::BinRel::Eq => rty::BinOp::Eq,
295 fixpoint::BinRel::Ne => rty::BinOp::Ne,
296 fixpoint::BinRel::Gt => rty::BinOp::Gt(rty::Sort::Str),
303 fixpoint::BinRel::Ge => rty::BinOp::Ge(rty::Sort::Str),
304 fixpoint::BinRel::Lt => rty::BinOp::Lt(rty::Sort::Str),
305 fixpoint::BinRel::Le => rty::BinOp::Le(rty::Sort::Str),
306 };
307 Ok(rty::Expr::binary_op(binrel, e1, e2))
308 } else {
309 Err(FixpointParseError::UIFRelArityMismatch(fargs.len()))
310 }
311 }
312 fixpoint::Expr::Var(fixpoint::Var::Global(global_var, _))
313 | fixpoint::Expr::Var(fixpoint::Var::Const(global_var, _)) => {
314 if let Some(const_key) = self.ecx.const_env.const_map_rev.get(global_var) {
315 match const_key {
316 ConstKey::PrimOp(_) => {
320 unreachable!(
321 "Should have been handled by is_curried_primop_app"
322 )
323 }
324 ConstKey::Cast(sort1, sort2) => {
325 if fargs.len() != 1 {
326 Err(FixpointParseError::CastArityMismatch(fargs.len()))
327 } else {
328 Ok(rty::Expr::cast(
329 sort1.clone(),
330 sort2.clone(),
331 self.fixpoint_to_expr(&fargs[0])?,
332 ))
333 }
334 }
335 ConstKey::Alias(assoc_id, generic_args) => {
336 let lowered_args: flux_rustc_bridge::ty::GenericArgs =
337 generic_args.lower(self.genv.tcx()).unwrap();
338 let generic_args = rty::refining::Refiner::default_for_item(
339 self.genv,
340 assoc_id.parent(),
341 )
342 .unwrap()
343 .refine_generic_args(assoc_id.parent(), &lowered_args)
344 .unwrap();
345 let alias_reft =
346 rty::AliasReft { assoc_id: *assoc_id, args: generic_args };
347 let args = fargs
348 .iter()
349 .map(|farg| self.fixpoint_to_expr(farg))
350 .try_collect()?;
351 Ok(rty::Expr::alias(alias_reft, args))
352 }
353 ConstKey::WKVar(..) => {
354 unreachable!("WKVars should not appear in global vars");
355 }
356 ConstKey::RustConst(..) | ConstKey::Lambda(..) => {
357 self.fixpoint_app_to_expr(fhead, fargs)
359 }
360 }
361 } else {
362 Err(FixpointParseError::NoGlobalVar(*global_var))
363 }
364 }
365 fhead => self.fixpoint_app_to_expr(fhead, fargs),
366 }
367 }
368 fixpoint::Expr::Neg(fexpr) => {
369 let e = self.fixpoint_to_expr(fexpr)?;
370 Ok(rty::Expr::neg(&e))
371 }
372 fixpoint::Expr::BinaryOp(fbinop, boxed_args) => {
373 let binop = match fbinop {
374 fixpoint::BinOp::Add => rty::BinOp::Add(rty::Sort::Int),
378 fixpoint::BinOp::Sub => rty::BinOp::Sub(rty::Sort::Int),
379 fixpoint::BinOp::Mul => rty::BinOp::Mul(rty::Sort::Int),
380 fixpoint::BinOp::Div => rty::BinOp::Div(rty::Sort::Int),
381 fixpoint::BinOp::Mod => rty::BinOp::Mod(rty::Sort::Int),
382 };
383 let [fe1, fe2] = &**boxed_args;
384 let e1 = self.fixpoint_to_expr(fe1)?;
385 let e2 = self.fixpoint_to_expr(fe2)?;
386 Ok(rty::Expr::binary_op(binop, e1, e2))
387 }
388 fixpoint::Expr::IfThenElse(boxed_args) => {
389 let [fe1, fe2, fe3] = &**boxed_args;
390 let e1 = self.fixpoint_to_expr(fe1)?;
391 let e2 = self.fixpoint_to_expr(fe2)?;
392 let e3 = self.fixpoint_to_expr(fe3)?;
393 Ok(rty::Expr::ite(e1, e2, e3))
394 }
395 fixpoint::Expr::And(fexprs) => {
396 let exprs: Vec<rty::Expr> = fexprs
397 .iter()
398 .map(|fexpr| self.fixpoint_to_expr(fexpr))
399 .try_collect()?;
400 Ok(rty::Expr::and_from_iter(exprs))
401 }
402 fixpoint::Expr::Or(fexprs) => {
403 let exprs: Vec<rty::Expr> = fexprs
404 .iter()
405 .map(|fexpr| self.fixpoint_to_expr(fexpr))
406 .try_collect()?;
407 Ok(rty::Expr::or_from_iter(exprs))
408 }
409 fixpoint::Expr::Not(fexpr) => {
410 let e = self.fixpoint_to_expr(fexpr)?;
411 Ok(rty::Expr::not(&e))
412 }
413 fixpoint::Expr::Imp(boxed_args) => {
414 let [fe1, fe2] = &**boxed_args;
415 let e1 = self.fixpoint_to_expr(fe1)?;
416 let e2 = self.fixpoint_to_expr(fe2)?;
417 Ok(rty::Expr::binary_op(rty::BinOp::Imp, e1, e2))
418 }
419 fixpoint::Expr::Iff(boxed_args) => {
420 let [fe1, fe2] = &**boxed_args;
421 let e1 = self.fixpoint_to_expr(fe1)?;
422 let e2 = self.fixpoint_to_expr(fe2)?;
423 Ok(rty::Expr::binary_op(rty::BinOp::Iff, e1, e2))
424 }
425 fixpoint::Expr::Atom(fbinrel, boxed_args) => {
426 let binrel = match fbinrel {
427 fixpoint::BinRel::Eq => rty::BinOp::Eq,
428 fixpoint::BinRel::Ne => rty::BinOp::Ne,
429 fixpoint::BinRel::Gt => rty::BinOp::Gt(rty::Sort::Int),
437 fixpoint::BinRel::Ge => rty::BinOp::Ge(rty::Sort::Int),
438 fixpoint::BinRel::Lt => rty::BinOp::Lt(rty::Sort::Int),
439 fixpoint::BinRel::Le => rty::BinOp::Le(rty::Sort::Int),
440 };
441 let [fe1, fe2] = &**boxed_args;
442 let e1 = self.fixpoint_to_expr(fe1)?;
443 let e2 = self.fixpoint_to_expr(fe2)?;
444 Ok(rty::Expr::binary_op(binrel, e1, e2))
445 }
446 fixpoint::Expr::Let(_var, _boxed_args) => {
447 todo!("Convert `var` in e2 to locally nameless var, then fill in sort");
454 }
456 fixpoint::Expr::ThyFunc(itf) => Ok(rty::Expr::global_func(SpecFuncKind::Thy(*itf))),
457 fixpoint::Expr::IsCtor(var, fe) => {
458 let (def_id, variant_idx) = match var {
459 fixpoint::Var::DataCtor(adt_id, variant_idx) => {
460 let def_id = self.scx.adt_sorts[adt_id.as_usize()];
461 Ok((def_id, *variant_idx))
462 }
463 _ => Err(FixpointParseError::WrongVarInIsCtor(*var)),
464 }?;
465 let e = self.fixpoint_to_expr(fe)?;
466 Ok(rty::Expr::is_ctor(def_id, variant_idx, e))
467 }
468 fixpoint::Expr::Quantifier(q, binder, body) => {
469 let expr = self.fixpoint_to_bind_expr(binder, body)?;
470 match q {
471 fixpoint::Quantifier::Exists => Ok(rty::Expr::exists(expr)),
472 fixpoint::Quantifier::Forall => Ok(rty::Expr::forall(expr)),
473 }
474 }
475 fixpoint::Expr::WKVar(fixpoint::WKVar { wkvid, args }) => {
476 let e_args: Vec<rty::Expr> = args
477 .iter()
478 .map(|fexpr| self.fixpoint_to_expr(fexpr))
479 .try_collect()?;
480 if let Some(const_key) = self.ecx.const_env.wkvar_map_rev.get(wkvid) {
481 match const_key {
482 ConstKey::WKVar(wkvid, self_args) => {
483 Ok(rty::Expr::wkvar(rty::WKVar {
484 wkvid: wkvid.clone(),
485 self_args: *self_args,
486 args: List::from_vec(e_args),
487 }))
488 }
489 _ => {
490 unreachable!("Weak KVar has a const_key that is not a wkvid");
491 }
492 }
493 } else {
494 unreachable!("missing weak kvar {:?} in const_env", wkvid);
495 }
496 }
497 }
498 }
499
500 fn fixpoint_to_bind_expr(
501 &mut self,
502 binder: &[(fixpoint::Var, fixpoint::Sort)],
503 body: &fixpoint::Expr,
504 ) -> Result<rty::Binder<rty::Expr>, FixpointParseError> {
505 let mut vars = vec![];
506 let mut sorts = vec![];
507 for (var, sort) in binder {
508 let fixpoint::Var::Local(local_var) = var else {
509 return Err(FixpointParseError::WrongVarInBinder(*var));
510 };
511 vars.push(*local_var);
512 sorts.push(self.fixpoint_to_sort(sort)?);
513 }
514 self.ecx.local_var_env.push_layer(vars);
515 let body = self.fixpoint_to_expr(body)?;
516 self.ecx.local_var_env.pop_layer();
517 Ok(Binder::bind_with_sorts(body, &sorts))
518 }
519
520 fn fixpoint_app_to_expr(
521 &mut self,
522 fhead: &fixpoint::Expr,
523 fargs: &[fixpoint::Expr],
524 ) -> Result<rty::Expr, FixpointParseError> {
525 let head = self.fixpoint_to_expr(fhead)?;
526 let args = fargs
527 .iter()
528 .map(|farg| self.fixpoint_to_expr(farg))
529 .try_collect()?;
530 Ok(rty::Expr::app(head, List::empty(), args))
531 }
532}
533
534#[derive(Debug)]
535pub enum FixpointParseError {
536 UIFRelArityMismatch(usize),
539 TupleCtorArityMismatch(usize, usize),
541 ProjArityMismatch(usize),
543 NoGlobalVar(fixpoint::GlobalVar),
544 CastArityMismatch(usize),
546 PrimOpArityMismatch(usize),
547 WrongVarInIsCtor(fixpoint::Var),
550 WrongVarInBinder(fixpoint::Var),
552 UnknownAdt(DefId),
553 InvalidDefKindForCtor,
554}