use std::fmt::Debug;
use crate::ast::identifiers::GlobalId;
use crate::ast::*;
use crate::ast::{diagnostics::*, visitors::*};
use crate::phase::Phase;
use crate::names::rust_primitives::hax::explicit_monadic::*;
#[derive(Default, Debug)]
pub struct ExplicitMonadic;
#[setup_error_handling_struct]
#[derive(Default)]
struct ExplicitMonadicVisitor;
#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash, Ord, PartialOrd)]
enum MonadicStatus {
Computation,
Value,
}
impl Phase for ExplicitMonadic {
fn apply(&self, items: &mut Vec<Item>) {
ExplicitMonadicVisitor::default().visit(items)
}
}
impl ExplicitMonadicVisitor {
fn wrap_app(expr: &Expr, head_id: GlobalId) -> Box<ExprKind> {
let expr = expr.clone();
Box::new(ExprKind::App {
head: Expr {
kind: Box::new(ExprKind::GlobalId(head_id)),
ty: Ty(Box::new(TyKind::Arrow {
inputs: vec![expr.ty.clone()],
output: expr.ty.clone(),
})),
meta: Metadata {
span: expr.meta.span,
attributes: vec![],
},
},
args: vec![expr],
generic_args: vec![],
bounds_impls: vec![],
trait_: None,
})
}
fn coerce(&mut self, expr: &mut Expr, from: MonadicStatus, to: MonadicStatus) {
if from == to {
return;
}
expr.kind = ExplicitMonadicVisitor::wrap_app(
expr,
match to {
MonadicStatus::Computation => pure,
MonadicStatus::Value => lift,
},
);
}
}
impl VisitorWithContext for ExplicitMonadicVisitor {
fn context(&self) -> Context {
Context::Phase(stringify!(ExplicitMonadic).into())
}
}
impl ExplicitMonadicVisitor {
fn visit_expr_coerce(&mut self, constraint: MonadicStatus, expr: &mut Expr) {
let opt_status = match &mut *expr.kind {
ExprKind::If {
condition,
then,
else_,
} => {
self.visit_expr_coerce(MonadicStatus::Value, condition);
[Some(then), else_.as_mut()]
.into_iter()
.flatten()
.for_each(|branch| self.visit_expr_coerce(MonadicStatus::Computation, branch));
Some(MonadicStatus::Computation)
}
ExprKind::Match { scrutinee, arms } => {
self.visit_expr_coerce(MonadicStatus::Value, scrutinee);
arms.iter_mut().for_each(|arm| {
if let Some(Guard {
kind: GuardKind::IfLet { rhs, .. },
..
}) = &mut arm.guard
{
self.visit_expr_coerce(MonadicStatus::Value, rhs);
};
self.visit_expr_coerce(MonadicStatus::Computation, &mut arm.body)
});
Some(MonadicStatus::Computation)
}
ExprKind::Block { body, .. } => {
self.visit_expr_coerce(constraint, body);
None
}
ExprKind::Break { .. }
| ExprKind::Return { .. }
| ExprKind::Continue { .. }
| ExprKind::Loop { .. } => {
unreachable_by_invariant!(Functionalize_loops)
}
ExprKind::Let { lhs: _, rhs, body } => {
self.visit_expr_coerce(MonadicStatus::Computation, rhs);
self.visit_expr_coerce(MonadicStatus::Computation, body);
Some(MonadicStatus::Computation)
}
ExprKind::App { head, args, .. } => {
self.visit_expr_coerce(MonadicStatus::Value, head);
args.iter_mut()
.for_each(|arg| self.visit_expr_coerce(MonadicStatus::Value, arg));
if let ExprKind::GlobalId(head) = &*head.kind
&& head.is_projector()
{
Some(MonadicStatus::Value)
} else if args.is_empty() {
Some(MonadicStatus::Value)
} else {
Some(MonadicStatus::Computation)
}
}
ExprKind::Array(exprs) => {
exprs
.iter_mut()
.for_each(|expr| self.visit_expr_coerce(MonadicStatus::Value, expr));
Some(MonadicStatus::Value)
}
ExprKind::Construct { fields, base, .. } => {
fields
.iter_mut()
.map(|(_, e)| e)
.chain(base.iter_mut())
.for_each(|expr| self.visit_expr_coerce(MonadicStatus::Value, expr));
Some(MonadicStatus::Value)
}
ExprKind::Assign { value: inner, .. }
| ExprKind::Borrow { inner, .. }
| ExprKind::AddressOf { inner, .. } => {
self.visit_expr_coerce(MonadicStatus::Value, inner);
Some(MonadicStatus::Value)
}
ExprKind::Ascription { e, ty } => {
self.visit_expr_coerce(MonadicStatus::Value, e);
self.visit(ty);
Some(MonadicStatus::Value)
}
ExprKind::Closure {
params: _,
body,
captures,
} => {
captures
.iter_mut()
.for_each(|capture| self.visit_expr_coerce(MonadicStatus::Value, capture));
self.visit_expr_coerce(MonadicStatus::Computation, body);
Some(MonadicStatus::Value)
}
ExprKind::Literal(_)
| ExprKind::GlobalId(_)
| ExprKind::LocalId(_)
| ExprKind::Quote { .. }
| ExprKind::Error(_) => Some(MonadicStatus::Value),
ExprKind::Resugared(_) => {
unreachable!("Resugarings should happen after phases")
}
};
if let Some(status) = opt_status {
self.coerce(expr, status, constraint)
}
}
}
impl AstVisitorMut for ExplicitMonadicVisitor {
setup_error_handling_impl!();
fn visit_expr(&mut self, x: &mut Expr) {
self.visit_expr_coerce(MonadicStatus::Computation, x)
}
fn visit_ty(&mut self, x: &mut Ty) {
if let TyKind::Array { length, .. } = x.kind_mut() {
self.visit_expr_coerce(MonadicStatus::Value, length);
};
}
fn visit_generic_value(&mut self, x: &mut GenericValue) {
if let GenericValue::Expr(expr) = x {
self.visit_expr_coerce(MonadicStatus::Value, expr);
};
}
}