use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use ocas_core::FastHashMap;
use crate::matcher::{Bindings, MatchError, MatchValue, match_pattern};
use crate::pattern::{Pattern, PatternAlloc, WildcardLevel};
pub struct Rule<'a> {
pattern: Pattern<'a>,
replacement: Replacement<'a>,
condition: Option<Condition<'a>>,
template: Option<Pattern<'a>>,
}
type Replacement<'a> = Box<dyn Fn(&Bindings<'a>, &AtomArena<'a>) -> Atom<'a> + 'a>;
type Condition<'a> = Box<dyn Fn(&Bindings<'a>) -> bool + 'a>;
impl<'a> Rule<'a> {
pub fn new<F>(pattern: Pattern<'a>, replacement: F) -> Self
where
F: Fn(&Bindings<'a>, &AtomArena<'a>) -> Atom<'a> + 'a,
{
Self {
pattern,
replacement: Box::new(replacement),
condition: None,
template: None,
}
}
pub fn with_condition<F>(mut self, condition: F) -> Self
where
F: Fn(&Bindings<'a>) -> bool + 'a,
{
self.condition = Some(Box::new(condition));
self
}
pub fn apply(&self, ctx: &'a AtomArena<'a>, atom: Atom<'a>) -> Option<Atom<'a>> {
match match_pattern(self.pattern.clone(), atom) {
Ok(bindings) => {
if let Some(cond) = &self.condition
&& !cond(&bindings)
{
return None;
}
if let Some(tmpl) = &self.template {
let next = instantiate(ctx, tmpl, &bindings)?;
Some(ocas_atom::normalize::normalize(ctx, next))
} else {
Some((self.replacement)(&bindings, ctx))
}
}
Err(MatchError::NoMatch)
| Err(MatchError::InconsistentBinding)
| Err(MatchError::BudgetExhausted) => None,
}
}
pub fn from_template(
ctx: &'a AtomArena<'a>,
alloc: &'a impl PatternAlloc<'a>,
pattern: &str,
template: &str,
) -> Rule<'a> {
let pat = pattern_from_str(ctx, alloc, pattern);
let tmpl = pattern_from_str(ctx, alloc, template);
Rule {
pattern: pat,
replacement: Box::new(|_, ctx| ctx.num(0)),
condition: None,
template: Some(tmpl),
}
}
}
pub fn instantiate<'a>(
ctx: &'a AtomArena<'a>,
template: &Pattern<'a>,
bindings: &Bindings<'a>,
) -> Option<Atom<'a>> {
match template {
Pattern::Literal(a) => Some(*a),
Pattern::Wildcard(name, WildcardLevel::Single) => match bindings.get(*name)? {
MatchValue::Single(v) => Some(*v),
MatchValue::Sequence(_) => None,
},
Pattern::Wildcard(name, level) => {
let _ = (name, level);
None
}
Pattern::Add(pats) => {
let args = instantiate_args(ctx, pats, bindings)?;
Some(ctx.add(&args))
}
Pattern::Mul(pats) => {
let args = instantiate_args(ctx, pats, bindings)?;
Some(ctx.mul(&args))
}
Pattern::Fun(name, pats) => {
let args = instantiate_args(ctx, pats, bindings)?;
Some(ctx.fun(name.as_str(), &args))
}
Pattern::Pow(p_box) => {
let (p_base, p_exp) = &**p_box;
let base = instantiate(ctx, p_base, bindings)?;
let exp = instantiate(ctx, p_exp, bindings)?;
Some(ctx.pow(base, exp))
}
}
}
fn instantiate_args<'a>(
ctx: &'a AtomArena<'a>,
pats: &[Pattern<'a>],
bindings: &Bindings<'a>,
) -> Option<Vec<Atom<'a>>> {
let mut out = Vec::with_capacity(pats.len());
for pat in pats {
match pat {
Pattern::Wildcard(name, WildcardLevel::Sequence | WildcardLevel::NullSequence) => {
match bindings.get(*name)? {
MatchValue::Sequence(slice) => out.extend_from_slice(slice),
MatchValue::Single(_) => return None,
}
}
_ => out.push(instantiate(ctx, pat, bindings)?),
}
}
Some(out)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum HeadKey {
Fun(Symbol),
Add,
Mul,
Pow,
Any,
}
pub fn head_of(atom: Atom<'_>) -> HeadKey {
match atom.node() {
AtomNode::Fun(name, _) => HeadKey::Fun(*name),
AtomNode::Add(_) => HeadKey::Add,
AtomNode::Mul(_) => HeadKey::Mul,
AtomNode::Pow(_, _) => HeadKey::Pow,
AtomNode::Num(_) | AtomNode::Var(_) => HeadKey::Any,
}
}
fn pattern_head_key(pattern: &Pattern<'_>) -> HeadKey {
match pattern {
Pattern::Fun(name, _) => HeadKey::Fun(*name),
Pattern::Add(_) => HeadKey::Add,
Pattern::Mul(_) => HeadKey::Mul,
Pattern::Pow(_) => HeadKey::Pow,
Pattern::Literal(_) | Pattern::Wildcard(_, _) => HeadKey::Any,
}
}
pub struct RuleTable<'a> {
by_head: FastHashMap<HeadKey, Vec<Rule<'a>>>,
}
impl<'a> RuleTable<'a> {
pub fn from_rules(rules: Vec<Rule<'a>>) -> Self {
let mut by_head: FastHashMap<HeadKey, Vec<Rule<'a>>> = FastHashMap::default();
for rule in rules {
let key = pattern_head_key(&rule.pattern);
by_head.entry(key).or_default().push(rule);
}
Self { by_head }
}
pub fn apply(&self, ctx: &'a AtomArena<'a>, atom: Atom<'a>) -> Option<Atom<'a>> {
let key = head_of(atom);
if let Some(rules) = self.by_head.get(&key) {
for rule in rules {
if let Some(next) = rule.apply(ctx, atom) {
return Some(next);
}
}
}
if key != HeadKey::Any
&& let Some(rules) = self.by_head.get(&HeadKey::Any)
{
for rule in rules {
if let Some(next) = rule.apply(ctx, atom) {
return Some(next);
}
}
}
None
}
}
fn pattern_from_str<'a>(
ctx: &'a AtomArena<'a>,
alloc: &'a impl PatternAlloc<'a>,
s: &str,
) -> Pattern<'a> {
let atom = ocas_parse::parse(ctx, s).expect("built-in rule pattern is valid");
Pattern::from_atom(alloc, atom)
}
macro_rules! binding_single {
($bindings:expr, $name:expr) => {
match $bindings.get(ocas_atom::Symbol::new($name)) {
Some(crate::matcher::MatchValue::Single(atom)) => *atom,
_ => panic!(concat!("expected single binding for '", $name, "'")),
}
};
}
pub fn add_zero<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "x_ + 0"), |bindings, _ctx| {
binding_single!(bindings, "x")
})
}
pub fn add_zero_left<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "0 + x_"), |bindings, _ctx| {
binding_single!(bindings, "x")
})
}
pub fn mul_zero<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "x_ * 0"), |_bindings, ctx| {
ctx.num(0)
})
}
pub fn mul_zero_left<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "0 * x_"), |_bindings, ctx| {
ctx.num(0)
})
}
pub fn mul_one<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "x_ * 1"), |bindings, _ctx| {
binding_single!(bindings, "x")
})
}
pub fn mul_one_left<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "1 * x_"), |bindings, _ctx| {
binding_single!(bindings, "x")
})
}
pub fn add_same<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "x_ + x_"), |bindings, ctx| {
let x = binding_single!(bindings, "x");
ctx.mul(&[ctx.num(2), x])
})
}
pub fn pow_zero<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "x_ ^ 0"), |_bindings, ctx| {
ctx.num(1)
})
}
pub fn pow_one<'a>(ctx: &'a AtomArena<'a>, alloc: &'a impl PatternAlloc<'a>) -> Rule<'a> {
Rule::new(pattern_from_str(ctx, alloc, "x_ ^ 1"), |bindings, _ctx| {
binding_single!(bindings, "x")
})
}
pub fn default_rules<'a>(
ctx: &'a AtomArena<'a>,
alloc: &'a impl PatternAlloc<'a>,
) -> Vec<Rule<'a>> {
vec![
add_zero(ctx, alloc),
add_zero_left(ctx, alloc),
mul_zero(ctx, alloc),
mul_zero_left(ctx, alloc),
mul_one(ctx, alloc),
mul_one_left(ctx, alloc),
add_same(ctx, alloc),
pow_zero(ctx, alloc),
pow_one(ctx, alloc),
]
}
#[cfg(test)]
mod tests {
use super::*;
use ocas_atom::AtomArena;
use ocas_core::arena::Arena;
struct VecAlloc;
impl<'a> PatternAlloc<'a> for VecAlloc {
fn alloc_slice(&self, _items: &[Pattern<'a>]) -> &'a [Pattern<'a>] {
unreachable!()
}
}
fn pat_atom<'a>(ctx: &'a AtomArena<'a>, alloc: &'a VecAlloc, s: &'a str) -> Pattern<'a> {
let atom = ocas_parse::parse(ctx, s).unwrap();
Pattern::from_atom(alloc, atom)
}
#[test]
fn add_zero_applies() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let rule = add_zero(&ctx, &alloc);
let x = ctx.var("x");
let atom = ctx.add(&[x, ctx.num(0)]);
let result = rule.apply(&ctx, atom).unwrap();
assert_eq!(result, x);
}
#[test]
fn mul_zero_applies() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let rule = mul_zero(&ctx, &alloc);
let x = ctx.var("x");
let atom = ctx.mul(&[x, ctx.num(0)]);
let result = rule.apply(&ctx, atom).unwrap();
assert_eq!(result, ctx.num(0));
}
#[test]
fn add_same_applies() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let rule = add_same(&ctx, &alloc);
let x = ctx.var("x");
let atom = ctx.add(&[x, x]);
let result = rule.apply(&ctx, atom).unwrap();
assert_eq!(result.to_string(), "2*x");
}
#[test]
fn rule_with_condition_can_block() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let pat = pat_atom(&ctx, &alloc, "x_");
let rule = Rule::new(pat, |bindings, _ctx| binding_single!(bindings, "x"))
.with_condition(|bindings| {
matches!(
bindings.get(ocas_atom::Symbol::new("x")),
Some(crate::matcher::MatchValue::Single(a)) if matches!(a.node(), ocas_atom::AtomNode::Num(_))
)
});
let x = ctx.var("x");
assert!(rule.apply(&ctx, x).is_none());
assert!(rule.apply(&ctx, ctx.num(5)).is_some());
}
#[test]
fn from_template_instantiates_single_wildcards() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let rule = Rule::from_template(&ctx, &alloc, "x_^n_", "x_^(n_+1)*(n_+1)^-1");
let atom = ocas_parse::parse(&ctx, "x^3").unwrap();
let result = rule.apply(&ctx, atom).unwrap();
assert_eq!(result.to_string(), "(4^-1)*(x^4)");
assert!(rule.apply(&ctx, ctx.var("y")).is_none());
}
#[test]
fn from_template_splices_sequence_wildcards() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let rule = Rule::from_template(&ctx, &alloc, "f(a__, b_)", "g(a__, b_)");
let atom = ocas_parse::parse(&ctx, "f(x, y, z)").unwrap();
let result = rule.apply(&ctx, atom).unwrap();
assert_eq!(result.to_string(), "g(x, y, z)");
}
#[test]
fn instantiate_unbound_wildcard_is_none() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let tmpl = pat_atom(&ctx, &alloc, "x_^2");
let bindings = Bindings::new();
assert!(instantiate(&ctx, &tmpl, &bindings).is_none());
}
#[test]
fn rule_table_equals_linear_scan_on_default_rules() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let alloc = VecAlloc;
let rules = default_rules(&ctx, &alloc);
let table = RuleTable::from_rules(rules);
let samples = [
"x + 0", "0 + x", "x*0", "0*x", "x*1", "1*x", "x + x", "y^0", "y^1", "2*x + 0", "x*y",
"0",
];
for s in samples {
let atom = ocas_parse::parse(&ctx, s).unwrap();
let mut linear = atom;
for rule in default_rules(&ctx, &alloc) {
if let Some(next) = rule.apply(&ctx, linear) {
linear = next;
}
}
let indexed = table.apply(&ctx, atom).unwrap_or(atom);
assert_eq!(
indexed, linear,
"table vs linear disagree on {s}: {indexed} vs {linear}"
);
}
}
}