summaryrefslogtreecommitdiff
path: root/src/match_rules.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/match_rules.rs')
-rw-r--r--src/match_rules.rs432
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));
+ }
+}