hax_rust_engine/phase/
explicit_monadic.rs

1use 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/// Monadic Phase
11///
12/// This module defines a phase that makes the monadic encoding explicit by introducing calls to hax
13/// primitives (`pure` and `lift`) when necessary.
14///
15/// # Details
16///
17/// In backends with a monadic encoding (Lean for instance), rust computations that can *crash* are
18/// wrapped in an error Monad (say `RustM`): a function `fn f(x:u32) -> u32` will be extracted to
19/// something like `def f (x:u32) : RustM u32`. There are two challenges in this encoding :
20///
21/// 1. Some expressions cannot panic (literals, consts, constructors for enums, etc) and should be
22///    wrapped in the monad[^coe]. This phase inserts explicit calls to `pure` to that aim.
23///
24/// 2. Language constructs (if-then-else, `match`, etc.) and rust functions still expect rust values
25///    as input, not monadic ones. This phase inserts explicit calls to `lift` to materialize the
26///    sub-expressions that return a monadic result where a value is expected. The Lean backend turns
27///    them into explicit lifts `(← ..)`, which implicitly introduces a monadic bind
28///
29/// This phase expects all function and closure bodies to be monadic computations by default.
30///
31/// [^coe]: While implicit coercions can sometime be enough, they can also badly interact with
32/// inference, typically when dealing with branches (like if-then-else) where some branches are
33/// pure and some are not.
34#[derive(Default, Debug)]
35pub struct ExplicitMonadic;
36
37/// Stateless visitor
38#[setup_error_handling_struct]
39#[derive(Default)]
40struct ExplicitMonadicVisitor;
41
42/// Status of a rust expression. Computations are possibly panicking, while values are pure
43#[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    /// Helper while waiting for a proper ast API. Wraps an expression in an application node, where
57    /// the head is a global id
58    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    /// Helper to coerce a expression into a given status. `from` should be the status of `expr`
80    fn coerce(&mut self, expr: &mut Expr, from: MonadicStatus, to: MonadicStatus) {
81        // If the status is already correct, nothing to do.
82        if from == to {
83            return;
84        }
85        expr.kind = ExplicitMonadicVisitor::wrap_app(
86            expr,
87            match to {
88                // from = Value, to = Computation : we insert `pure`
89                MonadicStatus::Computation => pure,
90                // from = Computation, to = Value : we insert `lift`
91                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        // Expression can force a status (returned as `Some(...)`), or be "transparent" (typically
106        // for control-flow) and just propagate the constraint.
107        let opt_status = match &mut *expr.kind {
108            // Control flow nodes
109            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            // Opaque nodes
146            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                    // Constructors for structures and enums are values
159                    Some(MonadicStatus::Value)
160                } else {
161                    // Other function calls are computations
162                    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        // Entry points are functions (items and impl items), which start with a `do` block,
222        // therefore a monadic computation
223        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}