use egglog::constraint::{SimpleTypeConstraint, TypeConstraint};
use egglog::prelude::*;
use egglog::sort::I64Sort;
use egglog::*;
use egglog_ast::generic_ast::Literal;
#[test]
fn test_add_primitive_validator() {
let mut egraph = EGraph::default();
let validator = |termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
let a = match termdag.get(args[0]) {
Term::Lit(Literal::Int(n)) => *n,
_ => return None,
};
let b = match termdag.get(args[1]) {
Term::Lit(Literal::Int(n)) => *n,
_ => return None,
};
let result = a + b;
let result_term = termdag.lit(Literal::Int(result));
Some(result_term)
};
add_primitive_with_validator!(
&mut egraph,
"test-add" = |a: i64, b: i64| -> i64 { a + b },
validator
);
egraph
.parse_and_run_program(None, "(check (= (test-add 2 3) 5))")
.unwrap();
}
#[test]
fn test_add_pure_primitive_with_validator() {
let mut egraph = EGraph::default();
use egglog::PureState;
#[derive(Clone)]
struct TestAdd;
impl Primitive for TestAdd {
fn name(&self) -> &str {
"test-add"
}
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 TestAdd {
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))
}
}
let validator =
std::sync::Arc::new(|termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
let a = match termdag.get(args[0]) {
Term::Lit(Literal::Int(n)) => *n,
_ => return None,
};
let b = match termdag.get(args[1]) {
Term::Lit(Literal::Int(n)) => *n,
_ => return None,
};
let result = a + b;
let result_term = termdag.lit(Literal::Int(result));
Some(result_term)
});
egraph.add_pure_primitive(TestAdd, Some(validator));
egraph
.parse_and_run_program(None, "(check (= (test-add 2 3) 5))")
.unwrap();
}