use super::{Rewrite, Rule};
use crate::ast::{Action, Actions, Expr, Fact};
use crate::*;
use egglog_ast::span::Span;
pub(crate) fn desugar_command(
command: Command,
parser: &mut Parser,
proof_testing: bool,
) -> Result<Vec<NCommand>, Error> {
let res = match command {
Command::Function {
span,
name,
schema,
merge,
hidden,
let_binding,
term_constructor,
unextractable,
} => {
let mut fdecl = FunctionDecl::function(span, name, schema, merge);
fdecl.internal_hidden = hidden;
fdecl.internal_let = let_binding;
fdecl.term_constructor = term_constructor;
if fdecl.term_constructor.is_some() {
fdecl.unextractable = unextractable;
} else if unextractable {
fdecl.unextractable = true;
}
vec![NCommand::Function(fdecl)]
}
Command::Constructor {
span,
name,
schema,
cost,
unextractable,
hidden,
let_binding,
term_constructor,
} => {
let mut fdecl =
FunctionDecl::constructor(span, name, schema, cost, unextractable, hidden);
fdecl.internal_let = let_binding;
fdecl.term_constructor = term_constructor;
std::iter::once(NCommand::Function(fdecl)).collect()
}
Command::Relation { span, name, inputs } => desugar_relation(parser, span, name, inputs),
Command::Datatype {
span,
name,
variants,
} => desugar_datatype(span, name, variants),
Command::Datatypes { span: _, datatypes } => {
let mut res = vec![];
for datatype in datatypes.iter() {
let span = datatype.0.clone();
let name = datatype.1.clone();
if let Subdatatypes::Variants(..) = datatype.2 {
res.push(NCommand::Sort {
span,
name,
presort_and_args: None,
uf: None,
proof_func: None,
container_rebuild: None,
proof_constructors: None,
unionable: true,
});
}
}
let (variants_vec, sorts): (Vec<_>, Vec<_>) = datatypes
.into_iter()
.partition(|datatype| matches!(datatype.2, Subdatatypes::Variants(..)));
for sort in sorts {
let span = sort.0.clone();
let name = sort.1;
let Subdatatypes::NewSort(sort, args) = sort.2 else {
unreachable!()
};
res.push(NCommand::Sort {
span,
name,
presort_and_args: Some((sort, args)),
uf: None,
proof_func: None,
container_rebuild: None,
proof_constructors: None,
unionable: true,
});
}
for variants in variants_vec {
let datatype = variants.1;
let Subdatatypes::Variants(variants) = variants.2 else {
unreachable!();
};
for variant in variants {
res.push(NCommand::Function(FunctionDecl::constructor(
variant.span,
variant.name,
Schema {
input: variant.types,
output: datatype.clone(),
},
variant.cost,
false,
false,
)));
}
}
res
}
Command::Rewrite(ruleset, mut rewrite, subsume) => {
let name = RuleLikeCommand::Rewrite(&ruleset, &mut rewrite, subsume).resolve_name();
vec![desugar_rewrite(ruleset, name, rewrite, subsume, parser)?]
}
Command::BiRewrite(ruleset, mut rewrite) => {
let name = RuleLikeCommand::BiRewrite(&ruleset, &mut rewrite).resolve_name();
desugar_birewrite(ruleset, name, rewrite, parser)?
}
Command::Include(_span, _file) => {
unreachable!("Include commands should be expanded before desugaring")
}
Command::Rule { mut rule } => {
rule.name = RuleLikeCommand::Rule(&mut rule).resolve_name();
vec![NCommand::NormRule { rule }]
}
Command::Sort {
span,
name,
presort_and_args,
uf,
proof_func,
container_rebuild,
proof_constructors,
unionable,
} => vec![NCommand::Sort {
span,
name,
presort_and_args,
uf,
proof_func,
container_rebuild,
proof_constructors,
unionable,
}],
Command::AddRuleset(span, name) => vec![NCommand::AddRuleset(span, name)],
Command::UnstableCombinedRuleset(span, name, subrulesets) => {
vec![NCommand::UnstableCombinedRuleset(span, name, subrulesets)]
}
Command::Action(action) => vec![NCommand::CoreAction(action)],
Command::RunSchedule(sched) => {
vec![NCommand::RunSchedule(sched.clone())]
}
Command::PrintOverallStatistics(span, file) => {
vec![NCommand::PrintOverallStatistics(span, file.clone())]
}
Command::Extract(span, expr, variants) => vec![NCommand::Extract(span, expr, variants)],
Command::Check(span, facts) => {
if proof_testing {
desugar_prove(parser, span.clone(), facts.clone())
} else {
vec![NCommand::Check(span, facts)]
}
}
Command::PrintFunction(span, symbol, size, file, mode) => {
vec![NCommand::PrintFunction(span, symbol, size, file, mode)]
}
Command::PrintSize(span, symbol) => vec![NCommand::PrintSize(span, symbol)],
Command::Output { span, file, exprs } => {
vec![NCommand::Output { span, file, exprs }]
}
Command::Push(num) => {
vec![NCommand::Push(num)]
}
Command::Pop(span, num) => {
vec![NCommand::Pop(span, num)]
}
Command::Fail(span, cmd) => {
if let Command::Include(..) = *cmd {
return Err(Error::DesugarError(
span.clone(),
"include is not allowed inside (fail ...)".to_string(),
));
}
let mut desugared = desugar_command(*cmd, parser, proof_testing)?;
let Some(last) = desugared.pop() else {
return Err(Error::DesugarError(
span.clone(),
"the command inside (fail ...) expands to no commands".to_string(),
));
};
desugared.push(NCommand::Fail(span, Box::new(last)));
return Ok(desugared);
}
Command::Input { span, name, file } => {
vec![NCommand::Input { span, name, file }]
}
Command::UserDefined(span, name, args) => {
vec![NCommand::UserDefined(span, name, args)]
}
Command::Prove(span, query) => desugar_prove(parser, span, query),
Command::ProveExists(span, constructor) => {
vec![NCommand::ProveExists(span, constructor)]
}
};
Ok(res)
}
enum RuleLikeCommand<'a> {
Rule(&'a mut Rule),
Rewrite(&'a str, &'a mut Rewrite, bool),
BiRewrite(&'a str, &'a mut Rewrite),
}
impl RuleLikeCommand<'_> {
fn resolve_name(mut self) -> String {
let supplied_name = match &mut self {
Self::Rule(rule) => std::mem::take(&mut rule.name),
Self::Rewrite(_, rw, _) | Self::BiRewrite(_, rw) => std::mem::take(&mut rw.name),
};
if !supplied_name.is_empty() {
return supplied_name;
}
self.to_string().replace('\"', "'")
}
}
impl std::fmt::Display for RuleLikeCommand<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Rule(rule) => std::fmt::Display::fmt(&**rule, f),
Self::Rewrite(ruleset, rw, subsume) => rw.fmt_with_ruleset(f, ruleset, false, *subsume),
Self::BiRewrite(ruleset, rw) => rw.fmt_with_ruleset(f, ruleset, true, false),
}
}
}
fn desugar_prove(parser: &mut Parser, span: Span, query: Vec<Fact>) -> Vec<NCommand> {
let fresh_sort = parser.symbol_gen.fresh("ExistsSort");
let constructor_name = parser.symbol_gen.fresh("ExistsConstructor");
let ruleset = parser.symbol_gen.fresh("exists");
let name = parser.symbol_gen.fresh("prove_exists_rule");
vec![
NCommand::Sort {
span: span.clone(),
name: fresh_sort.clone(),
presort_and_args: None,
uf: None,
proof_func: None,
container_rebuild: None,
proof_constructors: None,
unionable: false,
},
NCommand::Function(FunctionDecl::constructor(
span.clone(),
constructor_name.clone(),
Schema {
input: vec![],
output: fresh_sort.clone(),
},
None,
false,
true, )),
NCommand::AddRuleset(span.clone(), ruleset.clone()),
NCommand::NormRule {
rule: Rule {
span: span.clone(),
body: query,
head: Actions::singleton(Action::Expr(
span.clone(),
Expr::Call(span.clone(), constructor_name.clone(), vec![]),
)),
ruleset: ruleset.clone(),
name,
eval_mode: RuleEvalMode::Seminaive,
no_decomp: false,
include_subsumed: false,
},
},
NCommand::RunSchedule(GenericSchedule::Run(
span.clone(),
GenericRunConfig {
ruleset,
until: None,
},
)),
NCommand::ProveExists(span, constructor_name),
]
}
fn desugar_datatype(span: Span, name: String, variants: Vec<Variant>) -> Vec<NCommand> {
vec![NCommand::Sort {
span: span.clone(),
name: name.clone(),
presort_and_args: None,
uf: None,
proof_func: None,
container_rebuild: None,
proof_constructors: None,
unionable: true,
}]
.into_iter()
.chain(variants.into_iter().map(|variant| {
NCommand::Function(FunctionDecl::constructor(
variant.span,
variant.name,
Schema {
input: variant.types,
output: name.clone(),
},
variant.cost,
variant.unextractable,
false,
))
}))
.collect()
}
fn desugar_rewrite(
ruleset: String,
name: String,
rewrite: Rewrite,
subsume: bool,
parser: &mut Parser,
) -> Result<NCommand, Error> {
let span = rewrite.span;
let lhs = rewrite.lhs;
let rhs = rewrite.rhs;
let conditions = rewrite.conditions;
let var = parser.symbol_gen.fresh("rewrite_var__");
let mut head = Actions::singleton(Action::Union(
span.clone(),
Expr::Var(span.clone(), var.clone()),
rhs,
));
if subsume {
match &lhs {
Expr::Call(_, f, args) => {
head.0.push(Action::Change(
span.clone(),
Change::Subsume,
f.clone(),
args.to_vec(),
));
}
_ => {
return Err(Error::DesugarError(
lhs.span(),
"subsumed rewrite must have a function call on the lhs".to_string(),
));
}
}
}
Ok(NCommand::NormRule {
rule: Rule {
span: span.clone(),
body: [Fact::Eq(span.clone(), Expr::Var(span, var), lhs)]
.into_iter()
.chain(conditions)
.collect(),
head,
ruleset,
name,
eval_mode: RuleEvalMode::Seminaive,
no_decomp: false,
include_subsumed: false,
},
})
}
fn desugar_birewrite(
ruleset: String,
rewrite_name: String,
rewrite: Rewrite,
parser: &mut Parser,
) -> Result<Vec<NCommand>, Error> {
let mut reverse = rewrite.clone();
std::mem::swap(&mut reverse.lhs, &mut reverse.rhs);
let forward = desugar_rewrite(
ruleset.clone(),
format!("{rewrite_name}=>"),
rewrite,
false,
parser,
)?;
let backward = desugar_rewrite(ruleset, format!("{rewrite_name}<="), reverse, false, parser)?;
Ok(vec![forward, backward])
}
fn desugar_relation(
parser: &mut Parser,
span: Span,
name: String,
inputs: Vec<String>,
) -> Vec<NCommand> {
let dashes_removed = name.replace('-', "");
let fresh_sort = parser.symbol_gen.fresh(&format!("{dashes_removed}Sort"));
vec![
NCommand::Sort {
span: span.clone(),
name: fresh_sort.clone(),
presort_and_args: None,
uf: None,
proof_func: None,
container_rebuild: None,
proof_constructors: None,
unionable: false,
},
NCommand::Function(FunctionDecl::constructor(
span,
name,
Schema {
input: inputs,
output: fresh_sort,
},
None,
false,
false,
)),
]
}
pub fn rule_name<Head, Leaf>(command: &GenericCommand<Head, Leaf>) -> String
where
Head: Clone + Display,
Leaf: Clone + PartialEq + Eq + Hash + Display,
{
command.to_string().replace('\"', "'")
}