use crate::ast::ResolvedNCommand;
use crate::core::ResolvedCall;
use crate::*;
use egglog_ast::generic_ast::GenericExpr;
pub(crate) fn proof_form(
prog: Vec<ResolvedNCommand>,
fresh: &mut SymbolGen,
) -> Vec<ResolvedNCommand> {
prog.into_iter()
.map(|cmd| proof_form_cmd(cmd, fresh))
.collect()
}
fn proof_form_cmd(cmd: ResolvedNCommand, fresh: &mut SymbolGen) -> ResolvedNCommand {
cmd.visit_queries(&mut |query| {
let mut new_query = vec![];
for fact in query {
let rewritten = proof_form_fact(fact, &mut new_query, fresh);
new_query.push(rewritten);
}
new_query
})
}
fn proof_form_fact(
fact: ResolvedFact,
res: &mut Vec<ResolvedFact>,
fresh: &mut SymbolGen,
) -> ResolvedFact {
match fact {
ResolvedFact::Eq(
span,
ResolvedExpr::Call(span2, head @ ResolvedCall::Func(_), args),
ResolvedExpr::Var(span3, v),
) if head.is_custom_func() => {
let mut new_args = vec![];
for arg in args {
new_args.push(proof_form_expr(arg, res, fresh));
}
ResolvedFact::Eq(
span,
ResolvedExpr::Call(span2, head, new_args),
ResolvedExpr::Var(span3, v),
)
}
GenericFact::Eq(span, generic_expr, generic_expr2) => GenericFact::Eq(
span,
proof_form_expr(generic_expr, res, fresh),
proof_form_expr(generic_expr2, res, fresh),
),
GenericFact::Fact(generic_expr) => {
GenericFact::Fact(proof_form_expr(generic_expr, res, fresh))
}
}
}
fn proof_form_expr(
fact: ResolvedExpr,
res: &mut Vec<ResolvedFact>,
fresh: &mut SymbolGen,
) -> ResolvedExpr {
match fact {
ref fact @ ResolvedExpr::Call(
ref span,
ref head @ ResolvedCall::Func(ref func_type),
ref args,
) if head.is_custom_func() => {
let new_args = args
.iter()
.map(|expr| proof_form_expr(expr.clone(), res, fresh))
.collect();
let resolved = GenericExpr::Var(
span.clone(),
ResolvedVar {
name: fresh.fresh("n"),
sort: func_type.output.clone(),
is_global_ref: false,
},
);
res.push(ResolvedFact::Eq(
span.clone(),
ResolvedExpr::Call(span.clone(), head.clone(), new_args),
resolved.clone(),
));
log::warn!(
"Input program not in proof normal form! All function calls must be top-level in query.
Original fact: {fact}
New top level fact: {}
Replace with new variable {}
",
res.last().unwrap(),
resolved
);
resolved
}
ResolvedExpr::Call(span, head @ ResolvedCall::Primitive(_), args) => {
let mut new_args = vec![];
for arg in args {
match arg {
ref arg_expr @ ResolvedExpr::Call(
ref arg_span,
ResolvedCall::Func(ref func_type),
ref inner_args,
) => {
let normalized_inner_args: Vec<_> = inner_args
.iter()
.map(|e| proof_form_expr(e.clone(), res, fresh))
.collect();
let fresh_var = GenericExpr::Var(
arg_span.clone(),
ResolvedVar {
name: fresh.fresh("v"),
sort: func_type.output.clone(),
is_global_ref: false,
},
);
res.push(ResolvedFact::Eq(
arg_span.clone(),
ResolvedExpr::Call(
arg_span.clone(),
match arg_expr {
ResolvedExpr::Call(_, call, _) => call.clone(),
_ => unreachable!(),
},
normalized_inner_args,
),
fresh_var.clone(),
));
new_args.push(fresh_var);
}
other => {
new_args.push(proof_form_expr(other, res, fresh));
}
}
}
ResolvedExpr::Call(span, head, new_args)
}
ResolvedExpr::Call(span, head, args) => {
let mut new_args = vec![];
for arg in args {
let normalized = proof_form_expr(arg, res, fresh);
let lift = matches!(
&normalized,
ResolvedExpr::Call(_, ResolvedCall::Primitive(p), _)
if p.output().is_eq_container_sort()
);
if lift {
let (arg_span, sort) = match &normalized {
ResolvedExpr::Call(s, ResolvedCall::Primitive(p), _) => {
(s.clone(), p.output().clone())
}
_ => unreachable!(),
};
let fresh_var = GenericExpr::Var(
arg_span.clone(),
ResolvedVar {
name: fresh.fresh("v"),
sort,
is_global_ref: false,
},
);
res.push(ResolvedFact::Eq(arg_span, normalized, fresh_var.clone()));
new_args.push(fresh_var);
} else {
new_args.push(normalized);
}
}
ResolvedExpr::Call(span, head, new_args)
}
ResolvedExpr::Lit(..) | ResolvedExpr::Var(..) => fact,
}
}