#![doc = include_str!("lib.md")]
pub mod api;
pub mod ast;
#[cfg(feature = "bin")]
mod cli;
mod command_macro;
pub mod constraint;
mod core;
mod exec_state;
pub mod extract;
pub mod prelude;
mod proofs;
pub mod scheduler;
mod serialize;
pub mod sort;
mod termdag;
mod typechecking;
pub mod util;
pub use command_macro::{CommandMacro, CommandMacroRegistry};
extern crate self as egglog;
pub use ast::{ResolvedExpr, ResolvedFact, ResolvedVar};
#[cfg(feature = "bin")]
pub use cli::*;
use constraint::{Constraint, Problem, SimpleTypeConstraint, TypeConstraint};
use core::CoreActionContext;
use core::ResolvedAtomTerm;
pub use core::{Atom, AtomTerm};
pub use core::{ResolvedCall, SpecializedPrimitive};
pub use core_relations::{BaseValue, ContainerValue, Value};
use core_relations::{ExecutionState, ExternalFunctionId, make_external_func};
use csv::Writer;
pub use egglog_add_primitive::add_literal_prim;
pub use egglog_add_primitive::add_primitive;
pub use egglog_add_primitive::add_primitive_with_validator;
use egglog_ast::generic_ast::{Change, GenericExpr, Literal};
use egglog_ast::span::Span;
use egglog_ast::util::ListDisplay;
use egglog_bridge::{ColumnTy, QueryEntry};
use egglog_core_relations as core_relations;
use egglog_numeric_id as numeric_id;
use egglog_reports::{ReportLevel, RunReport};
pub use exec_state::{
Context, Core, Enode, FullState, FunctionEntry, PureState, Read, ReadState, Write, WriteState,
};
use extract::{DefaultCost, Extractor, TreeAdditiveCostModel};
use indexmap::map::Entry;
use log::{Level, log_enabled};
use numeric_id::DenseIdMap;
use prelude::*;
pub use proofs::proof_encoding_helpers::{file_supports_proofs, program_supports_proofs};
pub mod proof {
pub use crate::proofs::proof_format::{Justification, Proof, ProofId, ProofStore, Proposition};
}
use scheduler::{SchedulerId, SchedulerRecord};
pub use serialize::{SerializeConfig, SerializeOutput, SerializedNode};
use sort::*;
use std::any::{Any, TypeId};
use std::fmt::{Debug, Display, Formatter};
use std::fs::File;
use std::hash::Hash;
use std::io::{Read as _, Write as _};
use std::iter::once;
use std::ops::Deref;
use std::path::PathBuf;
use std::sync::Arc;
pub use termdag::{OrdTerm, Term, TermDag, TermId};
use thiserror::Error;
pub use typechecking::FuncType;
pub use typechecking::PrimitiveValidator;
pub use typechecking::TypeError;
pub use typechecking::TypeInfo;
use util::*;
use crate::ast::desugar::desugar_command;
use crate::ast::*;
use crate::core::{GenericActionsExt, ResolvedRuleExt};
use crate::proofs::proof_encoding::{EncodingState, ProofInstrumentor};
use crate::proofs::proof_encoding_helpers::{
ProofEncodingUnsupportedReason, command_supports_proof_encoding,
};
use crate::proofs::proof_extraction::ProveExistsError;
use crate::proofs::proof_format::{ProofId, ProofStore};
use crate::proofs::proof_normal_form::proof_form;
pub const GLOBAL_NAME_PREFIX: &str = "$";
pub type ArcSort = Arc<dyn Sort>;
pub trait Primitive: Send + Sync + 'static {
fn name(&self) -> &str;
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint>;
}
pub trait PurePrim: Primitive {
fn apply<'a, 'db>(&self, state: PureState<'a, 'db>, args: &[Value]) -> Option<Value>;
}
pub trait WritePrim: Primitive {
fn apply<'a, 'db>(&self, state: WriteState<'a, 'db>, args: &[Value]) -> Option<Value>;
}
pub trait ReadPrim: Primitive {
fn apply<'a, 'db>(&self, state: ReadState<'a, 'db>, args: &[Value]) -> Option<Value>;
}
pub trait FullPrim: Primitive {
fn apply<'a, 'db>(&self, state: FullState<'a, 'db>, args: &[Value]) -> Option<Value>;
}
pub trait UserDefinedCommandOutput: Debug + std::fmt::Display + Send + Sync {}
impl<T> UserDefinedCommandOutput for T where T: Debug + std::fmt::Display + Send + Sync {}
#[derive(Clone, Debug)]
#[allow(clippy::large_enum_variant)]
pub enum CommandOutput {
PrintFunctionSize(usize),
PrintAllFunctionsSize(Vec<(String, usize)>),
ExtractBest(TermDag, DefaultCost, TermId),
ExtractVariants(TermDag, Vec<TermId>),
ProveExists {
proof_store: ProofStore,
proof_id: ProofId,
},
OverallStatistics(RunReport),
PrintFunction(Function, TermDag, Vec<(TermId, TermId)>, PrintFunctionMode),
RunSchedule(RunReport),
UserDefined(Arc<dyn UserDefinedCommandOutput>),
}
impl CommandOutput {
pub fn snapshot_stable_under_proof_encoding(outputs: &[CommandOutput]) -> String {
outputs
.iter()
.filter_map(|output| match output {
CommandOutput::OverallStatistics(_) => None,
CommandOutput::PrintFunction(..) => None,
CommandOutput::ExtractBest(_, cost, _) => {
Some(format!("(extraction-costs {cost})\n"))
}
CommandOutput::ExtractVariants(..) => None,
other => Some(other.to_string()),
})
.collect::<Vec<_>>()
.join("")
}
}
impl std::fmt::Display for CommandOutput {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
CommandOutput::PrintFunctionSize(size) => writeln!(f, "{size}"),
CommandOutput::PrintAllFunctionsSize(names_and_sizes) => {
write!(f, "(")?;
for (i, (name, size)) in names_and_sizes.iter().enumerate() {
if i > 0 {
write!(f, " ")?;
}
write!(f, "({name} {size})")?;
if i < names_and_sizes.len() - 1 {
writeln!(f)?;
}
}
writeln!(f, ")")
}
CommandOutput::ExtractBest(termdag, _cost, term) => {
writeln!(f, "{}", termdag.to_string(*term))
}
CommandOutput::ExtractVariants(termdag, terms) => {
writeln!(f, "(")?;
for expr in terms {
writeln!(f, " {}", termdag.to_string(*expr))?;
}
writeln!(f, ")")
}
CommandOutput::ProveExists {
proof_store,
proof_id,
} => writeln!(f, "{}", proof_store.proof_to_string(*proof_id)),
CommandOutput::OverallStatistics(run_report) => {
write!(f, "Overall statistics:\n{run_report}")
}
CommandOutput::PrintFunction(function, termdag, terms_and_outputs, mode) => {
let out_is_unit = function.func_type.output.name() == UnitSort.name();
if *mode == PrintFunctionMode::CSV {
let mut wtr = Writer::from_writer(vec![]);
for (term_id, output) in terms_and_outputs {
let term = termdag.get(*term_id);
match term {
Term::App(name, children) => {
let mut values = vec![name.clone()];
for child_id in children {
values.push(termdag.to_string(*child_id));
}
if !out_is_unit {
values.push(termdag.to_string(*output));
}
wtr.write_record(&values).map_err(|_| std::fmt::Error)?;
}
_ => panic!("Expect function_to_dag to return a list of apps."),
}
}
let csv_bytes = wtr.into_inner().map_err(|_| std::fmt::Error)?;
f.write_str(&String::from_utf8(csv_bytes).map_err(|_| std::fmt::Error)?)
} else {
writeln!(f, "(")?;
for (term, output) in terms_and_outputs.iter() {
write!(f, " {}", termdag.to_string(*term))?;
if !out_is_unit {
write!(f, " -> {}", termdag.to_string(*output))?;
}
writeln!(f)?;
}
writeln!(f, ")")
}
}
CommandOutput::RunSchedule(_report) => Ok(()),
CommandOutput::UserDefined(output) => {
write!(f, "{}", *output)
}
}
}
}
trait ExtensionStateValue: Any + dyn_clone::DynClone + Send + Sync {}
impl<T> ExtensionStateValue for T where T: Any + Clone + Send + Sync {}
dyn_clone::clone_trait_object!(ExtensionStateValue);
#[derive(Clone)]
pub struct EGraph {
backend: egglog_bridge::EGraph,
pub parser: Parser,
names: check_shadowing::Names,
pushed_egraph: Option<Box<Self>>,
functions: IndexMap<String, Function>,
rulesets: IndexMap<String, Ruleset>,
pub fact_directory: Option<PathBuf>,
pub seminaive: bool,
pub no_decomp: bool,
type_info: TypeInfo,
overall_run_report: RunReport,
schedulers: DenseIdMap<SchedulerId, SchedulerRecord>,
commands: IndexMap<String, Arc<dyn UserDefinedCommand>>,
extension_state: HashMap<TypeId, Box<dyn ExtensionStateValue>>,
strict_mode: bool,
warned_about_global_prefix: bool,
command_macros: CommandMacroRegistry,
proof_state: EncodingState,
proof_check_program: Vec<ResolvedNCommand>,
}
pub trait UserDefinedCommand: Send + Sync {
fn update(&self, egraph: &mut EGraph, args: &[Expr]) -> Result<Vec<CommandOutput>, Error>;
}
#[derive(Clone)]
pub struct Function {
decl: ResolvedFunctionDecl,
func_type: Arc<FuncType>,
can_subsume: bool,
backend_id: egglog_bridge::FunctionId,
}
impl Function {
pub fn name(&self) -> &str {
&self.decl.name
}
pub fn func_type(&self) -> &FuncType {
&self.func_type
}
pub fn can_subsume(&self) -> bool {
self.can_subsume
}
pub fn is_let_binding(&self) -> bool {
self.decl.internal_let
}
pub fn is_hidden(&self) -> bool {
self.decl.internal_hidden
}
pub fn term_constructor(&self) -> Option<&str> {
self.decl.term_constructor.as_deref()
}
}
impl Debug for Function {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Function")
.field("decl", &self.decl)
.field("func_type", &self.func_type)
.finish()
}
}
impl Default for EGraph {
fn default() -> Self {
let mut parser = Parser::default();
let proof_state = EncodingState::new(&mut parser.symbol_gen);
let mut eg = Self {
backend: Default::default(),
parser,
names: Default::default(),
pushed_egraph: Default::default(),
functions: Default::default(),
rulesets: Default::default(),
fact_directory: None,
seminaive: true,
no_decomp: false,
overall_run_report: Default::default(),
type_info: Default::default(),
schedulers: Default::default(),
commands: Default::default(),
extension_state: Default::default(),
strict_mode: false,
warned_about_global_prefix: false,
command_macros: Default::default(),
proof_state,
proof_check_program: vec![],
};
add_base_sort(&mut eg, UnitSort, span!()).unwrap();
add_base_sort(&mut eg, StringSort, span!()).unwrap();
add_base_sort(&mut eg, BoolSort, span!()).unwrap();
add_base_sort(&mut eg, I64Sort, span!()).unwrap();
add_base_sort(&mut eg, F64Sort, span!()).unwrap();
add_base_sort(&mut eg, BigIntSort, span!()).unwrap();
add_base_sort(&mut eg, BigRatSort, span!()).unwrap();
eg.type_info.add_presort::<MapSort>(span!()).unwrap();
eg.type_info.add_presort::<SetSort>(span!()).unwrap();
eg.type_info.add_presort::<VecSort>(span!()).unwrap();
eg.type_info.add_presort::<FunctionSort>(span!()).unwrap();
eg.type_info.add_presort::<MultiSetSort>(span!()).unwrap();
eg.type_info.add_presort::<PairSort>(span!()).unwrap();
let neq_validator = |termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
if args.len() == 2 && args[0] != args[1] {
Some(termdag.lit(Literal::Unit))
} else {
None
}
};
add_primitive_with_validator!(
&mut eg,
"!=" = |a: #, b: #| -?> () {
(a != b).then_some(())
},
neq_validator
);
add_primitive_with_validator!(
&mut eg,
"bool-!=" = |a: #, b: #| -> bool {
(a != b)
},
|termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
if args.len() == 2 {
Some(termdag.lit(Literal::Bool(args[0] != args[1])))
} else {
None
}
}
);
add_primitive!(&mut eg, "value-eq" = |a: #, b: #| -?> () {
(a == b).then_some(())
});
add_primitive!(&mut eg, "ordering-min" = |a: #, b: #| -> # {
if a < b { a } else { b }
});
add_primitive!(&mut eg, "ordering-max" = |a: #, b: #| -> # {
if a > b { a } else { b }
});
eg.rulesets
.insert("".into(), Ruleset::Rules(Default::default()));
eg
}
}
struct ResolvedNCommands {
desugared: Vec<ResolvedNCommand>,
desugared_before_proofs: Vec<ResolvedNCommand>,
}
struct ResolvedNCommandsWithOutput {
outputs: Vec<CommandOutput>,
resolved: Vec<ResolvedNCommand>,
resolved_before_proofs: Vec<ResolvedNCommand>,
}
#[derive(Debug, Error)]
#[error("Not found: {0}")]
pub struct NotFoundError(String);
impl EGraph {
pub fn new(num_threads: usize) -> Self {
EGraph::default().with_num_threads(num_threads)
}
pub fn new_with_term_encoding() -> Self {
let mut egraph = EGraph::default();
egraph.proof_state.original_typechecking = Some(Box::new(egraph.clone()));
egraph
}
pub fn new_with_proofs() -> Self {
let mut egraph = EGraph::new_with_term_encoding();
egraph.proof_state.proofs_enabled = true;
egraph
}
#[cfg(feature = "bin")]
pub(crate) fn with_term_encoding_enabled(mut self) -> Self {
self.proof_state.original_typechecking = Some(Box::new(self.clone()));
self
}
#[cfg(feature = "bin")]
pub(crate) fn with_proofs_enabled(mut self) -> Self {
self = self.with_term_encoding_enabled();
self.proof_state.proofs_enabled = true;
self
}
pub fn with_proof_testing(mut self) -> Self {
self.proof_state.proof_testing = true;
self
}
pub fn with_num_threads(mut self, num_threads: usize) -> Self {
self.set_num_threads(num_threads);
self
}
pub fn set_num_threads(&mut self, num_threads: usize) {
self.backend.set_num_threads(num_threads);
if let Some(original) = &mut self.proof_state.original_typechecking {
original.set_num_threads(num_threads);
}
}
pub fn num_threads(&self) -> usize {
self.backend.num_threads()
}
pub fn extension_state<T>(&self) -> Option<&T>
where
T: Send + Sync + 'static,
{
let value = self.extension_state.get(&TypeId::of::<T>())?;
(value.as_ref() as &dyn Any).downcast_ref()
}
pub fn extension_state_or_default<T>(&mut self) -> &mut T
where
T: Default + Clone + Send + Sync + 'static,
{
let value = self
.extension_state
.entry(TypeId::of::<T>())
.or_insert_with(|| Box::new(T::default()));
(value.as_mut() as &mut dyn Any)
.downcast_mut()
.expect("extension state entry must have the requested type")
}
pub fn type_info(&mut self) -> &mut TypeInfo {
&mut self.type_info
}
pub fn command_macros(&self) -> &CommandMacroRegistry {
&self.command_macros
}
pub fn command_macros_mut(&mut self) -> &mut CommandMacroRegistry {
&mut self.command_macros
}
pub fn add_command(
&mut self,
name: String,
command: Arc<dyn UserDefinedCommand>,
) -> Result<(), Error> {
if self.commands.contains_key(&name)
|| self.functions.contains_key(&name)
|| self.type_info.get_prims(&name).is_some()
{
return Err(Error::CommandAlreadyExists(name, span!()));
}
self.commands.insert(name.clone(), command);
self.parser.add_user_defined(name)?;
Ok(())
}
pub fn set_strict_mode(&mut self, strict_mode: bool) {
self.strict_mode = strict_mode;
}
pub fn strict_mode(&self) -> bool {
self.strict_mode
}
#[doc(hidden)]
pub fn ensure_no_reserved_symbols(&mut self, should_ensure: bool) {
self.parser.ensure_no_reserved_symbols = should_ensure;
}
fn ensure_global_name_prefix(&mut self, span: &Span, name: &str) -> Result<(), TypeError> {
if name.starts_with(GLOBAL_NAME_PREFIX) {
return Ok(());
}
if self.strict_mode {
Err(TypeError::GlobalMissingPrefix {
name: name.to_owned(),
span: span.clone(),
})
} else {
self.warn_missing_global_prefix(span, name)?;
Ok(())
}
}
fn warn_missing_global_prefix(
&mut self,
span: &Span,
canonical_name: &str,
) -> Result<(), TypeError> {
if self.strict_mode {
return Err(TypeError::GlobalMissingPrefix {
name: format!("{GLOBAL_NAME_PREFIX}{canonical_name}"),
span: span.clone(),
});
}
if self.warned_about_global_prefix {
return Ok(());
}
self.warned_about_global_prefix = true;
log::warn!(
"{span}\nGlobal `{canonical_name}` should start with `{GLOBAL_NAME_PREFIX}`. Enable `--strict-mode` to turn this warning into an error. Suppressing additional warnings of this type."
);
Ok(())
}
fn warn_prefixed_non_globals(
&mut self,
span: &Span,
canonical_name: &str,
) -> Result<(), TypeError> {
if self.strict_mode {
return Err(TypeError::NonGlobalPrefixed {
name: canonical_name.to_string(),
span: span.clone(),
});
}
if self.warned_about_global_prefix {
return Ok(());
}
self.warned_about_global_prefix = true;
log::warn!(
"{span}\nNon-global `{canonical_name}` should not start with `{GLOBAL_NAME_PREFIX}`. Enable `--strict-mode` to turn this warning into an error. Suppressing additional warnings of this type."
);
Ok(())
}
pub fn push(&mut self) {
let prev_prev: Option<Box<Self>> = self.pushed_egraph.take();
let mut prev = self.clone();
prev.pushed_egraph = prev_prev;
self.pushed_egraph = Some(Box::new(prev));
}
pub fn pop(&mut self) -> Result<(), Error> {
match self.pushed_egraph.take() {
Some(mut e) => {
std::mem::swap(&mut self.overall_run_report, &mut e.overall_run_report);
std::mem::swap(&mut self.parser.symbol_gen, &mut e.parser.symbol_gen);
*self = *e;
Ok(())
}
None => Err(Error::Pop(span!())),
}
}
fn translate_expr_to_mergefn(
&self,
expr: &ResolvedExpr,
) -> Result<egglog_bridge::MergeFn, Error> {
match expr {
GenericExpr::Lit(_, literal) => {
let val = literal_to_value(&self.backend, literal);
Ok(egglog_bridge::MergeFn::Const(val))
}
GenericExpr::Var(span, resolved_var) => match resolved_var.name.as_str() {
"old" => Ok(egglog_bridge::MergeFn::Old),
"new" => Ok(egglog_bridge::MergeFn::New),
_ => Err(TypeError::Unbound(resolved_var.name.clone(), span.clone()).into()),
},
GenericExpr::Call(_, ResolvedCall::Func(f), args) => {
let translated_args = args
.iter()
.map(|arg| self.translate_expr_to_mergefn(arg))
.collect::<Result<Vec<_>, _>>()?;
Ok(egglog_bridge::MergeFn::Function(
self.functions[&f.name].backend_id,
translated_args,
))
}
GenericExpr::Call(_, ResolvedCall::Primitive(p), args) => {
let mut translated_args = args
.iter()
.map(|arg| self.translate_expr_to_mergefn(arg))
.collect::<Result<Vec<_>, _>>()?;
if p.name() == "unstable-fn" {
let Some(GenericExpr::Lit(span, Literal::String(name))) = args.first() else {
return Err(Error::BackendError(
"expected string literal after `unstable-fn`".into(),
));
};
let resolved = resolve_function_container_target_with_context(
&self.backend,
&self.functions,
&self.type_info,
name,
p,
self.backend
.action_registry()
.read()
.unwrap()
.default_panic_id(),
crate::Context::Write,
span,
)?;
translated_args[0] =
egglog_bridge::MergeFn::Const(self.backend.base_values().get(resolved));
}
Ok(egglog_bridge::MergeFn::Primitive(
p.external_id(crate::Context::Write),
translated_args,
))
}
}
}
fn declare_function(&mut self, decl: &ResolvedFunctionDecl) -> Result<(), Error> {
let func_type = match self.type_info.get_func_type(&decl.name) {
Some(func_type) => {
debug_assert!(
func_type.subtype == decl.subtype
&& func_type.input.len() == decl.schema.input.len()
&& func_type
.input
.iter()
.zip(&decl.schema.input)
.all(|(sort, name)| sort.name() == name)
&& func_type.output.name() == decl.schema.output,
"recorded signature for {} disagrees with its declaration",
decl.name
);
func_type.clone()
}
None => {
let get_sort = |name: &String| match self.type_info.get_sort_by_name(name) {
Some(sort) => Ok(sort.clone()),
None => Err(Error::TypeError(TypeError::UndefinedSort(
name.to_owned(),
decl.span.clone(),
))),
};
let func_type = Arc::new(FuncType {
name: decl.name.clone(),
subtype: decl.subtype,
input: decl
.schema
.input
.iter()
.map(get_sort)
.collect::<Result<Vec<_>, _>>()?,
output: get_sort(&decl.schema.output)?,
});
self.type_info.declare_func_type(func_type.clone());
func_type
}
};
let can_subsume = match decl.subtype {
FunctionSubtype::Constructor => true,
FunctionSubtype::Custom => decl.term_constructor.is_some(),
};
use egglog_bridge::{DefaultVal, MergeFn};
let backend_id = self.backend.add_table(egglog_bridge::FunctionConfig {
schema: func_type
.input
.iter()
.chain([&func_type.output])
.map(|sort| sort.column_ty(&self.backend))
.collect(),
default: match decl.subtype {
FunctionSubtype::Constructor => DefaultVal::FreshId,
FunctionSubtype::Custom => DefaultVal::Fail,
},
merge: match decl.subtype {
FunctionSubtype::Constructor => MergeFn::UnionId,
FunctionSubtype::Custom => match &decl.merge {
None => MergeFn::AssertEq,
Some(expr) => self.translate_expr_to_mergefn(expr)?,
},
},
name: decl.name.to_string(),
can_subsume,
});
let function = Function {
decl: decl.clone(),
func_type,
can_subsume,
backend_id,
};
let old = self.functions.insert(decl.name.clone(), function);
if old.is_some() {
panic!(
"Typechecking should have caught function already bound: {}",
decl.name
);
}
Ok(())
}
pub fn print_function(
&mut self,
sym: &str,
n: Option<usize>,
file: Option<(File, PathBuf)>,
span: Span,
mode: PrintFunctionMode,
) -> Result<Option<CommandOutput>, Error> {
let n = match n {
Some(n) => {
log::info!("Printing up to {n} tuples of function {sym} as {mode}");
n
}
None => {
log::info!("Printing all tuples of function {sym} as {mode}");
usize::MAX
}
};
let (terms, outputs, termdag) = self.function_to_dag(sym, n, true)?;
let f = self
.functions
.get(sym)
.unwrap();
let terms_and_outputs: Vec<_> = terms.into_iter().zip(outputs.unwrap()).collect();
let output = CommandOutput::PrintFunction(f.clone(), termdag, terms_and_outputs, mode);
match file {
Some((mut file, path)) => {
log::info!("Writing output to file");
file.write_all(output.to_string().as_bytes())
.map_err(|e| Error::IoError(path, e, span.clone()))?;
Ok(None)
}
None => Ok(Some(output)),
}
}
#[doc(hidden)]
pub fn set_proof_checking_program(
&mut self,
prog: Vec<Command>,
proof_testing: bool,
) -> Result<(), Error> {
let mut proof_check_eg = EGraph::new_with_proofs();
if proof_testing {
proof_check_eg = proof_check_eg.with_proof_testing();
}
let resolved = proof_check_eg.process_program_internal(prog, false)?;
self.proof_check_program = resolved.resolved_before_proofs;
Ok(())
}
pub fn print_size(&self, sym: Option<&str>) -> Result<CommandOutput, Error> {
if let Some(sym) = sym {
let f = self
.functions
.values()
.find(|f| f.decl.term_constructor.as_deref() == Some(sym))
.or_else(|| self.functions.get(sym))
.ok_or(TypeError::UnboundFunction(sym.to_owned(), span!()))?;
if f.decl.internal_hidden || f.decl.internal_let {
return Err(TypeError::UnboundFunction(sym.to_owned(), span!()).into());
}
let size = self.backend.table_size(f.backend_id);
log::info!("Function {sym} has size {size}");
Ok(CommandOutput::PrintFunctionSize(size))
} else {
let mut lens = self
.functions
.iter()
.filter(|(_, f)| !f.decl.internal_hidden && !f.decl.internal_let)
.map(|(sym, f)| {
let name = f
.decl
.term_constructor
.clone()
.unwrap_or_else(|| sym.clone());
(name, self.backend.table_size(f.backend_id))
})
.collect::<Vec<_>>();
lens.sort_by_key(|(name, _)| name.clone());
if log_enabled!(Level::Info) {
for (sym, len) in &lens {
log::info!("Function {sym} has size {len}");
}
}
Ok(CommandOutput::PrintAllFunctionsSize(lens))
}
}
fn run_schedule(&mut self, sched: &ResolvedSchedule) -> Result<RunReport, Error> {
match sched {
ResolvedSchedule::Run(span, config) => self.run_rules(span, config),
ResolvedSchedule::Repeat(_span, limit, sched) => {
let mut report = RunReport::default();
for _i in 0..*limit {
let rec = self.run_schedule(sched)?;
let can_stop = rec.can_stop;
report.union(rec);
if can_stop {
break;
}
}
Ok(report)
}
ResolvedSchedule::Saturate(_span, sched) => {
let mut report = RunReport::default();
let mut i = 0usize;
loop {
i += 1;
log::debug!(
"Saturate iteration {i} start: {}",
Self::schedule_for_log(sched)
);
let rec = self.run_schedule(sched)?;
let updated = rec.updated;
log::debug!(
"Saturate iteration {i} end: {}",
Self::run_report_debug_summary(&rec)
);
report.union(rec);
if !updated {
log::debug!("Saturate reached fixpoint after {i} iteration(s)");
break;
}
}
Ok(report)
}
ResolvedSchedule::Sequence(_span, scheds) => {
let mut report = RunReport::default();
for sched in scheds {
report.union(self.run_schedule(sched)?);
}
Ok(report)
}
}
}
fn run_rules(&mut self, span: &Span, config: &ResolvedRunConfig) -> Result<RunReport, Error> {
log::debug!("Running ruleset: {}", config.ruleset);
let mut report: RunReport = Default::default();
let GenericRunConfig { ruleset, until } = config;
if !self.rulesets.contains_key(ruleset) {
return Err(Error::NoSuchRuleset(ruleset.clone(), span.clone()));
}
if let Some(facts) = until
&& self.check_facts(span, facts).is_ok()
{
log::info!(
"Breaking early because of facts:\n {}!",
ListDisplay(facts, "\n")
);
return Ok(report);
}
let subreport = self.step_rules(ruleset)?;
report.union(subreport);
if log_enabled!(Level::Debug) {
log::debug!(
"Finished ruleset {ruleset}: database size {}, {}",
self.num_tuples(),
Self::run_report_debug_summary(&report)
);
}
Ok(report)
}
fn run_report_debug_summary(report: &RunReport) -> String {
let mut rules = report
.num_matches_per_rule
.iter()
.filter(|(_, matches)| **matches > 0)
.collect::<Vec<_>>();
rules.sort_by(|(_, left), (_, right)| right.cmp(left));
let top_rules = rules
.into_iter()
.take(5)
.map(|(rule, matches)| {
format!("{}={matches}", Self::truncate_for_log(rule.as_ref(), 80))
})
.collect::<Vec<_>>()
.join(", ");
format!(
"updated={}, can_stop={}, iterations={}, top_matches=[{}]",
report.updated,
report.can_stop,
report.iterations.len(),
top_rules
)
}
fn schedule_for_log(sched: &ResolvedSchedule) -> String {
Self::truncate_for_log(&sched.to_string(), 160)
}
fn truncate_for_log(s: &str, limit: usize) -> String {
let mut s = s.replace('\n', " ");
if s.len() > limit {
s.truncate(limit);
s.push_str("...");
}
s
}
pub fn step_rules(&mut self, ruleset: &str) -> Result<RunReport, Error> {
fn collect_rule_ids(
ruleset: &str,
rulesets: &IndexMap<String, Ruleset>,
ids: &mut Vec<egglog_bridge::RuleId>,
) {
match &rulesets[ruleset] {
Ruleset::Rules(rules) => {
for (_, id) in rules.values() {
ids.push(*id);
}
}
Ruleset::Combined(sub_rulesets) => {
for sub_ruleset in sub_rulesets {
collect_rule_ids(sub_ruleset, rulesets, ids);
}
}
}
}
let mut rule_ids = Vec::new();
collect_rule_ids(ruleset, &self.rulesets, &mut rule_ids);
let iteration_report = self
.backend
.run_rules(&rule_ids, Some(&self.type_info))
.map_err(|e| Error::BackendError(e.to_string()))?;
Ok(RunReport::singleton(ruleset, iteration_report))
}
fn add_rule(&mut self, rule: ast::ResolvedRule) -> Result<String, Error> {
let core_rule = rule.to_canonicalized_core_rule(
&self.type_info,
&mut self.parser.symbol_gen,
self.proof_state.original_typechecking.is_none(),
)?;
let (query, actions) = (&core_rule.body, &core_rule.head);
let seminaive = self.seminaive && !rule.eval_mode.is_naive();
let no_decomp = self.no_decomp || rule.no_decomp;
let requires_read_context = !seminaive
|| matches!(
rule.eval_mode,
RuleEvalMode::Naive | RuleEvalMode::UnsafeSeminaive
);
let rule_id = {
let mut rb = self.backend.new_rule(&rule.name, seminaive);
rb.set_no_decomp(no_decomp);
let mut translator =
BackendRule::new(rb, &self.functions, &self.type_info, requires_read_context);
translator.query(query, rule.include_subsumed)?;
translator.actions(actions)?;
translator.build()
};
if let Some(rules) = self.rulesets.get_mut(&rule.ruleset) {
match rules {
Ruleset::Rules(rules) => {
match rules.entry(rule.name.clone()) {
indexmap::map::Entry::Occupied(_) => {
return Err(Error::RuleAlreadyExists(rule.name, rule.span));
}
indexmap::map::Entry::Vacant(e) => e.insert((core_rule, rule_id)),
};
Ok(rule.name)
}
Ruleset::Combined(_) => Err(Error::CombinedRulesetError(rule.ruleset, rule.span)),
}
} else {
Err(Error::NoSuchRuleset(rule.ruleset, rule.span))
}
}
fn eval_actions(&mut self, actions: &ResolvedActions) -> Result<(), Error> {
let mut binding = IndexSet::default();
let mut ctx = CoreActionContext::new(
&self.type_info,
&mut binding,
&mut self.parser.symbol_gen,
self.proof_state.original_typechecking.is_none(),
);
let (actions, _) = actions.to_core_actions(&mut ctx)?;
let mut translator = BackendRule::new(
self.backend.new_rule("eval_actions", false),
&self.functions,
&self.type_info,
true, );
translator.actions(&actions)?;
let id = translator.build();
let result = self.backend.run_rules(&[id], Some(&self.type_info));
self.backend.free_rule(id);
match result {
Ok(_) => Ok(()),
Err(e) => Err(Error::BackendError(e.to_string())),
}
}
pub fn get_function_names(&self) -> Vec<String> {
self.functions.keys().cloned().collect()
}
pub fn functions_iter(&self) -> impl Iterator<Item = (&String, &Function)> {
self.functions.iter()
}
fn with_execution_state<R>(&self, f: impl FnOnce(&mut ExecutionState<'_>) -> R) -> R {
self.backend.with_execution_state(Some(&self.type_info), f)
}
fn with_execution_state_tracked<R>(
&self,
f: impl FnOnce(&mut ExecutionState<'_>) -> R,
) -> (R, bool) {
self.backend
.with_execution_state_tracked(Some(&self.type_info), f)
}
pub fn read<R>(&self, f: impl FnOnce(ReadState<'_, '_>) -> R) -> R {
let registry = self.backend.action_registry().clone();
let guard = registry.read().unwrap();
self.with_execution_state_tracked(|es| f(ReadState::wrap(es, &guard, Context::Read)))
.0
}
pub fn function_entries(
&self,
name: &str,
f: impl FnMut(FunctionEntry<'_>),
) -> Result<(), Error> {
self.read(|rs| rs.function_entries(name, f))
}
pub fn function_entries_while(
&self,
name: &str,
f: impl FnMut(FunctionEntry<'_>) -> bool,
) -> Result<(), Error> {
self.read(|rs| rs.function_entries_while(name, f))
}
pub fn constructor_enodes(&self, name: &str, f: impl FnMut(Enode<'_>)) -> Result<(), Error> {
self.read(|rs| rs.constructor_enodes(name, f))
}
pub fn constructor_enodes_while(
&self,
name: &str,
f: impl FnMut(Enode<'_>) -> bool,
) -> Result<(), Error> {
self.read(|rs| rs.constructor_enodes_while(name, f))
}
pub fn clear_function(&mut self, func_name: &str) -> Result<(), Error> {
let backend_id = self
.functions
.get(func_name)
.ok_or_else(|| TypeError::UnboundFunction(func_name.to_string(), span!()))?
.backend_id;
self.backend.clear_table(backend_id);
Ok(())
}
pub fn eval_expr(&mut self, expr: &Expr) -> Result<(ArcSort, Value), Error> {
let span = expr.span();
let command = Command::Action(Action::Expr(span.clone(), expr.clone()));
let resolved = self.resolve_command(command)?;
if self.are_proofs_enabled() {
self.proof_check_program
.extend(resolved.desugared_before_proofs);
}
let resolved_commands = resolved.desugared;
if resolved_commands.len() != 1 {
return Err(Error::BackendError(
"eval_expr expects a single resolved command".to_string(),
));
}
let Some(resolved_command) = resolved_commands.into_iter().next() else {
return Err(Error::BackendError(
"eval_expr expects a single resolved command".to_string(),
));
};
let resolved_expr = match resolved_command {
ResolvedNCommand::CoreAction(ResolvedAction::Expr(_, resolved_expr)) => resolved_expr,
cmd => {
return Err(Error::BackendError(format!(
"eval_expr: unexpected resolved command: {cmd:?}"
)));
}
};
let sort = resolved_expr.output_type();
let value = self.eval_resolved_expr(span, &resolved_expr)?;
Ok((sort, value))
}
pub fn typecheck_expr_with_bindings_and_output(
&mut self,
expr: &Expr,
bindings: &[(String, Span, ArcSort)],
output_sort: ArcSort,
context: Context,
) -> Result<ResolvedExpr, TypeError> {
let mut binding_map = IndexMap::default();
binding_map.reserve(bindings.len());
for (name, span, sort) in bindings {
if binding_map
.insert(name.as_str(), (span.clone(), sort.clone()))
.is_some()
{
return Err(TypeError::AlreadyDefined(name.clone(), span.clone()));
}
}
let resolved = self.type_info.typecheck_expr_with_output(
&mut self.parser.symbol_gen,
expr,
&binding_map,
output_sort,
context,
)?;
Ok(remove_globals::remove_globals_expr(resolved))
}
pub fn prepare_unstable_fn_targets_for_eval(
&mut self,
expr: &ResolvedExpr,
) -> Result<(ResolvedExpr, Vec<(String, Value)>), Error> {
let mut bindings = Vec::new();
let expr = self.prepare_unstable_fn_targets_for_eval_inner(expr, &mut bindings)?;
Ok((expr, bindings))
}
fn prepare_unstable_fn_targets_for_eval_inner(
&mut self,
expr: &ResolvedExpr,
bindings: &mut Vec<(String, Value)>,
) -> Result<ResolvedExpr, Error> {
match expr {
ResolvedExpr::Lit(..) | ResolvedExpr::Var(..) => Ok(expr.clone()),
ResolvedExpr::Call(span, resolved_call, children) => {
if let ResolvedCall::Primitive(prim) = resolved_call
&& prim.name() == "unstable-fn"
{
let Some(ResolvedExpr::Lit(target_span, Literal::String(name))) =
children.first()
else {
return Err(Error::BackendError(format!(
"{}\nunstable-fn requires a literal string function name",
children
.first()
.map(ResolvedExpr::span)
.unwrap_or_else(|| Span::Panic)
)));
};
let panic_id = self.backend.new_panic(format!(
"unstable-fn over `{name}` was applied in a context where its wrapped \
function is not valid for this call site, if in a rule, add :naive."
));
let resolved_function = resolve_function_container_target_with_context(
&self.backend,
&self.functions,
&self.type_info,
name,
prim,
panic_id,
crate::Context::Full,
target_span,
)?;
let fn_value = self.backend.base_values().get(resolved_function);
let binding_name = self.parser.symbol_gen.fresh("unstable_fn_target");
bindings.push((binding_name.clone(), fn_value));
let mut prepared_children = Vec::with_capacity(children.len());
prepared_children.push(ResolvedExpr::Var(
target_span.clone(),
ResolvedVar {
name: binding_name,
sort: children[0].output_type(),
is_global_ref: false,
},
));
for child in &children[1..] {
prepared_children.push(
self.prepare_unstable_fn_targets_for_eval_inner(child, bindings)?,
);
}
return Ok(ResolvedExpr::Call(
span.clone(),
resolved_call.clone(),
prepared_children,
));
}
let prepared_children = children
.iter()
.map(|child| self.prepare_unstable_fn_targets_for_eval_inner(child, bindings))
.collect::<Result<Vec<_>, _>>()?;
Ok(ResolvedExpr::Call(
span.clone(),
resolved_call.clone(),
prepared_children,
))
}
}
}
fn eval_resolved_expr(&mut self, span: Span, expr: &ResolvedExpr) -> Result<Value, Error> {
let unit_id = self.backend.base_values().get_ty::<()>();
let unit_val = self.backend.base_values().get(());
let result: egglog_bridge::SideChannel<Value> = Default::default();
let result_ref = result.clone();
let ext_id = self
.backend
.register_external_func(Box::new(make_external_func(move |_es, vals| {
debug_assert!(vals.len() == 1);
*result_ref.lock().unwrap() = Some(vals[0]);
Some(unit_val)
})));
let mut translator = BackendRule::new(
self.backend.new_rule("eval_resolved_expr", false),
&self.functions,
&self.type_info,
true, );
let result_var = ResolvedVar {
name: self.parser.symbol_gen.fresh("eval_resolved_expr"),
sort: expr.output_type(),
is_global_ref: false,
};
let actions = ResolvedActions::singleton(ResolvedAction::Let(
span.clone(),
result_var.clone(),
expr.clone(),
));
let mut binding = IndexSet::default();
let mut ctx = CoreActionContext::new(
&self.type_info,
&mut binding,
&mut self.parser.symbol_gen,
self.proof_state.original_typechecking.is_none(),
);
let actions = actions.to_core_actions(&mut ctx)?.0;
translator.actions(&actions)?;
let arg = translator.entry(&ResolvedAtomTerm::Var(span.clone(), result_var));
translator.rb.call_external_func(
ext_id,
&[arg],
egglog_bridge::ColumnTy::Base(unit_id),
|| "this function will never panic".to_string(),
);
let id = translator.build();
let rule_result = self.backend.run_rules(&[id], Some(&self.type_info));
self.backend.free_rule(id);
self.backend.free_external_func(ext_id);
let _ = rule_result.map_err(|e| {
Error::BackendError(format!("Failed to evaluate expression '{expr}': {e}"))
})?;
let result = result.lock().unwrap().unwrap();
Ok(result)
}
fn add_combined_ruleset(&mut self, name: String, rulesets: Vec<String>) {
match self.rulesets.entry(name.clone()) {
Entry::Occupied(_) => panic!("Ruleset '{name}' was already present"),
Entry::Vacant(e) => e.insert(Ruleset::Combined(rulesets)),
};
}
fn add_ruleset(&mut self, name: String) {
match self.rulesets.entry(name.clone()) {
Entry::Occupied(_) => panic!("Ruleset '{name}' was already present"),
Entry::Vacant(e) => e.insert(Ruleset::Rules(Default::default())),
};
}
fn check_facts(&mut self, span: &Span, facts: &[ResolvedFact]) -> Result<(), Error> {
let fresh_name = self.parser.symbol_gen.fresh("check_facts");
let fresh_ruleset = self.parser.symbol_gen.fresh("check_facts_ruleset");
let rule = ast::ResolvedRule {
span: span.clone(),
head: ResolvedActions::default(),
body: facts.to_vec(),
name: fresh_name.clone(),
ruleset: fresh_ruleset.clone(),
eval_mode: RuleEvalMode::default(),
no_decomp: false,
include_subsumed: false,
};
let core_rule = rule.to_canonicalized_core_rule(
&self.type_info,
&mut self.parser.symbol_gen,
self.proof_state.original_typechecking.is_none(),
)?;
let query = core_rule.body;
let ext_sc = egglog_bridge::SideChannel::default();
let ext_sc_ref = ext_sc.clone();
let ext_id = self
.backend
.register_external_func(Box::new(make_external_func(move |exec_state, _| {
*ext_sc_ref.lock().unwrap() = Some(());
exec_state.trigger_early_stop();
Some(Value::new_const(0))
})));
let mut translator = BackendRule::new(
self.backend.new_rule("check_facts", false),
&self.functions,
&self.type_info,
true, );
translator.query(&query, true)?;
translator
.rb
.call_external_func(ext_id, &[], egglog_bridge::ColumnTy::Id, || {
"this function will never panic".to_string()
});
let id = translator.build();
let run_result = self.backend.run_rules(&[id], Some(&self.type_info));
self.backend.free_rule(id);
self.backend.free_external_func(ext_id);
run_result.map_err(|e| Error::BackendError(e.to_string()))?;
let ext_sc_val = ext_sc.lock().unwrap().take();
let matched = matches!(ext_sc_val, Some(()));
if !matched {
Err(Error::CheckError(
facts.iter().map(|f| f.clone().make_unresolved()).collect(),
span.clone(),
))
} else {
Ok(())
}
}
fn run_command(&mut self, command: ResolvedNCommand) -> Result<Vec<CommandOutput>, Error> {
match command {
ResolvedNCommand::Sort {
name,
uf,
proof_func,
proof_constructors,
..
} => {
if let Some((uf_ctor, uf_index)) = uf {
self.proof_state.uf_parent.insert(name.clone(), uf_ctor);
if let Some(uf_index) = uf_index {
self.proof_state.uf_function.insert(name.clone(), uf_index);
}
}
if let Some(proof_func_name) = proof_func {
self.proof_state
.proof_func_parent
.insert(name.clone(), proof_func_name);
}
if let Some(pc) = proof_constructors {
let names = &mut self.proof_state.proof_names;
names.proof_datatype = name.clone();
names.congr_constructor = pc.congr;
names.eq_trans_constructor = pc.trans;
names.eq_sym_constructor = pc.sym;
names.container_normalize_constructor = pc.normalize;
}
log::info!("Declared sort {name}.")
}
ResolvedNCommand::Function(fdecl) => {
self.declare_function(&fdecl)?;
log::info!("Declared {} {}.", fdecl.subtype, fdecl.name)
}
ResolvedNCommand::AddRuleset(_span, name) => {
self.add_ruleset(name.clone());
log::info!("Declared ruleset {name}.");
}
ResolvedNCommand::UnstableCombinedRuleset(_span, name, others) => {
self.add_combined_ruleset(name.clone(), others);
log::info!("Declared ruleset {name}.");
}
ResolvedNCommand::NormRule { rule } => {
let name = rule.name.clone();
self.add_rule(rule)?;
log::info!("Declared rule {name}.")
}
ResolvedNCommand::RunSchedule(sched) => {
let report = self.run_schedule(&sched)?;
log::info!("Ran schedule {sched}.");
log::info!("Report: {report}");
self.overall_run_report.union(report.clone());
return Ok(vec![CommandOutput::RunSchedule(report)]);
}
ResolvedNCommand::PrintOverallStatistics(span, file) => match file {
None => {
log::info!("Printed overall statistics");
return Ok(vec![CommandOutput::OverallStatistics(
self.overall_run_report.clone(),
)]);
}
Some(path) => {
let mut file = std::fs::File::create(&path)
.map_err(|e| Error::IoError(path.clone().into(), e, span.clone()))?;
log::info!("Printed overall statistics to json file {path}");
serde_json::to_writer(&mut file, &self.overall_run_report).map_err(|e| {
Error::BackendError(format!("failed writing statistics: {e}"))
})?;
}
},
ResolvedNCommand::Check(span, facts) => {
self.check_facts(&span, &facts)?;
log::info!("Checked fact {facts:?}.");
}
ResolvedNCommand::CoreAction(action) => match &action {
ResolvedAction::Let(_, name, contents) => {
panic!("Globals should have been desugared away: {name} = {contents}")
}
_ => {
self.eval_actions(&ResolvedActions::new(vec![action.clone()]))?;
}
},
ResolvedNCommand::Extract(span, expr, variants) => {
let sort = expr.output_type();
let x = self.eval_resolved_expr(span.clone(), &expr)?;
let n = self.eval_resolved_expr(span, &variants)?;
let n: i64 = self.backend.base_values().unwrap(n);
let mut termdag = TermDag::default();
let extractor = Extractor::compute_costs_from_rootsorts(
Some(vec![sort]),
self,
TreeAdditiveCostModel::default(),
);
return if n == 0 {
if let Some((cost, term)) = extractor.extract_best(self, &mut termdag, x) {
if log_enabled!(Level::Info) {
log::info!("extracted with cost {cost}: {}", termdag.to_string(term));
}
Ok(vec![CommandOutput::ExtractBest(termdag, cost, term)])
} else {
Err(Error::ExtractError(
"Unable to find any valid extraction (likely due to subsume or delete)"
.to_string(),
))
}
} else {
if n < 0 {
return Err(Error::ExtractError(
"cannot extract a negative number of variants".to_string(),
));
}
let terms: Vec<TermId> = extractor
.extract_variants(self, &mut termdag, x, n as usize)
.iter()
.map(|e| e.1)
.collect();
if log_enabled!(Level::Info) {
let expr_str = expr.to_string();
log::info!("extracted {} variants for {expr_str}", terms.len());
}
Ok(vec![CommandOutput::ExtractVariants(termdag, terms)])
};
}
ResolvedNCommand::Push(n) => {
(0..n).for_each(|_| self.push());
log::info!("Pushed {n} levels.")
}
ResolvedNCommand::Pop(span, n) => {
for _ in 0..n {
self.pop().map_err(|err| {
if let Error::Pop(_) = err {
Error::Pop(span.clone())
} else {
err
}
})?;
}
log::info!("Popped {n} levels.")
}
ResolvedNCommand::PrintFunction(span, f, n, file, mode) => {
let file = file
.map(|file| {
let path: PathBuf = file.into();
match std::fs::File::create(&path) {
Ok(f) => Ok((f, path)),
Err(e) => Err(Error::IoError(path, e, span.clone())),
}
})
.transpose()?;
return self
.print_function(&f, n, file, span.clone(), mode)
.map_err(|e| match e {
Error::TypeError(TypeError::UnboundFunction(f, _)) => {
Error::TypeError(TypeError::UnboundFunction(f, span.clone()))
}
_ => e,
})
.map(|opt| opt.into_iter().collect());
}
ResolvedNCommand::PrintSize(span, f) => {
let res = self.print_size(f.as_deref()).map_err(|e| match e {
Error::TypeError(TypeError::UnboundFunction(f, _)) => {
Error::TypeError(TypeError::UnboundFunction(f, span.clone()))
}
_ => e,
})?;
return Ok(vec![res]);
}
ResolvedNCommand::Fail(span, c) => {
let result = self.run_command(*c);
if let Err(e) = result {
log::info!("Command failed as expected: {e}");
} else {
return Err(Error::ExpectFail(span));
}
}
ResolvedNCommand::Input { span, name, file } => {
self.input_file(&name, file, span)?;
}
ResolvedNCommand::Output { span, file, exprs } => {
let mut filename = self.fact_directory.clone().unwrap_or_default();
filename.push(file.as_str());
let mut f = File::options()
.append(true)
.create(true)
.open(&filename)
.map_err(|e| Error::IoError(filename.clone(), e, span.clone()))?;
let extractor = Extractor::compute_costs_from_rootsorts(
None,
self,
TreeAdditiveCostModel::default(),
);
let mut termdag: TermDag = Default::default();
use std::io::Write;
for expr in exprs {
let value = self.eval_resolved_expr(span.clone(), &expr)?;
let expr_type = expr.output_type();
let term = match extractor.extract_best_with_sort(
self,
&mut termdag,
value,
expr_type,
) {
Some((_, term)) => term,
None => return Err(Error::ExtractError(
"Unable to find any valid extraction (likely due to subsume or delete)"
.to_string(),
)),
};
writeln!(f, "{}", termdag.to_string(term))
.map_err(|e| Error::IoError(filename.clone(), e, span.clone()))?;
}
log::info!("Output to '{filename:?}'.")
}
ResolvedNCommand::UserDefined(_span, name, exprs) => {
let command = self
.commands
.get(&name)
.ok_or_else(|| {
NotFoundError(format!("Unrecognized user-defined command: {name}"))
})?
.clone();
return command.update(self, &exprs);
}
ResolvedNCommand::ProveExists(span, resolved_call) => {
let mut instrument = ProofInstrumentor { egraph: self };
let (proof_store, proof_id) =
instrument
.prove_exists(&resolved_call)
.map_err(|error| Error::ProofError {
span: span.clone(),
error,
})?;
return Ok(vec![CommandOutput::ProveExists {
proof_store,
proof_id,
}]);
}
};
Ok(vec![])
}
fn input_file(&mut self, func_name: &str, file: String, span: Span) -> Result<(), Error> {
let function_type = self.type_info.get_func_type(func_name).ok_or_else(|| {
Error::TypeError(TypeError::UnboundFunction(
func_name.to_string(),
span.clone(),
))
})?;
let func = self.functions.get_mut(func_name).unwrap();
let mut filename = self.fact_directory.clone().unwrap_or_default();
filename.push(file.as_str());
for t in &func.func_type.input {
match t.name() {
"i64" | "f64" | "String" => {}
s => return Err(Error::UnsupportedInputType(s.to_string(), span.clone())),
}
}
if function_type.subtype != FunctionSubtype::Constructor {
match func.func_type.output.name() {
"i64" | "String" | "Unit" => {}
s => return Err(Error::UnsupportedInputType(s.to_string(), span.clone())),
}
}
log::info!("Opening file '{filename:?}'...");
let mut f =
File::open(&filename).map_err(|e| Error::IoError(filename.clone(), e, span.clone()))?;
let mut contents = String::new();
f.read_to_string(&mut contents)
.map_err(|e| Error::IoError(filename.clone(), e, span.clone()))?;
let mut parsed_contents: Vec<Vec<Value>> = Vec::with_capacity(contents.lines().count());
let mut row_schema = func.func_type.input.clone();
if function_type.subtype == FunctionSubtype::Custom {
row_schema.push(func.func_type.output.clone());
}
log::debug!("{row_schema:?}");
let unit_val = self.backend.base_values().get(());
for line in contents.lines() {
let mut it = line.split('\t').map(|s| s.trim());
let mut row: Vec<Value> = Vec::with_capacity(row_schema.len());
for sort in row_schema.iter() {
if let Some(raw) = it.next() {
let val = match sort.name() {
"i64" => {
if let Ok(i) = raw.parse::<i64>() {
self.backend.base_values().get(i)
} else {
return Err(Error::InputFileFormatError(file));
}
}
"f64" => {
if let Ok(f) = raw.parse::<f64>() {
self.backend
.base_values()
.get::<F>(core_relations::Boxed::new(f.into()))
} else {
return Err(Error::InputFileFormatError(file));
}
}
"String" => self.backend.base_values().get::<S>(raw.to_string().into()),
"Unit" => unit_val,
_ => panic!("Unreachable"),
};
row.push(val);
} else {
break;
}
}
if row.is_empty() {
continue;
}
if row.len() != row_schema.len() || it.next().is_some() {
return Err(Error::InputFileFormatError(file));
}
parsed_contents.push(row);
}
log::debug!("Successfully loaded file.");
let num_facts = parsed_contents.len();
let table_action = egglog_bridge::TableAction::new(&self.backend, func.backend_id);
if function_type.subtype != FunctionSubtype::Constructor {
self.with_execution_state(|es| {
for row in parsed_contents.iter() {
table_action.insert(es, row.iter().copied());
}
Some(unit_val)
});
} else {
self.with_execution_state(|es| {
for row in parsed_contents.iter() {
table_action.lookup_or_insert(es, row);
}
Some(unit_val)
});
}
self.backend.flush_updates();
log::info!("Read {num_facts} facts into {func_name} from '{file}'.");
Ok(())
}
pub fn are_proofs_enabled(&self) -> bool {
self.proof_state.proofs_enabled
}
fn resolve_command_before_proofs(
&mut self,
command: Command,
) -> Result<Vec<ResolvedNCommand>, Error> {
let desugared = desugar_command(command, &mut self.parser, self.proof_state.proof_testing)?;
if let Some(original_typechecking) = self.proof_state.original_typechecking.as_mut() {
let typechecked = original_typechecking.typecheck_program(&desugared)?;
for command in &typechecked {
if let Err(reason) = command_supports_proof_encoding(
&command.to_command(),
&original_typechecking.type_info,
) {
let command_text = format!("{}", command.to_command());
return Err(Error::UnsupportedProofCommand {
command: command_text,
reason,
});
}
}
Ok(proof_form(typechecked, &mut self.parser.symbol_gen))
} else {
let mut typechecked = self.typecheck_program(&desugared)?;
typechecked = remove_globals::remove_globals(typechecked, &mut self.parser.symbol_gen);
for command in &typechecked {
self.names.check_shadowing(command)?;
}
Ok(typechecked)
}
}
fn resolve_command(&mut self, command: Command) -> Result<ResolvedNCommands, Error> {
let resolved_before_proofs = self.resolve_command_before_proofs(command)?;
if self.proof_state.original_typechecking.is_none() {
Ok(ResolvedNCommands {
desugared: resolved_before_proofs,
desugared_before_proofs: vec![],
})
} else {
let typechecked_no_globals = proof_global_remover::remove_globals(
resolved_before_proofs.clone(),
&mut self.parser.symbol_gen,
);
for command in &typechecked_no_globals {
self.names.check_shadowing(command)?;
}
let term_encoding_added =
ProofInstrumentor::add_term_encoding(self, typechecked_no_globals);
let mut new_typechecked = vec![];
for new_cmd in term_encoding_added {
let desugared =
desugar_command(new_cmd, &mut self.parser, self.proof_state.proof_testing)?;
for cmd in &desugared {
log::trace!("Desugared term encoding: {}", cmd.to_command());
}
let desugared_typechecked = self.typecheck_program(&desugared)?;
let desugared_typechecked = remove_globals::remove_globals(
desugared_typechecked,
&mut self.parser.symbol_gen,
);
new_typechecked.extend(desugared_typechecked);
}
Ok(ResolvedNCommands {
desugared: new_typechecked,
desugared_before_proofs: resolved_before_proofs,
})
}
}
fn process_program_internal(
&mut self,
program: Vec<Command>,
run_commands: bool,
) -> Result<ResolvedNCommandsWithOutput, Error> {
let mut outputs = Vec::new();
let mut desugared_before_proofs = Vec::new();
let mut desugared = Vec::new();
for before_expanded_command in program {
let macro_type_info = self
.proof_state
.original_typechecking
.as_ref()
.map(|egraph| &egraph.type_info)
.unwrap_or(&self.type_info);
let macro_expanded = self.command_macros.apply(
before_expanded_command,
&mut self.parser.symbol_gen,
macro_type_info,
)?;
for command in macro_expanded {
if let Command::Include(span, file) = &command {
let s = std::fs::read_to_string(file)
.map_err(|e| Error::IoError(file.clone().into(), e, span.clone()))?;
let included_program = self
.parser
.get_program_from_string(Some(file.clone()), &s)?;
let resolved = self.process_program_internal(included_program, run_commands)?;
outputs.extend(resolved.outputs);
desugared.extend(resolved.resolved);
desugared_before_proofs.extend(resolved.resolved_before_proofs);
} else {
let resolved = self.resolve_command(command)?;
if run_commands && self.are_proofs_enabled() {
self.proof_check_program
.extend(resolved.desugared_before_proofs.clone());
}
desugared_before_proofs.extend(resolved.desugared_before_proofs);
desugared.extend(resolved.desugared.clone());
for processed in resolved.desugared {
if run_commands
|| matches!(
processed,
ResolvedNCommand::Push(_) | ResolvedNCommand::Pop(_, _)
)
{
let result = self.run_command(processed)?;
outputs.extend(result);
}
}
}
}
}
Ok(ResolvedNCommandsWithOutput {
outputs,
resolved_before_proofs: desugared_before_proofs,
resolved: desugared,
})
}
pub fn run_program(&mut self, program: Vec<Command>) -> Result<Vec<CommandOutput>, Error> {
let res = self.process_program_internal(program, true)?;
Ok(res.outputs)
}
pub fn resolve_program(
&mut self,
filename: Option<String>,
input: &str,
) -> Result<Vec<ResolvedCommand>, Error> {
let parsed = self.parser.get_program_from_string(filename, input)?;
let res = self.process_program_internal(parsed, false)?;
Ok(res.resolved.into_iter().map(|c| c.to_command()).collect())
}
pub fn parse_program(
&mut self,
filename: Option<String>,
input: &str,
) -> Result<Vec<Command>, Error> {
let parsed = self.parser.get_program_from_string(filename, input)?;
Ok(parsed)
}
pub fn parse_and_run_program(
&mut self,
filename: Option<String>,
input: &str,
) -> Result<Vec<CommandOutput>, Error> {
let parsed = self.parser.get_program_from_string(filename, input)?;
self.run_program(parsed)
}
pub fn num_tuples(&self) -> usize {
self.functions
.values()
.map(|f| self.backend.table_size(f.backend_id))
.sum()
}
pub fn get_sort<S: Sort>(&self) -> Arc<S> {
self.type_info.get_sort()
}
pub fn get_sort_by<S: Sort>(&self, f: impl Fn(&Arc<S>) -> bool) -> Arc<S> {
self.type_info.get_sort_by(f)
}
pub fn get_sorts<S: Sort>(&self) -> Vec<Arc<S>> {
self.type_info.get_sorts()
}
pub fn get_sorts_by<S: Sort>(&self, f: impl Fn(&Arc<S>) -> bool) -> Vec<Arc<S>> {
self.type_info.get_sorts_by(f)
}
pub fn get_arcsort_by(&self, f: impl Fn(&ArcSort) -> bool) -> ArcSort {
self.type_info.get_arcsort_by(f)
}
pub fn get_arcsort_for_value_type<T: 'static>(&self) -> ArcSort {
self.type_info.get_arcsort_for_value_type::<T>()
}
pub fn get_arcsorts_by(&self, f: impl Fn(&ArcSort) -> bool) -> Vec<ArcSort> {
self.type_info.get_arcsorts_by(f)
}
pub fn get_sort_by_name(&self, sym: &str) -> Option<&ArcSort> {
self.type_info.get_sort_by_name(sym)
}
pub fn get_overall_run_report(&self) -> &RunReport {
&self.overall_run_report
}
pub fn value_to_base<T: BaseValue>(&self, x: Value) -> T {
self.backend.base_values().unwrap::<T>(x)
}
pub fn base_to_value<T: BaseValue>(&self, x: T) -> Value {
self.backend.base_values().get::<T>(x)
}
pub fn value_to_container<T: ContainerValue>(
&self,
x: Value,
) -> Option<impl Deref<Target = T>> {
self.backend.container_values().get_val::<T>(x)
}
pub fn container_to_value<T: ContainerValue>(&mut self, x: T) -> Value {
self.with_execution_state(|state| {
self.backend.container_values().register_val::<T>(x, state)
})
}
pub fn get_size(&self, func: &str) -> usize {
let function_id = self.functions.get(func).unwrap().backend_id;
self.backend.table_size(function_id)
}
pub fn get_function(&self, name: &str) -> Option<&Function> {
self.functions.get(name)
}
pub fn has_command(&self, name: &str) -> bool {
self.commands.contains_key(name)
}
pub fn run_user_defined_command(
&mut self,
name: &str,
args: &[Expr],
) -> Result<Vec<CommandOutput>, Error> {
self.run_command(ResolvedNCommand::UserDefined(
span!(),
name.to_string(),
args.to_vec(),
))
}
pub fn set_report_level(&mut self, level: ReportLevel) {
self.backend.set_report_level(level);
}
pub fn dump_debug_info(&self) {
self.backend.dump_debug_info();
}
pub fn update<R>(
&mut self,
f: impl FnOnce(FullState<'_, '_>) -> Result<R, Error>,
) -> Result<R, Error> {
if self.are_proofs_enabled() {
return Err(Error::ProofsIncompatibleApi {
api: "EGraph::update",
reason: "writes inside the closure bypass the proof-encoding pipeline,\n\
so any rule derivations resting on them would be unverifiable.",
});
}
self.update_unchecked(f)
}
pub(crate) fn update_unchecked<R>(
&mut self,
f: impl FnOnce(FullState<'_, '_>) -> Result<R, Error>,
) -> Result<R, Error> {
let registry = self.backend.action_registry().clone();
let guard = registry.read().unwrap();
let (result, changed) =
self.with_execution_state_tracked(|es| f(FullState::wrap(es, &guard, Context::Full)));
drop(guard);
if changed {
self.backend.flush_updates();
}
result
}
pub fn query(
&mut self,
vars: &[(&str, ArcSort)],
facts: ast::Facts<String, String>,
) -> Result<Vec<HashMap<String, Value>>, Error> {
if self.are_proofs_enabled() {
return Err(Error::ProofsIncompatibleApi {
api: "EGraph::query",
reason: "the underlying rust_rule callback has no proof-encoding validator,\n\
so query matches cannot be verified.",
});
}
use std::sync::{Arc, Mutex};
let names: Arc<[String]> = vars.iter().map(|(n, _)| (*n).to_owned()).collect();
let results: Arc<Mutex<Vec<HashMap<String, Value>>>> = Arc::new(Mutex::new(Vec::new()));
let results_weak = Arc::downgrade(&results);
let names_for_cb = names.clone();
let ruleset = self.parser.symbol_gen.fresh("query_ruleset");
prelude::add_ruleset(self, &ruleset)?;
let outcome = (|| -> Result<_, Error> {
prelude::rust_rule(self, "query", &ruleset, vars, facts, move |_, values| {
let arc = results_weak.upgrade().unwrap();
let mut results = arc.lock().unwrap();
let map: HashMap<String, Value> = names_for_cb
.iter()
.zip(values.iter().copied())
.map(|(n, v)| (n.clone(), v))
.collect();
results.push(map);
Some(())
})?;
prelude::run_ruleset(self, &ruleset)?;
Ok(())
})();
if let Some(Ruleset::Rules(rules)) = self.rulesets.swap_remove(&ruleset) {
for (_, rule) in rules {
self.backend.free_rule(rule.1);
}
}
outcome?;
let Some(mutex) = Arc::into_inner(results) else {
panic!("`results_weak` outlived the callback");
};
Ok(mutex.into_inner().unwrap())
}
}
pub use crate::api::{ApiError, FromValue, FromValues, IntoValue, IntoValues, RawValues};
#[allow(clippy::too_many_arguments)]
fn resolve_function_container_target_with_context(
backend: &egglog_bridge::EGraph,
functions: &IndexMap<String, Function>,
type_info: &TypeInfo,
name: &str,
primitive: &core::SpecializedPrimitive,
panic_id: ExternalFunctionId,
ctx: crate::Context,
span: &Span,
) -> Result<ResolvedFunction, Error> {
let Some(target_function) = type_info
.get_sorts::<FunctionSort>()
.into_iter()
.find(|function| function.name() == primitive.output().name())
else {
return Err(Error::BackendError(format!(
"`unstable-fn` output sort `{}` is not a function sort",
primitive.output().name()
)));
};
let partial_arcsorts: Vec<_> = primitive.input().iter().skip(1).cloned().collect();
let remaining_inputs = target_function.inputs();
let output = target_function.output();
let id = if let Some(func) = functions.get(name) {
let func_type = type_info.get_func_type(name).ok_or_else(|| {
Error::BackendError(format!(
"`unstable-fn` references `{name}`, which has no resolved type"
))
})?;
let expected_inputs = partial_arcsorts
.iter()
.chain(remaining_inputs)
.collect::<Vec<_>>();
let inputs_match = func_type.input.len() == expected_inputs.len()
&& func_type
.input
.iter()
.zip(&expected_inputs)
.all(|(actual, expected)| actual.name() == expected.name());
if !inputs_match || func_type.output.name() != output.name() {
let expected_input_names = expected_inputs
.iter()
.map(|sort| sort.name())
.collect::<Vec<_>>()
.join(", ");
let actual_input_names = func_type
.input
.iter()
.map(|sort| sort.name())
.collect::<Vec<_>>()
.join(", ");
return Err(Error::BackendError(format!(
"`unstable-fn` reference `{name}` expected ({}) -> {}, found ({}) -> {}",
expected_input_names,
output.name(),
actual_input_names,
func_type.output.name(),
)));
}
let action = egglog_bridge::TableAction::new(backend, func.backend_id);
match func_type.subtype {
ast::FunctionSubtype::Constructor => ResolvedFunctionId::Constructor(action),
ast::FunctionSubtype::Custom => ResolvedFunctionId::Function(action),
}
} else if let Some(primitives) = type_info.get_prims(name) {
let signature: Vec<_> = partial_arcsorts
.iter()
.chain(remaining_inputs)
.chain(once(&output))
.cloned()
.collect();
let candidates: Vec<_> = primitives
.iter()
.filter(|primitive| primitive.accept(&signature, type_info))
.collect();
let mut ambiguous_ctx = None;
let context_ids = enum_map::EnumMap::from_fn(|runtime_ctx| {
let mut ids = candidates
.iter()
.filter_map(|primitive| primitive.context_ids[runtime_ctx]);
match (ids.next(), ids.next()) {
(None, _) => None,
(Some(id), None) => Some(id),
(Some(_), Some(_)) => {
ambiguous_ctx = Some(runtime_ctx);
None
}
}
});
if let Some(runtime_ctx) = ambiguous_ctx {
return Err(TypeError::AmbiguousPrimitive {
name: name.to_owned(),
ctx: runtime_ctx,
span: span.clone(),
}
.into());
}
if !context_ids.iter().any(|(_, id)| id.is_some()) {
return Err(TypeError::UnresolvedPrimitive {
name: name.to_owned(),
ctx,
span: span.clone(),
}
.into());
}
ResolvedFunctionId::Primitive { context_ids }
} else {
return Err(TypeError::UnresolvedPrimitive {
name: name.to_owned(),
ctx,
span: span.clone(),
}
.into());
};
Ok(ResolvedFunction {
id,
partial_arcsorts,
name: name.to_owned(),
panic_id,
})
}
struct BackendRule<'a> {
rb: egglog_bridge::RuleBuilder<'a>,
entries: HashMap<core::ResolvedAtomTerm, QueryEntry>,
functions: &'a IndexMap<String, Function>,
type_info: &'a TypeInfo,
requires_read_context: bool,
}
impl<'a> BackendRule<'a> {
fn new(
rb: egglog_bridge::RuleBuilder<'a>,
functions: &'a IndexMap<String, Function>,
type_info: &'a TypeInfo,
requires_read_context: bool,
) -> BackendRule<'a> {
BackendRule {
rb,
functions,
type_info,
requires_read_context,
entries: Default::default(),
}
}
fn query_context(&self) -> crate::Context {
if self.requires_read_context {
crate::Context::Read
} else {
crate::Context::Pure
}
}
fn action_context(&self) -> crate::Context {
if self.requires_read_context {
crate::Context::Full
} else {
crate::Context::Write
}
}
fn entry(&mut self, x: &core::ResolvedAtomTerm) -> QueryEntry {
self.entries
.entry(x.clone())
.or_insert_with(|| match x {
core::GenericAtomTerm::Var(_, v) => self
.rb
.new_var_named(v.sort.column_ty(self.rb.egraph()), &v.name),
core::GenericAtomTerm::Literal(_, l) => literal_to_entry(self.rb.egraph(), l),
core::GenericAtomTerm::Global(..) => {
panic!("Globals should have been desugared")
}
})
.clone()
}
fn func(&self, f: &typechecking::FuncType) -> egglog_bridge::FunctionId {
self.functions[&f.name].backend_id
}
fn prim(
&mut self,
prim: &core::SpecializedPrimitive,
args: &[core::ResolvedAtomTerm],
ctx: crate::Context,
) -> Result<(ExternalFunctionId, Vec<QueryEntry>, ColumnTy), Error> {
let resolved_id = prim.external_id(ctx);
let mut qe_args = self.args(args);
if prim.name() == "unstable-fn" {
let core::ResolvedAtomTerm::Literal(ref span, Literal::String(ref name)) = args[0]
else {
return Err(Error::BackendError(
"expected a string literal as the first argument to `unstable-fn`".to_string(),
));
};
let panic_id = self.rb.new_panic(format!(
"unstable-fn over `{name}` was applied in a context where its wrapped \
function is not valid for this call site, if in a rule, add :naive."
));
let resolved = resolve_function_container_target_with_context(
self.rb.egraph(),
self.functions,
self.type_info,
name,
prim,
panic_id,
ctx,
span,
)?;
qe_args[0] = self.rb.egraph().base_value_constant(resolved);
}
Ok((
resolved_id,
qe_args,
prim.output().column_ty(self.rb.egraph()),
))
}
fn args<'b>(
&mut self,
args: impl IntoIterator<Item = &'b core::ResolvedAtomTerm>,
) -> Vec<QueryEntry> {
args.into_iter().map(|x| self.entry(x)).collect()
}
fn query(
&mut self,
query: &core::Query<ResolvedCall, ResolvedVar>,
include_subsumed: bool,
) -> Result<(), Error> {
for atom in &query.atoms {
match &atom.head {
ResolvedCall::Func(f) => {
let f = self.func(f);
let args = self.args(&atom.args);
let is_subsumed = match include_subsumed {
true => None,
false => Some(false),
};
self.rb.query_table(f, &args, is_subsumed).unwrap();
}
ResolvedCall::Primitive(p) => {
let ctx = self.query_context();
let (p, args, ty) = self.prim(p, &atom.args, ctx)?;
self.rb.query_prim(p, &args, ty).unwrap()
}
}
}
Ok(())
}
fn actions(&mut self, actions: &core::ResolvedCoreActions) -> Result<(), Error> {
for action in &actions.0 {
match action {
core::GenericCoreAction::Let(span, v, f, args) => {
let v = core::GenericAtomTerm::Var(span.clone(), v.clone());
let y = match f {
ResolvedCall::Func(f) => {
let name = f.name.clone();
let f = self.func(f);
let args = self.args(args);
let span = span.clone();
self.rb.lookup(f, &args, move || {
format!("{span}: lookup of function {name} failed")
})
}
ResolvedCall::Primitive(p) => {
let name = p.name().to_owned();
let ctx = self.action_context();
let (p, args, ty) = self.prim(p, args, ctx)?;
let span = span.clone();
self.rb.call_external_func(p, &args, ty, move || {
format!("{span}: call of primitive {name} failed")
})
}
};
self.entries.insert(v, y.into());
}
core::GenericCoreAction::LetAtomTerm(span, v, x) => {
let v = core::GenericAtomTerm::Var(span.clone(), v.clone());
let x = self.entry(x);
self.entries.insert(v, x);
}
core::GenericCoreAction::Set(_, f, xs, y) => match f {
ResolvedCall::Primitive(..) => panic!("runtime primitive set!"),
ResolvedCall::Func(f) => {
let f = self.func(f);
let args = self.args(xs.iter().chain([y]));
self.rb.set(f, &args)
}
},
core::GenericCoreAction::Change(span, change, f, args) => match f {
ResolvedCall::Primitive(..) => panic!("runtime primitive change!"),
ResolvedCall::Func(f) => {
let name = f.name.clone();
let can_subsume = self.functions[&f.name].can_subsume;
let f = self.func(f);
let args = self.args(args);
match change {
Change::Delete => self.rb.remove(f, &args),
Change::Subsume if can_subsume => self.rb.subsume(f, &args),
Change::Subsume => {
return Err(Error::SubsumeMergeError(name, span.clone()));
}
}
}
},
core::GenericCoreAction::Union(_, x, y) => {
let x = self.entry(x);
let y = self.entry(y);
self.rb.union(x, y)
}
core::GenericCoreAction::Panic(_, message) => self.rb.panic(message.clone()),
}
}
Ok(())
}
fn build(self) -> egglog_bridge::RuleId {
self.rb.build()
}
}
fn literal_to_entry(egraph: &egglog_bridge::EGraph, l: &Literal) -> QueryEntry {
match l {
Literal::Int(x) => egraph.base_value_constant::<i64>(*x),
Literal::Float(x) => egraph.base_value_constant::<sort::F>(x.into()),
Literal::String(x) => egraph.base_value_constant::<sort::S>(sort::S::new(x.clone())),
Literal::Bool(x) => egraph.base_value_constant::<bool>(*x),
Literal::Unit => egraph.base_value_constant::<()>(()),
}
}
fn literal_to_value(egraph: &egglog_bridge::EGraph, l: &Literal) -> Value {
match l {
Literal::Int(x) => egraph.base_values().get::<i64>(*x),
Literal::Float(x) => egraph.base_values().get::<sort::F>(x.into()),
Literal::String(x) => egraph.base_values().get::<sort::S>(sort::S::new(x.clone())),
Literal::Bool(x) => egraph.base_values().get::<bool>(*x),
Literal::Unit => egraph.base_values().get::<()>(()),
}
}
#[derive(Debug, Error)]
pub enum Error {
#[error(transparent)]
ParseError(#[from] ParseError),
#[error(transparent)]
NotFoundError(#[from] NotFoundError),
#[error(transparent)]
TypeError(#[from] TypeError),
#[error(transparent)]
ApiError(#[from] crate::api::ApiError),
#[error("Errors:\n{}", ListDisplay(.0, "\n"))]
TypeErrors(Vec<TypeError>),
#[error("{}\nCheck failed: \n{}", .1, ListDisplay(.0, "\n"))]
CheckError(Vec<Fact>, Span),
#[error("{1}\nNo such ruleset: {0}")]
NoSuchRuleset(String, Span),
#[error(
"{1}\nAttempted to add a rule to combined ruleset {0}. Combined rulesets may only depend on other rulesets."
)]
CombinedRulesetError(String, Span),
#[error("{0}")]
BackendError(String),
#[error("{0}\nTried to pop too much")]
Pop(Span),
#[error("{0}\nCommand should have failed.")]
ExpectFail(Span),
#[error("{2}\nIO error: {0}: {1}")]
IoError(PathBuf, std::io::Error, Span),
#[error("{1}\nCannot subsume function with merge: {0}")]
SubsumeMergeError(String, Span),
#[error("extraction failure: {:?}", .0)]
ExtractError(String),
#[error("{span}\n{error}")]
ProofError {
span: Span,
#[source]
error: ProveExistsError,
},
#[error("{1}\n{2}\nShadowing is not allowed, but found {0}")]
Shadowing(String, Span, Span),
#[error("{1}\nCommand already exists: {0}")]
CommandAlreadyExists(String, Span),
#[error("{1}\nRule already exists: {0}")]
RuleAlreadyExists(String, Span),
#[error("{1}\nUnsupported type {0} for input")]
UnsupportedInputType(String, Span),
#[error("{0}\n{1}")]
DesugarError(Span, String),
#[error("Incorrect format in file '{0}'.")]
InputFileFormatError(String),
#[error(
"Command is not supported by the current proof term encoding implementation.\n\
Reason: {reason}\n\
This typically means the command uses constructs that cannot yet be represented as proof terms.\n\
Consider disabling proof term encoding for this run or rewriting the command to avoid unsupported features.\n\
Offending command: {command}"
)]
UnsupportedProofCommand {
command: String,
reason: ProofEncodingUnsupportedReason,
},
#[error(
"`{api}` is incompatible with proof mode: {reason} \
Disable proofs or make the operation a command in the syntax of the egglog language and use `EGraph::parse_and_run`."
)]
ProofsIncompatibleApi {
api: &'static str,
reason: &'static str,
},
}
#[cfg(test)]
mod tests {
use crate::constraint::SimpleTypeConstraint;
use crate::*;
use crate::PureState;
#[derive(Clone)]
struct InnerProduct {
vec: ArcSort,
}
impl Primitive for InnerProduct {
fn name(&self) -> &str {
"inner-product"
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn crate::constraint::TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![self.vec.clone(), self.vec.clone(), I64Sort.to_arcsort()],
span.clone(),
)
.into_box()
}
}
impl PurePrim for InnerProduct {
fn apply<'a, 'db>(&self, state: PureState<'a, 'db>, args: &[Value]) -> Option<Value> {
let mut sum = 0;
let vec1 = state
.container_values()
.get_val::<VecContainer>(args[0])
.unwrap();
let vec2 = state
.container_values()
.get_val::<VecContainer>(args[1])
.unwrap();
assert_eq!(vec1.data.len(), vec2.data.len());
for (a, b) in vec1.data.iter().zip(vec2.data.iter()) {
let a = state.base_values().unwrap::<i64>(*a);
let b = state.base_values().unwrap::<i64>(*b);
sum += a * b;
}
Some(state.base_values().get::<i64>(sum))
}
}
#[derive(Clone)]
struct FullOnly;
impl Primitive for FullOnly {
fn name(&self) -> &str {
"full-only"
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn crate::constraint::TypeConstraint> {
SimpleTypeConstraint::new(self.name(), vec![I64Sort.to_arcsort()], span.clone())
.into_box()
}
}
impl FullPrim for FullOnly {
fn apply<'a, 'db>(&self, state: FullState<'a, 'db>, _args: &[Value]) -> Option<Value> {
Some(state.base_values().get::<i64>(1))
}
}
#[test]
fn test_user_defined_primitive() {
let mut egraph = EGraph::default();
egraph
.parse_and_run_program(None, "(sort IntVec (Vec i64))")
.unwrap();
let int_vec_sort = egraph.get_arcsort_by(|s| {
s.value_type() == Some(std::any::TypeId::of::<VecContainer>())
&& s.inner_sorts()[0].name() == I64Sort.name()
});
egraph.add_pure_primitive(InnerProduct { vec: int_vec_sort }, None);
egraph
.parse_and_run_program(
None,
"
(let a (vec-of 1 2 3 4 5 6))
(let b (vec-of 6 5 4 3 2 1))
(check (= (inner-product a b) 56))
",
)
.unwrap();
}
#[test]
fn proof_support_accepts_container_sort_declarations() {
let mut egraph = EGraph::default();
let resolved = egraph
.resolve_program(None, "(datatype X (x))\n(sort XPair (Pair X i64))")
.unwrap();
assert!(program_supports_proofs(&resolved, &egraph.type_info));
let mut egraph = EGraph::default();
let resolved = egraph
.resolve_program(None, "(datatype X (x))\n(sort XFn (UnstableFn (X) X))")
.unwrap();
assert!(program_supports_proofs(&resolved, &egraph.type_info));
}
#[test]
fn proof_support_rejects_unstable_fn_primitives_without_validators() {
let mut egraph = EGraph::default();
let resolved = egraph
.resolve_program(
None,
r#"
(datatype X (x))
(sort XFn (UnstableFn (X) X))
(function id (X) X :merge old)
(let f (unstable-fn "id"))
"#,
)
.unwrap();
assert!(!program_supports_proofs(&resolved, &egraph.type_info));
}
#[test]
fn proof_support_accepts_set_primitive_validators() {
let mut egraph = EGraph::default();
let resolved = egraph
.resolve_program(
None,
r#"
(sort ISet (Set i64))
(function Shared () ISet :merge (set-intersect old new))
(check (= (set-insert (set-empty) 1) (set-of 1)))
(check (= (set-remove (set-of 1 2) 2) (set-of 1)))
(check (= (set-length (set-of 1 2)) 2))
(check (set-contains (set-of 1 2) 1))
(check (set-not-contains (set-of 1 2) 3))
(check (= (set-union (set-of 1) (set-of 2)) (set-of 1 2)))
(check (= (set-diff (set-of 1 2) (set-of 2)) (set-of 1)))
(check (= (set-intersect (set-of 1 2) (set-of 2 3)) (set-of 2)))
"#,
)
.unwrap();
assert!(program_supports_proofs(&resolved, &egraph.type_info));
}
#[test]
fn proof_support_rejects_set_get() {
let mut egraph = EGraph::default();
let resolved = egraph
.resolve_program(
None,
r#"
(sort ISet (Set i64))
(check (= (set-get (set-of 1 2) 0) 1))
"#,
)
.unwrap();
assert!(!program_supports_proofs(&resolved, &egraph.type_info));
}
#[test]
fn test_typecheck_expr_with_bindings_and_output_rejects_mismatch() {
let mut egraph = EGraph::default();
let mut parser = crate::ast::Parser::default();
let expr = parser.get_expr_from_string(None, "(+ 1 2)").unwrap();
let resolved = egraph
.typecheck_expr_with_bindings_and_output(
&expr,
&[],
I64Sort.to_arcsort(),
Context::Pure,
)
.unwrap();
assert_eq!(resolved.output_type().name(), I64Sort.name());
let err = egraph
.typecheck_expr_with_bindings_and_output(
&expr,
&[],
BoolSort.to_arcsort(),
Context::Pure,
)
.unwrap_err();
match err {
TypeError::Mismatch {
expected, actual, ..
} => {
assert_eq!(expected.name(), BoolSort.name());
assert_eq!(actual.name(), I64Sort.name());
}
other => panic!("expected mismatch, got {other:?}"),
}
let literal = parser.get_expr_from_string(None, "1").unwrap();
let err = egraph
.typecheck_expr_with_bindings_and_output(
&literal,
&[],
BoolSort.to_arcsort(),
Context::Pure,
)
.unwrap_err();
match err {
TypeError::Mismatch {
expected, actual, ..
} => {
assert_eq!(expected.name(), BoolSort.name());
assert_eq!(actual.name(), I64Sort.name());
}
other => panic!("expected literal mismatch, got {other:?}"),
}
}
#[test]
fn test_typecheck_expr_with_bindings_and_output_uses_explicit_bindings() {
let mut egraph = EGraph::default();
let mut parser = crate::ast::Parser::default();
let expr = parser.get_expr_from_string(None, "(+ x 2)").unwrap();
let bindings = vec![("x".to_string(), span!(), I64Sort.to_arcsort())];
let resolved = egraph
.typecheck_expr_with_bindings_and_output(
&expr,
&bindings,
I64Sort.to_arcsort(),
Context::Pure,
)
.unwrap();
assert_eq!(resolved.output_type().name(), I64Sort.name());
}
#[test]
fn test_typecheck_expr_with_bindings_and_output_uses_context() {
let mut egraph = EGraph::default();
egraph.add_full_primitive(FullOnly, None);
let mut parser = crate::ast::Parser::default();
let expr = parser.get_expr_from_string(None, "(full-only)").unwrap();
let resolved = egraph
.typecheck_expr_with_bindings_and_output(
&expr,
&[],
I64Sort.to_arcsort(),
Context::Full,
)
.unwrap();
assert_eq!(resolved.output_type().name(), I64Sort.name());
let err = egraph
.typecheck_expr_with_bindings_and_output(
&expr,
&[],
I64Sort.to_arcsort(),
Context::Pure,
)
.unwrap_err();
match err {
TypeError::UnboundFunction(name, _) => assert_eq!(name, "full-only"),
other => panic!("expected unbound function, got {other:?}"),
}
}
#[test]
fn test_typecheck_expr_with_bindings_and_output_rejects_duplicate_bindings() {
let mut egraph = EGraph::default();
let mut parser = crate::ast::Parser::default();
let expr = parser.get_expr_from_string(None, "x").unwrap();
let bindings = vec![
("x".to_string(), span!(), I64Sort.to_arcsort()),
("x".to_string(), span!(), BoolSort.to_arcsort()),
];
let err = egraph
.typecheck_expr_with_bindings_and_output(
&expr,
&bindings,
I64Sort.to_arcsort(),
Context::Pure,
)
.unwrap_err();
match err {
TypeError::AlreadyDefined(name, _) => assert_eq!(name, "x"),
other => panic!("expected duplicate binding, got {other:?}"),
}
}
#[test]
fn test_typecheck_expr_with_bindings_and_output_rewrites_globals() {
let mut egraph = EGraph::default();
egraph.parse_and_run_program(None, "(let $x 1)").unwrap();
let mut parser = crate::ast::Parser::default();
let expr = parser.get_expr_from_string(None, "$x").unwrap();
let resolved = egraph
.typecheck_expr_with_bindings_and_output(
&expr,
&[],
I64Sort.to_arcsort(),
Context::Read,
)
.unwrap();
match resolved {
ResolvedExpr::Call(_, ResolvedCall::Func(func), children) => {
assert_eq!(func.name, "$x");
assert!(children.is_empty());
assert_eq!(func.output.name(), I64Sort.name());
}
other => panic!("expected global function call rewrite, got {other:?}"),
}
}
#[test]
fn test_egraph_send_sync() {
fn is_send<T: Send>(_t: &T) -> bool {
true
}
fn is_sync<T: Sync>(_t: &T) -> bool {
true
}
let egraph = EGraph::default();
assert!(is_send(&egraph) && is_sync(&egraph));
}
#[test]
fn test_extension_state_clones_and_restores_with_egraph() {
let mut egraph = EGraph::default();
assert_eq!(egraph.extension_state::<usize>(), None);
assert_eq!(egraph.clone().extension_state::<usize>(), None);
*egraph.extension_state_or_default::<usize>() = 1;
let mut cloned = egraph.clone();
assert_eq!(cloned.extension_state::<usize>(), Some(&1));
*cloned.extension_state_or_default::<usize>() = 2;
assert_eq!(egraph.extension_state::<usize>(), Some(&1));
egraph.push();
*egraph.extension_state_or_default::<usize>() = 3;
egraph.pop().unwrap();
assert_eq!(egraph.extension_state::<usize>(), Some(&1));
}
fn get_function(egraph: &EGraph, name: &str) -> Function {
egraph.functions.get(name).unwrap().clone()
}
fn get_value(egraph: &EGraph, name: &str) -> Value {
let mut out = None;
let id = get_function(egraph, name).backend_id;
egraph.backend.for_each(id, |row| out = Some(row.vals[0]));
out.unwrap()
}
#[test]
fn test_subsumed_unextractable_rebuild_arg() {
let mut egraph = EGraph::default();
egraph
.parse_and_run_program(
None,
r#"
(datatype Math)
(constructor container (Math) Math)
(constructor expensive () Math :cost 100)
(constructor cheap () Math)
(constructor cheap-1 () Math)
; we make the container cheap so that it will be extracted if possible, but then we mark it as subsumed
; so the (expensive) expr should be extracted instead
(let res (container (cheap)))
(union res (expensive))
(cheap)
(cheap-1)
(subsume (container (cheap)))
"#,
).unwrap();
let orig_cheap_value = get_value(&egraph, "cheap");
let orig_cheap_1_value = get_value(&egraph, "cheap-1");
assert_ne!(orig_cheap_value, orig_cheap_1_value);
egraph
.parse_and_run_program(
None,
r#"
(union (cheap-1) (cheap))
"#,
)
.unwrap();
let new_cheap_value = get_value(&egraph, "cheap");
let new_cheap_1_value = get_value(&egraph, "cheap-1");
assert_eq!(new_cheap_value, new_cheap_1_value);
assert!(new_cheap_value != orig_cheap_value || new_cheap_1_value != orig_cheap_1_value);
let outputs = egraph
.parse_and_run_program(
None,
r#"
(extract res)
"#,
)
.unwrap();
assert_eq!(outputs[0].to_string(), "(expensive)\n");
}
#[test]
fn test_subsumed_unextractable_rebuild_self() {
let mut egraph = EGraph::default();
egraph
.parse_and_run_program(
None,
r#"
(datatype Math)
(constructor container (Math) Math)
(constructor expensive () Math :cost 100)
(constructor cheap () Math)
(expensive)
(let x (cheap))
(subsume (cheap))
"#,
)
.unwrap();
let orig_cheap_value = get_value(&egraph, "cheap");
egraph
.parse_and_run_program(
None,
r#"
(union (expensive) x)
"#,
)
.unwrap();
let new_cheap_value = get_value(&egraph, "cheap");
assert_ne!(new_cheap_value, orig_cheap_value);
let res = egraph
.parse_and_run_program(
None,
r#"
(extract x)
"#,
)
.unwrap();
assert_eq!(res[0].to_string(), "(expensive)\n");
}
#[test]
fn test_run_undefined_ruleset_errors() {
let mut egraph = EGraph::default();
let err = egraph
.parse_and_run_program(None, "(ruleset test)\n(run test2 1)")
.unwrap_err();
assert!(matches!(err, Error::NoSuchRuleset(name, _) if name == "test2"));
}
#[test]
fn test_duplicate_rule_name_errors() {
let err = EGraph::default()
.parse_and_run_program(
None,
"(relation foo (i64))
(rule ((foo x)) ((foo (+ x 1))) :name \"r\")
(rule ((foo x)) ((foo (+ x 2))) :name \"r\")",
)
.unwrap_err();
assert!(matches!(err, Error::RuleAlreadyExists(..)));
}
#[test]
fn test_extract_negative_variants_errors() {
let err = EGraph::default()
.parse_and_run_program(
None,
"(sort Math)(constructor Num (i64) Math)(let x (Num 5))(extract x -1)",
)
.unwrap_err();
assert!(matches!(err, Error::ExtractError(..)));
}
#[test]
fn test_input_missing_file_errors() {
let err = EGraph::default()
.parse_and_run_program(
None,
"(function edge (i64) i64 :merge old)(input edge \"/no/such/file_xyz\")",
)
.unwrap_err();
assert!(matches!(err, Error::IoError(..)));
}
}