diff options
Diffstat (limited to 'src/match_rules.rs')
| -rw-r--r-- | src/match_rules.rs | 432 |
1 files changed, 395 insertions, 37 deletions
diff --git a/src/match_rules.rs b/src/match_rules.rs index 938ee08..fe58119 100644 --- a/src/match_rules.rs +++ b/src/match_rules.rs @@ -1,20 +1,20 @@ use std::collections::HashMap; -use crate::sexpr::{Atom, Expr, Variable, VariableType}; +use crate::sexpr::{Atom, Expr, Number, Variable, VariableType}; -#[derive(Debug, PartialEq, Eq, PartialOrd, Ord)] +#[derive(Debug, PartialEq, Eq, PartialOrd, Ord, Clone)] pub enum BindValue { - Integer(i64), + Number(Number), Expr(Expr), Ellipsis(Vec<Expr>) } -#[derive(Default)] +#[derive(Default, Clone)] pub struct Bindings<'b>(HashMap<&'b str, BindValue>); impl<'b> Bindings<'b> { - pub fn check_or_insert(&mut self, index: &'b String, bind_value: BindValue) -> bool { - if let Some(existing) = self.0.get(index.as_str()) { + pub fn check_or_insert(&mut self, index: &'b str, bind_value: BindValue) -> bool { + if let Some(existing) = self.0.get(index) { existing == &bind_value } else { self.0.insert(index, bind_value); @@ -29,45 +29,152 @@ impl std::fmt::Debug for Bindings<'_> { } } -pub fn args_match<'b>(p_args: &'b [Expr], args: &'b [Expr], bindings: &mut Bindings<'b>) -> bool { - const ELLIPSIS: Expr = Expr::Atom(Atom::Ellipsis); +pub fn substitute(template: &Expr, bindings: &Bindings<'_>) -> Expr { + match template { + Expr::Atom(Atom::Variable(variable)) => { + if let Some(value) = bindings.0.get(variable.index.as_str()) { + match (&variable.r#type, value) { + (VariableType::Number, BindValue::Number(i)) => Expr::Atom(Atom::Number(*i)), + (VariableType::Number, BindValue::Expr(Expr::Atom(Atom::Number(i)))) => { + Expr::Atom(Atom::Number(*i)) + } + (VariableType::Expr | VariableType::NonNumberExpr, BindValue::Expr(expr)) => { + expr.clone() + } + _ => template.clone(), + } + } else { + template.clone() + } + } + Expr::Atom(_) => template.clone(), + Expr::Application(items) => { + let mut rewritten = Vec::new(); + + for item in items { + if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Ellipsis, + index, + })) = item + { + if let Some(BindValue::Ellipsis(exprs)) = bindings.0.get(index.as_str()) { + rewritten.extend(exprs.clone()); + } + } else { + rewritten.push(substitute(item, bindings)); + } + } + + Expr::Application(rewritten) + } + Expr::OrderedList(items) => { + Expr::OrderedList(items.iter().map(|item| substitute(item, bindings)).collect()) + } + } +} + +fn match_with_ellipsis<'b>( + patterns: &[&'b Expr], + remaining: Vec<&'b Expr>, + bindings: &Bindings<'b>, +) -> Option<(Bindings<'b>, Vec<&'b Expr>)> { + if patterns.is_empty() { + return Some((bindings.clone(), remaining)); + } - if p_args.contains(&ELLIPSIS) { - let mut sorted_p: Vec<&Expr> = p_args.iter() - .filter(|p| **p != ELLIPSIS) - .collect(); - - let mut sorted_a: Vec<&Expr> = args.iter().collect(); + let p = patterns[0]; - if sorted_p.len() != sorted_a.len() { - return false; + for (idx, a) in remaining.iter().enumerate() { + let mut trial_bindings = bindings.clone(); + + if matches(p, a, &mut trial_bindings) { + let mut next_remaining = remaining.clone(); + next_remaining.remove(idx); + + if let Some((final_bindings, final_remaining)) = + match_with_ellipsis(&patterns[1..], next_remaining, &trial_bindings) + { + return Some((final_bindings, final_remaining)); + } } + } - sorted_p.sort(); - sorted_a.sort(); + None +} - std::iter::zip(sorted_p, sorted_a).all(|(p, a)| { - matches(p, a, bindings) - }) +fn is_ellipsis_pattern(p: &Expr) -> Option<&str> { + if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Ellipsis, + index, + })) = p + { + Some(index.as_str()) } else { - if p_args.len() != args.len() { - return false; + None + } +} + +fn ellipsis_name(p_args: &[Expr]) -> Result<Option<&str>, ()> { + let mut found: Option<&str> = None; + + for p in p_args { + if let Some(name) = is_ellipsis_pattern(p) { + if found.is_some() { + return Err(()); + } + found = Some(name); + } + } + + Ok(found) +} + +pub fn args_match<'b>(p_args: &'b [Expr], args: &'b [Expr], bindings: &mut Bindings<'b>) -> bool { + match ellipsis_name(p_args) { + Err(()) => false, + Ok(Some(name)) => { + let patterns: Vec<&Expr> = p_args.iter().filter(|p| is_ellipsis_pattern(p).is_none()).collect(); + let remaining: Vec<&Expr> = args.iter().collect(); + + if patterns.len() > remaining.len() { + return false; + } + + if let Some((mut final_bindings, final_remaining)) = + match_with_ellipsis(&patterns, remaining, bindings) + { + final_bindings.0.insert( + name, + BindValue::Ellipsis(final_remaining.into_iter().cloned().collect()), + ); + *bindings = final_bindings; + true + } else { + false + } + } + Ok(None) => { + if p_args.len() != args.len() { + return false; + } + std::iter::zip(p_args.iter(), args.iter()).all(|(a, b)| { + matches(a, b, bindings) + }) } - std::iter::zip(p_args.iter(), args.iter()).all(|(a, b)| { - matches(a, b, bindings) - }) } } pub fn matches<'b>(p: &'b Expr, expr: &'b Expr, bindings: &mut Bindings<'b>) -> bool { match expr { - Expr::Atom(Atom::Int(i)) => { + Expr::Atom(Atom::Number(i)) => { match p { - Expr::Atom(Atom::Int(pi)) => i == pi, - Expr::Atom(Atom::Variable(Variable { r#type: e @ (VariableType::Integer | VariableType::Expr), index })) => { - bindings.check_or_insert(index, match e { - VariableType::Integer => BindValue::Integer(*i), - VariableType::Expr => BindValue::Expr(Expr::Atom(Atom::Int(*i))) + Expr::Atom(Atom::Number(pi)) => i == pi, + Expr::Atom(Atom::Variable(Variable { r#type: e, index })) => { + bindings.check_or_insert(index.as_str(), match e { + VariableType::Number => BindValue::Number(*i), + VariableType::Expr => BindValue::Expr(Expr::Atom(Atom::Number(*i))), + VariableType::NonNumberExpr => return false, + VariableType::Ellipsis => return false, }) } _ => false @@ -76,8 +183,11 @@ pub fn matches<'b>(p: &'b Expr, expr: &'b Expr, bindings: &mut Bindings<'b>) -> Expr::Atom(Atom::Builtin(s)) => { match p { Expr::Atom(Atom::Builtin(ps)) => s == ps, - Expr::Atom(Atom::Variable(Variable { r#type: VariableType::Expr, index })) => { - bindings.check_or_insert(index, BindValue::Expr(Expr::Atom(Atom::Builtin(s.to_string())))) + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr | VariableType::NonNumberExpr, + index, + })) => { + bindings.check_or_insert(index.as_str(), BindValue::Expr(Expr::Atom(Atom::Builtin(s.to_string())))) } _ => false } @@ -89,15 +199,263 @@ pub fn matches<'b>(p: &'b Expr, expr: &'b Expr, bindings: &mut Bindings<'b>) -> f == pf && args_match(&p_exprs[1..], &exprs[1..], bindings) } else if exprs.len() == 1 { matches(p, &exprs[0], bindings) - } else if let Expr::Atom(Atom::Variable(Variable { r#type: VariableType::Expr, index })) = p { - bindings.check_or_insert(index, BindValue::Expr(expr.clone())) + } else if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr | VariableType::NonNumberExpr, + index, + })) = p { + bindings.check_or_insert(index.as_str(), BindValue::Expr(expr.clone())) } else { false } }, - Expr::Atom(Atom::Variable(_)) | Expr::Atom(Atom::Ellipsis) => { + Expr::OrderedList(exprs) => { + if let Expr::OrderedList(p_exprs) = p { + if p_exprs.len() != exprs.len() { + return false; + } + + std::iter::zip(p_exprs.iter(), exprs.iter()).all(|(p_expr, expr)| { + matches(p_expr, expr, bindings) + }) + } else if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr | VariableType::NonNumberExpr, + index, + })) = p { + bindings.check_or_insert(index.as_str(), BindValue::Expr(expr.clone())) + } else { + false + } + } + Expr::Atom(Atom::Variable(_)) => { // These should only appear in patterns, not in expressions being matched false } } } + +#[cfg(test)] +mod tests { + use super::*; + + fn num(i: f64) -> Expr { + Expr::Atom(Atom::Number(Number(i))) + } + + fn builtin(name: &str) -> Expr { + Expr::Atom(Atom::Builtin(name.to_string())) + } + + fn var_expr(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr, + index: name.to_string(), + })) + } + + fn var_int(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Number, + index: name.to_string(), + })) + } + + fn var_ellipsis(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Ellipsis, + index: name.to_string(), + })) + } + + fn var_non_number(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::NonNumberExpr, + index: name.to_string(), + })) + } + + fn app(items: Vec<Expr>) -> Expr { + Expr::Application(items) + } + + #[test] + fn ellipsis_allows_extra_arguments() { + let pattern = app(vec![ + builtin("+"), + num(3.0), + app(vec![builtin("*"), var_int("i"), var_expr("a")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("+"), + num(3.0), + app(vec![builtin("*"), num(7.0), builtin("foo")]), + builtin("bar"), + builtin("baz"), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + } + + #[test] + fn ellipsis_matching_backtracks_for_bindings() { + let pattern = app(vec![ + builtin("f"), + var_expr("x"), + app(vec![builtin("g"), var_expr("x")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("f"), + app(vec![builtin("g"), builtin("a")]), + builtin("a"), + builtin("tail"), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + } + + #[test] + fn substitute_replaces_expression_and_integer_variables() { + let pattern = app(vec![builtin("*"), var_int("i"), var_expr("x")]); + let expr = app(vec![builtin("*"), num(9.0), builtin("foo")]); + let mut bindings = Bindings::default(); + + assert!(matches(&pattern, &expr, &mut bindings)); + + let rhs = app(vec![builtin("+"), var_int("i"), var_expr("x")]); + let rewritten = substitute(&rhs, &bindings); + + assert_eq!(rewritten, app(vec![builtin("+"), num(9.0), builtin("foo")])); + } + + #[test] + fn substitute_splices_ellipsis_tail() { + let pattern = app(vec![ + builtin("+"), + app(vec![builtin("*"), var_int("i"), var_expr("a")]), + app(vec![builtin("*"), var_int("j"), var_expr("a")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("+"), + app(vec![builtin("*"), num(3.0), builtin("a")]), + app(vec![builtin("*"), num(2.0), builtin("a")]), + num(7.0), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + + let rhs = app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), var_int("i"), var_int("j")]), + var_expr("a"), + ]), + var_ellipsis("rest"), + ]); + + let rewritten = substitute(&rhs, &bindings); + + assert_eq!( + rewritten, + app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), num(3.0), num(2.0)]), + builtin("a"), + ]), + num(7.0), + ]) + ); + } + + #[test] + fn named_ellipsis_captures_unmatched_arguments() { + let pattern = app(vec![ + builtin("+"), + app(vec![builtin("*"), var_int("i"), var_expr("a")]), + app(vec![builtin("*"), var_int("j"), var_expr("a")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("+"), + builtin("x"), + app(vec![builtin("*"), num(3.0), builtin("a")]), + builtin("y"), + app(vec![builtin("*"), num(2.0), builtin("a")]), + num(7.0), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + + let rhs = app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), var_int("i"), var_int("j")]), + var_expr("a"), + ]), + var_ellipsis("rest"), + ]); + + let rewritten = substitute(&rhs, &bindings); + assert_eq!( + rewritten, + app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), num(3.0), num(2.0)]), + builtin("a"), + ]), + builtin("x"), + builtin("y"), + num(7.0), + ]) + ); + } + + #[test] + fn duplicate_named_ellipsis_is_rejected() { + let pattern = app(vec![builtin("+"), var_ellipsis("rest"), var_ellipsis("rest")]); + let expr = app(vec![builtin("+"), num(1.0), num(2.0)]); + + let mut bindings = Bindings::default(); + assert!(!matches(&pattern, &expr, &mut bindings)); + } + + #[test] + fn non_number_var_matches_symbol_and_application() { + let symbol_pattern = app(vec![builtin("*"), var_non_number("a"), var_non_number("a")]); + let symbol_expr = app(vec![builtin("*"), builtin("x"), builtin("x")]); + let mut bindings = Bindings::default(); + assert!(matches(&symbol_pattern, &symbol_expr, &mut bindings)); + + let app_expr = app(vec![ + builtin("*"), + app(vec![builtin("+"), builtin("x"), num(1.0)]), + app(vec![builtin("+"), builtin("x"), num(1.0)]), + ]); + let mut bindings = Bindings::default(); + assert!(matches(&symbol_pattern, &app_expr, &mut bindings)); + } + + #[test] + fn non_number_var_does_not_match_number_atom() { + let pattern = app(vec![builtin("*"), var_non_number("a"), var_non_number("a")]); + let expr = app(vec![builtin("*"), num(2.0), num(2.0)]); + + let mut bindings = Bindings::default(); + assert!(!matches(&pattern, &expr, &mut bindings)); + } +} |
