use egglog::add_primitive;
use egglog::ast::Span;
use egglog::constraint::{SimpleTypeConstraint, TypeConstraint};
use egglog::scheduler::{Matches, Scheduler};
use egglog::sort::{I64Sort, S, StringSort};
use egglog::{
EGraph, Error, FullPrim, FullState, Primitive, PurePrim, PureState, RawValues, Read, ReadPrim,
ReadState, TypeError, Value, WritePrim, WriteState, prelude::*,
};
#[track_caller]
fn assert_ambiguous_primitive<T: std::fmt::Debug>(result: Result<T, Error>) {
assert!(
matches!(
result,
Err(Error::TypeError(TypeError::AmbiguousPrimitive { .. }))
),
"expected TypeError::AmbiguousPrimitive, got: {result:?}"
);
}
#[derive(Clone)]
struct ChooseAllScheduler;
impl Scheduler for ChooseAllScheduler {
fn filter_matches(&mut self, _rule: &str, _ruleset: &str, matches: &mut Matches) -> bool {
matches.choose_all();
false
}
}
#[derive(Clone)]
struct PureAdd(&'static str);
impl Primitive for PureAdd {
fn name(&self) -> &str {
self.0
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![
I64Sort.to_arcsort(),
I64Sort.to_arcsort(),
I64Sort.to_arcsort(),
],
span.clone(),
)
.into_box()
}
}
impl PurePrim for PureAdd {
fn apply<'a, 'db>(&self, state: PureState<'a, 'db>, args: &[Value]) -> Option<Value> {
let a = state.base_values().unwrap::<i64>(args[0]);
let b = state.base_values().unwrap::<i64>(args[1]);
Some(state.base_values().get(a + b))
}
}
#[derive(Clone)]
struct WriteEcho(&'static str);
impl Primitive for WriteEcho {
fn name(&self) -> &str {
self.0
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![I64Sort.to_arcsort(), I64Sort.to_arcsort()],
span.clone(),
)
.into_box()
}
}
impl WritePrim for WriteEcho {
fn apply<'a, 'db>(&self, state: WriteState<'a, 'db>, args: &[Value]) -> Option<Value> {
let _ = state.base_values();
Some(args[0])
}
}
#[derive(Clone)]
struct ReadLookup {
name: &'static str,
table_name: &'static str,
}
impl Primitive for ReadLookup {
fn name(&self) -> &str {
self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![I64Sort.to_arcsort(), I64Sort.to_arcsort()],
span.clone(),
)
.into_box()
}
}
impl ReadPrim for ReadLookup {
fn apply<'a, 'db>(&self, state: ReadState<'a, 'db>, args: &[Value]) -> Option<Value> {
state
.lookup(self.table_name, RawValues(args.to_vec()))
.ok()
.flatten()
}
}
#[derive(Clone)]
struct ReadTableSize(&'static str);
impl Primitive for ReadTableSize {
fn name(&self) -> &str {
self.0
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![StringSort.to_arcsort(), I64Sort.to_arcsort()],
span.clone(),
)
.into_box()
}
}
impl ReadPrim for ReadTableSize {
fn apply<'a, 'db>(&self, state: ReadState<'a, 'db>, args: &[Value]) -> Option<Value> {
let table_name = state.base_values().unwrap::<S>(args[0]).0;
let size = state.table_size(&table_name).unwrap_or(0);
let size = i64::try_from(size).ok()?;
Some(state.base_values().get::<i64>(size))
}
}
#[derive(Clone)]
struct ReadAllTableSizes(&'static str);
impl Primitive for ReadAllTableSizes {
fn name(&self) -> &str {
self.0
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(self.name(), vec![I64Sort.to_arcsort()], span.clone()).into_box()
}
}
impl ReadPrim for ReadAllTableSizes {
fn apply<'a, 'db>(&self, state: ReadState<'a, 'db>, _args: &[Value]) -> Option<Value> {
let size: usize = state.table_sizes().into_iter().map(|(_, size)| size).sum();
let size = i64::try_from(size).ok()?;
Some(state.base_values().get::<i64>(size))
}
}
#[derive(Clone)]
struct FullEcho(&'static str);
impl Primitive for FullEcho {
fn name(&self) -> &str {
self.0
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![I64Sort.to_arcsort(), I64Sort.to_arcsort()],
span.clone(),
)
.into_box()
}
}
impl FullPrim for FullEcho {
fn apply<'a, 'db>(&self, state: FullState<'a, 'db>, args: &[Value]) -> Option<Value> {
let _ = state.base_values();
Some(args[0])
}
}
#[derive(Clone)]
struct PureEcho(&'static str);
impl Primitive for PureEcho {
fn name(&self) -> &str {
self.0
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![I64Sort.to_arcsort(), I64Sort.to_arcsort()],
span.clone(),
)
.into_box()
}
}
impl PurePrim for PureEcho {
fn apply<'a, 'db>(&self, _state: PureState<'a, 'db>, args: &[Value]) -> Option<Value> {
Some(args[0])
}
}
#[derive(Clone)]
struct ReadEcho(&'static str);
impl Primitive for ReadEcho {
fn name(&self) -> &str {
self.0
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![I64Sort.to_arcsort(), I64Sort.to_arcsort()],
span.clone(),
)
.into_box()
}
}
impl ReadPrim for ReadEcho {
fn apply<'a, 'db>(&self, _state: ReadState<'a, 'db>, args: &[Value]) -> Option<Value> {
Some(args[0])
}
}
#[test]
fn pure_primitive_accepted_everywhere() {
let mut egraph = EGraph::default();
egraph.add_pure_primitive(PureAdd("p-add"), None);
egraph
.parse_and_run_program(None, "(check (= (p-add 2 3) 5))")
.unwrap();
egraph
.parse_and_run_program(None, "(let $x (p-add 7 8))")
.unwrap();
egraph
.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(rule ((= y (p-add 1 2))) ((set (f y) (p-add 10 20))))\n\
(run 1)",
)
.unwrap();
}
#[test]
fn write_primitive_rejected_in_queries() {
let mut egraph = EGraph::default();
egraph.add_write_primitive(WriteEcho("w-echo"), None);
egraph
.parse_and_run_program(
None,
"(function g (i64) i64 :no-merge)\n\
(rule () ((set (g 0) (w-echo 42))))",
)
.unwrap();
let mut egraph2 = EGraph::default();
egraph2.add_write_primitive(WriteEcho("w-echo"), None);
let result = egraph2.parse_and_run_program(
None,
"(function g (i64) i64 :no-merge)\n\
(rule ((= x (w-echo 1))) ((set (g 0) x)))",
);
assert!(result.is_err(), "WritePrim must be rejected on a rule LHS");
let mut egraph3 = EGraph::default();
egraph3.add_write_primitive(WriteEcho("w-echo"), None);
let result = egraph3.parse_and_run_program(None, "(check (= (w-echo 1) 1))");
assert!(
result.is_err(),
"WritePrim must be rejected in `check` (Context::Read)"
);
}
#[test]
fn read_primitive_rejected_in_rule_contexts() {
let mut egraph = EGraph::default();
egraph.add_read_primitive(
ReadLookup {
name: "lookup-f",
table_name: "f",
},
None,
);
egraph
.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(set (f 7) 42)\n\
(check (= (lookup-f 7) 42))\n\
(let $r (lookup-f 7))\n\
(check (= $r 42))",
)
.unwrap();
let mut egraph2 = EGraph::default();
egraph2.add_read_primitive(
ReadLookup {
name: "lookup-f",
table_name: "f",
},
None,
);
let result = egraph2.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(function g (i64) i64 :no-merge)\n\
(rule ((= x (lookup-f 1))) ((set (g 0) x)))",
);
assert!(result.is_err(), "ReadPrim must be rejected on a rule LHS");
let mut egraph3 = EGraph::default();
egraph3.add_read_primitive(
ReadLookup {
name: "lookup-f",
table_name: "f",
},
None,
);
let result = egraph3.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(function g (i64) i64 :no-merge)\n\
(rule () ((set (g 0) (lookup-f 1))))",
);
assert!(result.is_err(), "ReadPrim must be rejected on a rule RHS");
let mut egraph4 = EGraph::default();
egraph4.add_read_primitive(
ReadLookup {
name: "lookup-f",
table_name: "f",
},
None,
);
egraph4
.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(function g (i64) i64 :no-merge)\n\
(set (f 1) 99)\n\
(rule ((= x (lookup-f 1))) ((set (g 0) x)) :naive)\n\
(run 1)\n\
(check (= (g 0) 99))",
)
.unwrap();
}
#[test]
fn read_primitive_can_observe_table_sizes() {
let mut egraph = EGraph::default();
egraph.add_read_primitive(ReadTableSize("table-size"), None);
egraph.add_read_primitive(ReadAllTableSizes("all-table-sizes"), None);
egraph
.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(set (f 1) 10)\n\
(set (f 2) 20)\n\
(check (= (table-size \"f\") 2))\n\
(check (= (all-table-sizes) 2))",
)
.unwrap();
}
#[test]
fn full_primitive_accepted_only_in_global_action() {
let mut egraph = EGraph::default();
egraph.add_full_primitive(FullEcho("f-echo"), None);
egraph
.parse_and_run_program(None, "(let $ff (f-echo 7))")
.unwrap();
let mut egraph2 = EGraph::default();
egraph2.add_full_primitive(FullEcho("f-echo"), None);
let result = egraph2.parse_and_run_program(None, "(check (= (f-echo 1) 1))");
assert!(
result.is_err(),
"FullPrim must be rejected in Context::Read (`check`)"
);
let mut egraph3 = EGraph::default();
egraph3.add_full_primitive(FullEcho("f-echo"), None);
let result = egraph3.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(rule ((= x (f-echo 1))) ((set (f 0) x)))",
);
assert!(result.is_err(), "FullPrim must be rejected on a rule LHS");
let mut egraph4 = EGraph::default();
egraph4.add_full_primitive(FullEcho("f-echo"), None);
let result = egraph4.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(rule () ((set (f 0) (f-echo 1))))",
);
assert!(result.is_err(), "FullPrim must be rejected on a rule RHS");
let mut egraph5 = EGraph::default();
egraph5.add_full_primitive(FullEcho("f-echo"), None);
egraph5
.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(function trigger () i64 :no-merge)\n\
(set (trigger) 1)\n\
(rule ((= _ (trigger))) ((set (f 0) (f-echo 5))) :naive)\n\
(run 1)\n\
(check (= (f 0) 5))",
)
.unwrap();
}
#[test]
fn merge_primitives_use_write_context() {
let mut egraph = EGraph::default();
egraph.add_write_primitive(WriteEcho("w-echo"), None);
egraph
.parse_and_run_program(None, "(function g () i64 :merge (w-echo old))")
.unwrap();
let mut egraph2 = EGraph::default();
egraph2.add_read_primitive(
ReadLookup {
name: "lookup-f",
table_name: "f",
},
None,
);
let result = egraph2.parse_and_run_program(
None,
"(function f (i64) i64 :no-merge)\n\
(function g () i64 :merge (lookup-f old))",
);
assert!(result.is_err(), "ReadPrim must be rejected in :merge");
let mut egraph3 = EGraph::default();
egraph3.add_full_primitive(FullEcho("f-echo"), None);
let result = egraph3.parse_and_run_program(None, "(function g () i64 :merge (f-echo old))");
assert!(result.is_err(), "FullPrim must be rejected in :merge");
}
#[test]
fn two_same_signature_registrations_error_on_use() {
let mut egraph = EGraph::default();
egraph.add_pure_primitive(PureAdd("dup-add"), None);
egraph.add_pure_primitive(PureAdd("dup-add"), None);
assert_ambiguous_primitive(egraph.parse_and_run_program(None, "(check (= (dup-add 1 2) 3))"));
}
#[test]
fn duplicate_primitive_in_rule_query_errors() {
let mut egraph = EGraph::default();
egraph.add_pure_primitive(PureAdd("dup-add"), None);
egraph.add_pure_primitive(PureAdd("dup-add"), None);
assert_ambiguous_primitive(egraph.parse_and_run_program(
None,
"(relation R (i64))\n\
(rule ((R x) (= y (dup-add x x))) ((R y)))",
));
}
#[test]
fn duplicate_primitive_in_rule_action_errors() {
let mut egraph = EGraph::default();
egraph.add_pure_primitive(PureAdd("dup-add"), None);
egraph.add_pure_primitive(PureAdd("dup-add"), None);
assert_ambiguous_primitive(egraph.parse_and_run_program(
None,
"(relation R (i64))\n\
(function out (i64) i64 :no-merge)\n\
(rule ((R x)) ((set (out x) (dup-add x x))))",
));
}
#[test]
fn duplicate_primitive_in_scheduled_rule_errors_and_restores() {
let mut egraph = EGraph::default();
egraph.add_pure_primitive(PureEcho("dup-echo"), None);
egraph
.parse_and_run_program(
None,
"(sort Fn (UnstableFn (i64) i64))\n\
(ruleset test)\n\
(relation R (i64))\n\
(relation S (i64))\n\
(rule ((R x)) ((let f (unstable-fn \"dup-echo\")) (S (unstable-app f x))) \
:ruleset test :name \"uses-dup\")\n\
(R 0)",
)
.unwrap();
egraph.add_pure_primitive(PureEcho("dup-echo"), None);
let scheduler_id = egraph.add_scheduler(Box::new(ChooseAllScheduler));
assert_ambiguous_primitive(egraph.step_rules_with_scheduler(scheduler_id, "test"));
assert_ambiguous_primitive(egraph.step_rules_with_scheduler(scheduler_id, "test"));
}
#[test]
#[should_panic(expected = "Expected exactly one sort for type `u32`")]
fn missing_sort_panics_with_type_name() {
let mut egraph = EGraph::default();
add_primitive!(&mut egraph, "u32-id" = |a: u32| -> i64 { a as i64 });
}
#[test]
fn unstable_fn_duplicate_primitive_registration_errors_on_build() {
let mut egraph = EGraph::default();
egraph.add_pure_primitive(PureEcho("dup-echo"), None);
egraph.add_pure_primitive(PureEcho("dup-echo"), None);
assert_ambiguous_primitive(egraph.parse_and_run_program(
None,
"(sort Fn (UnstableFn (i64) i64))\n\
(let $f (unstable-fn \"dup-echo\"))\n\
(check (= (unstable-app $f 7) 7))",
));
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum AppCtx {
Pure,
Read,
Write,
Full,
}
const ALL_CTXS: [AppCtx; 4] = [AppCtx::Pure, AppCtx::Read, AppCtx::Write, AppCtx::Full];
fn matrix_program(ctx: AppCtx) -> String {
let header = "(sort Fn (UnstableFn (i64) i64))\n\
(let $f (unstable-fn \"p\"))\n";
match ctx {
AppCtx::Pure => format!(
"{header}(function out (i64) i64 :no-merge)\n\
(rule ((= y (unstable-app $f 7))) ((set (out 0) y)))\n\
(run 1)\n\
(check (= (out 0) 7))"
),
AppCtx::Read => format!("{header}(check (= (unstable-app $f 7) 7))"),
AppCtx::Write => format!(
"{header}(function out (i64) i64 :no-merge)\n\
(rule () ((set (out 0) (unstable-app $f 7))))\n\
(run 1)\n\
(check (= (out 0) 7))"
),
AppCtx::Full => format!("{header}(let $r (unstable-app $f 7))\n(check (= $r 7))"),
}
}
fn run_matrix_cell(register: impl FnOnce(&mut EGraph), ctx: AppCtx) -> Result<(), String> {
let mut egraph = EGraph::default();
register(&mut egraph);
egraph
.parse_and_run_program(None, &matrix_program(ctx))
.map(|_| ())
.map_err(|e| e.to_string())
}
#[test]
fn unstable_app_dispatch_matrix() {
let cells: &[(&str, fn(&mut EGraph), &[AppCtx])] = &[
(
"pure",
|e: &mut EGraph| e.add_pure_primitive(PureEcho("p"), None),
&[AppCtx::Pure, AppCtx::Read, AppCtx::Write, AppCtx::Full],
),
(
"read",
|e: &mut EGraph| e.add_read_primitive(ReadEcho("p"), None),
&[AppCtx::Read, AppCtx::Full],
),
(
"write",
|e: &mut EGraph| e.add_write_primitive(WriteEcho("p"), None),
&[AppCtx::Write, AppCtx::Full],
),
(
"full",
|e: &mut EGraph| e.add_full_primitive(FullEcho("p"), None),
&[AppCtx::Full],
),
];
for (label, register, valid) in cells {
for &ctx in &ALL_CTXS {
let result = run_matrix_cell(*register, ctx);
let should_succeed = valid.contains(&ctx);
if should_succeed {
assert!(
result.is_ok(),
"{label} prim applied via unstable-app in {ctx:?} ctx should succeed; \
got error: {:?}",
result.err()
);
} else {
assert!(
result.is_err(),
"{label} prim applied via unstable-app in {ctx:?} ctx should fail \
(dispatch mismatch panic), but parse_and_run_program returned Ok"
);
}
}
}
}