use std::collections::HashMap; use crate::sexpr::{Atom, Expr, Number, Variable, VariableType}; #[derive(Debug, PartialEq, Eq, PartialOrd, Ord, Clone)] pub enum BindValue { Number(Number), Expr(Expr), Ellipsis(Vec) } #[derive(Default, Clone)] pub struct Bindings<'b>(HashMap<&'b str, BindValue>); impl<'b> Bindings<'b> { 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); true } } } impl std::fmt::Debug for Bindings<'_> { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{:?}", self.0) } } 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)); } let p = patterns[0]; 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)); } } } None } 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 { None } } fn ellipsis_name(p_args: &[Expr]) -> Result, ()> { 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) }) } } } pub fn matches<'b>(p: &'b Expr, expr: &'b Expr, bindings: &mut Bindings<'b>) -> bool { match expr { Expr::Atom(Atom::Number(i)) => { match p { 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 } } Expr::Atom(Atom::Builtin(s)) => { match p { Expr::Atom(Atom::Builtin(ps)) => s == ps, 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 } } Expr::Application(exprs) => { if let Expr::Atom(Atom::Builtin(f)) = &exprs[0] && let Expr::Application(p_exprs) = p && let Expr::Atom(Atom::Builtin(pf)) = &p_exprs[0] { 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 | VariableType::NonNumberExpr, index, })) = p { bindings.check_or_insert(index.as_str(), BindValue::Expr(expr.clone())) } else { false } }, 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::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)); } }