hax_rust_engine/phase/
explicit_monadic.rs1use std::fmt::Debug;
2
3use crate::ast::identifiers::GlobalId;
4use crate::ast::*;
5use crate::ast::{diagnostics::*, visitors::*};
6use crate::phase::Phase;
7
8use crate::names::rust_primitives::hax::explicit_monadic::*;
9
10#[derive(Default, Debug)]
35pub struct ExplicitMonadic;
36
37#[setup_error_handling_struct]
39#[derive(Default)]
40struct ExplicitMonadicVisitor;
41
42#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash, Ord, PartialOrd)]
44enum MonadicStatus {
45 Computation,
46 Value,
47}
48
49impl Phase for ExplicitMonadic {
50 fn apply(&self, items: &mut Vec<Item>) {
51 ExplicitMonadicVisitor::default().visit(items)
52 }
53}
54
55impl ExplicitMonadicVisitor {
56 fn wrap_app(expr: &Expr, head_id: GlobalId) -> Box<ExprKind> {
59 let expr = expr.clone();
60 Box::new(ExprKind::App {
61 head: Expr {
62 kind: Box::new(ExprKind::GlobalId(head_id)),
63 ty: Ty(Box::new(TyKind::Arrow {
64 inputs: vec![expr.ty.clone()],
65 output: expr.ty.clone(),
66 })),
67 meta: Metadata {
68 span: expr.meta.span.clone(),
69 attributes: vec![],
70 },
71 },
72 args: vec![expr],
73 generic_args: vec![],
74 bounds_impls: vec![],
75 trait_: None,
76 })
77 }
78
79 fn coerce(&mut self, expr: &mut Expr, from: MonadicStatus, to: MonadicStatus) {
81 if from == to {
83 return;
84 }
85 expr.kind = ExplicitMonadicVisitor::wrap_app(
86 expr,
87 match to {
88 MonadicStatus::Computation => pure,
90 MonadicStatus::Value => lift,
92 },
93 );
94 }
95}
96
97impl VisitorWithContext for ExplicitMonadicVisitor {
98 fn context(&self) -> Context {
99 Context::Phase(stringify!(ExplicitMonadic).into())
100 }
101}
102
103impl ExplicitMonadicVisitor {
104 fn visit_expr_coerce(&mut self, constraint: MonadicStatus, expr: &mut Expr) {
105 let opt_status = match &mut *expr.kind {
108 ExprKind::If {
110 condition,
111 then,
112 else_,
113 } => {
114 self.visit_expr_coerce(MonadicStatus::Value, condition);
115 [Some(then), else_.as_mut()]
116 .into_iter()
117 .flatten()
118 .for_each(|branch| self.visit_expr_coerce(MonadicStatus::Computation, branch));
119 Some(MonadicStatus::Computation)
120 }
121 ExprKind::Match { scrutinee, arms } => {
122 self.visit_expr_coerce(MonadicStatus::Value, scrutinee);
123 arms.iter_mut().for_each(|arm| {
124 if let Some(Guard {
125 kind: GuardKind::IfLet { rhs, .. },
126 ..
127 }) = &mut arm.guard
128 {
129 self.visit_expr_coerce(MonadicStatus::Value, rhs);
130 };
131 self.visit_expr_coerce(MonadicStatus::Computation, &mut arm.body)
132 });
133 Some(MonadicStatus::Computation)
134 }
135 ExprKind::Block { body, .. } => {
136 self.visit_expr_coerce(constraint, body);
137 None
138 }
139 ExprKind::Break { .. }
140 | ExprKind::Return { .. }
141 | ExprKind::Continue { .. }
142 | ExprKind::Loop { .. } => {
143 unreachable_by_invariant!(Functionalize_loops)
144 }
145 ExprKind::Let { lhs: _, rhs, body } => {
147 self.visit_expr_coerce(MonadicStatus::Computation, rhs);
148 self.visit_expr_coerce(MonadicStatus::Computation, body);
149 Some(MonadicStatus::Computation)
150 }
151 ExprKind::App { head, args, .. } => {
152 self.visit_expr_coerce(MonadicStatus::Value, head);
153 args.iter_mut()
154 .for_each(|arg| self.visit_expr_coerce(MonadicStatus::Value, arg));
155 if let ExprKind::GlobalId(head) = &*head.kind
156 && head.is_projector()
157 {
158 Some(MonadicStatus::Value)
160 } else {
161 Some(MonadicStatus::Computation)
163 }
164 }
165 ExprKind::Array(exprs) => {
166 exprs
167 .iter_mut()
168 .for_each(|expr| self.visit_expr_coerce(MonadicStatus::Value, expr));
169 Some(MonadicStatus::Value)
170 }
171 ExprKind::Construct { fields, base, .. } => {
172 fields
173 .iter_mut()
174 .map(|(_, e)| e)
175 .chain(base.iter_mut())
176 .for_each(|expr| self.visit_expr_coerce(MonadicStatus::Value, expr));
177 Some(MonadicStatus::Value)
178 }
179 ExprKind::Assign { value: inner, .. }
180 | ExprKind::Borrow { inner, .. }
181 | ExprKind::AddressOf { inner, .. }
182 | ExprKind::Deref(inner) => {
183 self.visit_expr_coerce(MonadicStatus::Value, inner);
184 Some(MonadicStatus::Value)
185 }
186 ExprKind::Ascription { e, ty } => {
187 self.visit_expr_coerce(MonadicStatus::Value, e);
188 self.visit(ty);
189 Some(MonadicStatus::Value)
190 }
191 ExprKind::Closure {
192 params: _,
193 body,
194 captures,
195 } => {
196 captures
197 .iter_mut()
198 .for_each(|capture| self.visit_expr_coerce(MonadicStatus::Value, capture));
199 self.visit_expr_coerce(MonadicStatus::Computation, body);
200 Some(MonadicStatus::Value)
201 }
202 ExprKind::Literal(_)
203 | ExprKind::GlobalId(_)
204 | ExprKind::LocalId(_)
205 | ExprKind::Quote { .. }
206 | ExprKind::Error(_) => Some(MonadicStatus::Value),
207 ExprKind::Resugared(_) => {
208 unreachable!("Resugarings should happen after phases")
209 }
210 };
211 if let Some(status) = opt_status {
212 self.coerce(expr, status, constraint)
213 }
214 }
215}
216
217impl AstVisitorMut for ExplicitMonadicVisitor {
218 setup_error_handling_impl!();
219
220 fn visit_expr(&mut self, x: &mut Expr) {
221 self.visit_expr_coerce(MonadicStatus::Computation, x)
224 }
225
226 fn visit_ty(&mut self, x: &mut Ty) {
227 if let TyKind::Array { length, .. } = x.kind_mut() {
228 self.visit_expr_coerce(MonadicStatus::Value, length);
229 };
230 }
231
232 fn visit_generic_value(&mut self, x: &mut GenericValue) {
233 if let GenericValue::Expr(expr) = x {
234 self.visit_expr_coerce(MonadicStatus::Value, expr);
235 };
236 }
237}