use egglog::CommandOutput;
use egglog::EGraph;
use egglog::prelude::*;
fn find_extract_best(outputs: &[CommandOutput]) -> String {
outputs
.iter()
.find(|o| matches!(o, CommandOutput::ExtractBest(..)))
.expect("No ExtractBest output found")
.to_string()
}
#[test]
fn test_extraction_same_with_proof_mode() {
let _ = env_logger::builder().is_test(true).try_init();
let program = r#"
(datatype Math
(Num i64)
(Add Math Math)
(Mul Math Math))
(rewrite (Add (Num a) (Num b)) (Num (+ a b)))
(rewrite (Mul (Num a) (Num b)) (Num (* a b)))
; commutativity
(rewrite (Add x y) (Add y x))
(rewrite (Mul x y) (Mul y x))
; associativity
(rewrite (Add (Add x y) z) (Add x (Add y z)))
(rewrite (Mul (Mul x y) z) (Mul x (Mul y z)))
; distributivity
(rewrite (Mul x (Add y z)) (Add (Mul x y) (Mul x z)))
(let expr (Mul (Add (Num 1) (Num 2)) (Num 3)))
(run 10)
"#;
let mut egraph_normal = EGraph::default();
egraph_normal.parse_and_run_program(None, program).unwrap();
let normal_output = egraph_normal
.parse_and_run_program(None, "(extract expr)")
.unwrap();
let normal_extracted = find_extract_best(&normal_output);
let mut egraph_proofs = EGraph::new_with_proofs();
egraph_proofs.parse_and_run_program(None, program).unwrap();
let proofs_output = egraph_proofs
.parse_and_run_program(None, "(extract expr)")
.unwrap();
let proofs_extracted = find_extract_best(&proofs_output);
assert_eq!(
normal_extracted, proofs_extracted,
"Extraction differs between normal mode and proof mode:\nNormal: {normal_extracted}\nProofs: {proofs_extracted}"
);
assert!(
normal_extracted.contains("Num") && normal_extracted.contains("9"),
"Expected (Num 9), got: {normal_extracted}"
);
}
#[test]
fn test_extraction_same_with_proof_mode_using_rule_macro() {
let _ = env_logger::builder().is_test(true).try_init();
let setup = r#"
(datatype Expr
(Var String)
(Lit i64)
(Add Expr Expr))
; Simplification rule that gives a unique result
(rewrite (Add (Lit a) (Lit b)) (Lit (+ a b)))
(let x (Add (Lit 1) (Lit 2)))
(run 10)
"#;
let mut egraph_normal = EGraph::default();
egraph_normal.parse_and_run_program(None, setup).unwrap();
add_ruleset(&mut egraph_normal, "my_rules").unwrap();
rule(
&mut egraph_normal,
"my_rules",
facts![(= (Add a b) e)],
actions![(union e (Add b a))],
)
.unwrap();
for _ in 0..5 {
run_ruleset(&mut egraph_normal, "my_rules").unwrap();
}
let normal_output = egraph_normal
.parse_and_run_program(None, "(extract x)")
.unwrap();
let normal_extracted = find_extract_best(&normal_output);
let mut egraph_proofs = EGraph::new_with_proofs();
egraph_proofs.parse_and_run_program(None, setup).unwrap();
add_ruleset(&mut egraph_proofs, "my_rules").unwrap();
rule(
&mut egraph_proofs,
"my_rules",
facts![(= (Add a b) e)],
actions![(union e (Add b a))],
)
.unwrap();
for _ in 0..5 {
run_ruleset(&mut egraph_proofs, "my_rules").unwrap();
}
let proofs_output = egraph_proofs
.parse_and_run_program(None, "(extract x)")
.unwrap();
let proofs_extracted = find_extract_best(&proofs_output);
assert_eq!(
normal_extracted, proofs_extracted,
"Extraction differs between normal mode and proof mode:\nNormal: {normal_extracted}\nProofs: {proofs_extracted}"
);
assert!(
normal_extracted.contains("Lit") && normal_extracted.contains("3"),
"Expected (Lit 3), got: {normal_extracted}"
);
}