use std::{
hash::Hash,
ops::{Add, BitOr, DerefMut, Div, Mul},
};
use ahash::{HashMap, HashSet};
use dyn_clone::DynClone;
use crate::{
OperationCount,
atom::{
Atom, AtomCore, AtomType, AtomView, Indeterminate, InlineNum, ListIterator, SliceType,
Symbol, representation::InlineVar,
},
coefficient::{Coefficient, CoefficientView},
domains::rational::Rational,
state::{RecycledAtom, Workspace},
transformer::{Transformer, TransformerError},
utils::{BorrowedOrOwned, Settable},
};
pub use crate::atom::AliasedAtom;
static ZERO: InlineNum = InlineNum::zero();
static ONE: InlineNum = InlineNum::one();
#[derive(Clone)]
pub enum Pattern {
Literal(Atom),
Wildcard(Symbol, bool),
Fn(Symbol, Vec<Pattern>),
Pow(Box<[Pattern; 2]>),
Mul(Vec<Pattern>),
Add(Vec<Pattern>),
Alternative(Vec<Pattern>),
Transformer(Box<(Option<Pattern>, Vec<Transformer>)>),
}
impl<T: Into<Pattern>> Mul<T> for &Pattern {
type Output = Pattern;
fn mul(self, rhs: T) -> Self::Output {
Workspace::get_local().with(|ws| self.mul(&rhs.into(), ws))
}
}
impl<T: Into<Pattern>> Mul<T> for Pattern {
type Output = Pattern;
fn mul(self, rhs: T) -> Self::Output {
Workspace::get_local().with(|ws| (&self).mul(&rhs.into(), ws))
}
}
impl<T: Into<Pattern>> Add<T> for &Pattern {
type Output = Pattern;
fn add(self, rhs: T) -> Self::Output {
Workspace::get_local().with(|ws| self.add(&rhs.into(), ws))
}
}
impl<T: Into<Pattern>> Add<T> for Pattern {
type Output = Pattern;
fn add(self, rhs: T) -> Self::Output {
Workspace::get_local().with(|ws| (&self).add(&rhs.into(), ws))
}
}
impl<T: Into<Pattern>> Div<T> for &Pattern {
type Output = Pattern;
fn div(self, rhs: T) -> Self::Output {
Workspace::get_local().with(|ws| self.div(&rhs.into(), ws))
}
}
impl<T: Into<Pattern>> Div<T> for Pattern {
type Output = Pattern;
fn div(self, rhs: T) -> Self::Output {
Workspace::get_local().with(|ws| (&self).div(&rhs.into(), ws))
}
}
impl<T: Into<Pattern>> BitOr<T> for &Pattern {
type Output = Pattern;
fn bitor(self, rhs: T) -> Self::Output {
self.clone().alternative(rhs.into()).unwrap()
}
}
impl<T: Into<Pattern>> BitOr<T> for Pattern {
type Output = Pattern;
fn bitor(self, rhs: T) -> Self::Output {
self.alternative(rhs.into()).unwrap()
}
}
impl From<Symbol> for Pattern {
fn from(symbol: Symbol) -> Pattern {
InlineVar::new(symbol).to_pattern()
}
}
impl From<Atom> for Pattern {
fn from(atom: Atom) -> Self {
Pattern::new(atom)
}
}
impl From<Indeterminate> for Pattern {
fn from(atom: Indeterminate) -> Self {
match atom {
Indeterminate::Symbol(s, _) => Pattern::from(s),
Indeterminate::Function(_, a) => Pattern::from(a),
}
}
}
impl std::fmt::Display for Pattern {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
if let Ok(a) = self.to_atom() {
a.fmt(f)
} else {
std::fmt::Debug::fmt(self, f)
}
}
}
pub trait MatchMap: Fn(&MatchStack) -> Atom + DynClone + Send + Sync {}
dyn_clone::clone_trait_object!(MatchMap);
impl<T: Clone + Send + Sync + Fn(&MatchStack) -> Atom> MatchMap for T {}
#[derive(Clone)]
pub enum ReplaceWith<'a> {
Pattern(BorrowedOrOwned<'a, Pattern>),
Map(Box<dyn MatchMap>),
}
impl<T: Into<Coefficient>> From<T> for ReplaceWith<'_> {
fn from(val: T) -> Self {
ReplaceWith::Pattern(BorrowedOrOwned::Owned(Atom::num(val.into()).into()))
}
}
impl From<Atom> for ReplaceWith<'_> {
fn from(val: Atom) -> Self {
ReplaceWith::Pattern(BorrowedOrOwned::Owned(val.into()))
}
}
impl From<Indeterminate> for ReplaceWith<'_> {
fn from(val: Indeterminate) -> Self {
ReplaceWith::Pattern(BorrowedOrOwned::Owned(val.into()))
}
}
impl<'a> From<&'a Pattern> for ReplaceWith<'a> {
fn from(val: &'a Pattern) -> Self {
ReplaceWith::Pattern(BorrowedOrOwned::Borrowed(val))
}
}
impl From<Pattern> for ReplaceWith<'_> {
fn from(val: Pattern) -> Self {
ReplaceWith::Pattern(BorrowedOrOwned::Owned(val))
}
}
impl std::fmt::Debug for ReplaceWith<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ReplaceWith::Pattern(p) => write!(f, "{p:?}"),
ReplaceWith::Map(_) => write!(f, "Map"),
}
}
}
impl std::fmt::Display for ReplaceWith<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ReplaceWith::Pattern(p) => write!(f, "{}", p.borrow()),
ReplaceWith::Map(_) => write!(f, "Map"),
}
}
}
#[derive(Debug, Clone)]
pub struct Replacement {
pub pat: Pattern,
pub rhs: ReplaceWith<'static>,
pub conditions: Option<Condition<PatternRestriction>>,
pub match_settings: MatchSettings,
}
impl std::fmt::Display for Replacement {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{} -> {}", self.pat, self.rhs)?;
if let Some(c) = &self.conditions {
write!(f, "; {c}")?;
}
Ok(())
}
}
impl Replacement {
pub fn new<P: Into<Pattern>, R: Into<ReplaceWith<'static>>>(pat: P, rhs: R) -> Self {
Replacement {
pat: pat.into(),
rhs: rhs.into(),
conditions: None,
match_settings: MatchSettings::default(),
}
}
pub fn non_greedy_wildcards(mut self, non_greedy_wildcards: Vec<Symbol>) -> Self {
self.match_settings.non_greedy_wildcards = non_greedy_wildcards;
self
}
pub fn level_range(mut self, level_range: (usize, Option<usize>)) -> Self {
self.match_settings.level_range = level_range;
self
}
pub fn min_level(mut self, min_level: usize) -> Self {
self.match_settings.level_range.0 = min_level;
self
}
pub fn max_level(mut self, max_level: usize) -> Self {
self.match_settings.level_range.1 = Some(max_level);
self
}
pub fn level_is_tree_depth(mut self, level_is_tree_depth: bool) -> Self {
self.match_settings.level_is_tree_depth = level_is_tree_depth;
self
}
pub fn partial(mut self, partial: bool) -> Self {
self.match_settings.partial = partial;
self
}
pub fn allow_new_wildcards_on_rhs(mut self, allow: bool) -> Self {
self.match_settings.allow_new_wildcards_on_rhs = allow;
self
}
pub fn rhs_cache_size(mut self, rhs_cache_size: usize) -> Self {
self.match_settings.rhs_cache_size = rhs_cache_size;
self
}
pub fn when<'b, R: Into<BorrowedOrOwned<'b, Condition<PatternRestriction>>>>(
mut self,
conditions: R,
) -> Self {
self.conditions = Some(conditions.into().yield_owned());
self
}
}
#[derive(Clone, Copy)]
pub struct BorrowedReplacement<'a> {
pub pattern: &'a Pattern,
pub rhs: &'a ReplaceWith<'a>,
pub conditions: Option<&'a Condition<PatternRestriction>>,
pub settings: Option<&'a MatchSettings>,
}
pub trait BorrowReplacement {
fn borrow(&self) -> BorrowedReplacement<'_>;
}
impl BorrowReplacement for Replacement {
fn borrow(&self) -> BorrowedReplacement<'_> {
BorrowedReplacement {
pattern: &self.pat,
rhs: &self.rhs,
conditions: self.conditions.as_ref(),
settings: Some(&self.match_settings),
}
}
}
impl BorrowReplacement for &Replacement {
fn borrow(&self) -> BorrowedReplacement<'_> {
BorrowedReplacement {
pattern: &self.pat,
rhs: &self.rhs,
conditions: self.conditions.as_ref(),
settings: Some(&self.match_settings),
}
}
}
impl BorrowReplacement for BorrowedReplacement<'_> {
fn borrow(&self) -> BorrowedReplacement<'_> {
*self
}
}
#[derive(Debug, Default, Copy, Clone)]
pub struct ReplaceSettings {
pub(crate) once: bool,
pub(crate) bottom_up: bool,
pub(crate) nested: bool,
}
impl ReplaceSettings {
pub fn new() -> Self {
Self::default()
}
pub fn once(mut self, once: bool) -> Self {
self.once = once;
self
}
pub fn bottom_up(mut self, bottom_up: bool) -> Self {
self.bottom_up = bottom_up;
self
}
pub fn nested(mut self, nested: bool) -> Self {
self.nested = nested;
self
}
}
#[derive(Debug, Clone)]
pub struct ReplaceBuilder<'a, 'b> {
target: AtomView<'a>,
pattern: BorrowedOrOwned<'b, Pattern>,
conditions: Option<BorrowedOrOwned<'b, Condition<PatternRestriction>>>,
match_settings: MatchSettings,
replace_settings: ReplaceSettings,
repeat: bool,
}
impl<'a, 'b> ReplaceBuilder<'a, 'b> {
pub fn new<T: Into<BorrowedOrOwned<'b, Pattern>>>(
target: AtomView<'a>,
replacement: T,
) -> Self {
ReplaceBuilder {
target,
pattern: replacement.into(),
conditions: None,
match_settings: MatchSettings::default(),
repeat: false,
replace_settings: ReplaceSettings::default(),
}
}
pub fn non_greedy_wildcards(mut self, non_greedy_wildcards: Vec<Symbol>) -> Self {
self.match_settings.non_greedy_wildcards = non_greedy_wildcards;
self
}
pub fn level_range(mut self, level_range: (usize, Option<usize>)) -> Self {
self.match_settings.level_range = level_range;
self
}
pub fn min_level(mut self, min_level: usize) -> Self {
self.match_settings.level_range.0 = min_level;
self
}
pub fn max_level(mut self, max_level: usize) -> Self {
self.match_settings.level_range.1 = Some(max_level);
self
}
pub fn level_is_tree_depth(mut self, level_is_tree_depth: bool) -> Self {
self.match_settings.level_is_tree_depth = level_is_tree_depth;
self
}
pub fn partial(mut self, partial: bool) -> Self {
self.match_settings.partial = partial;
self
}
pub fn allow_new_wildcards_on_rhs(mut self, allow: bool) -> Self {
self.match_settings.allow_new_wildcards_on_rhs = allow;
self
}
pub fn rhs_cache_size(mut self, rhs_cache_size: usize) -> Self {
self.match_settings.rhs_cache_size = rhs_cache_size;
self
}
pub fn set_optional(mut self, wildcard: Symbol) -> Self {
self.pattern = BorrowedOrOwned::Owned(self.pattern.yield_owned().set_optional(wildcard));
self
}
pub fn when<R: Into<BorrowedOrOwned<'b, Condition<PatternRestriction>>>>(
mut self,
conditions: R,
) -> Self {
self.conditions = Some(conditions.into());
self
}
pub fn repeat(mut self) -> Self {
self.repeat = true;
self
}
pub fn once(mut self) -> Self {
self.replace_settings.once = true;
self
}
pub fn bottom_up(mut self) -> Self {
self.replace_settings.bottom_up = true;
self
}
pub fn nested(mut self) -> Self {
self.replace_settings.nested = true;
self.bottom_up()
}
pub fn with<'c, R: Into<BorrowedOrOwned<'c, Pattern>>>(&mut self, rhs: R) -> Atom {
let rhs = ReplaceWith::Pattern(rhs.into());
let mut expr_ref = self.target;
let mut out = RecycledAtom::new();
let mut out2 = RecycledAtom::new();
let mut c = Condition::True;
for fn_name in self.pattern.borrow().get_wildcard_function_names() {
c = Condition::And(Box::new((
c,
fn_name.restrict(WildcardRestriction::IsAtomType(AtomType::Var)),
)));
}
if matches!(c, Condition::True) {
if let Some(s) = self.conditions.as_mut() {
let old = std::mem::replace(s, BorrowedOrOwned::Owned(Condition::True));
*s = BorrowedOrOwned::Owned(Condition::And(Box::new((old.yield_owned(), c))));
} else {
self.conditions = Some(BorrowedOrOwned::Owned(c));
}
}
while expr_ref.replace_into(
self.pattern.borrow(),
&rhs,
self.conditions.as_ref().map(|x| x.borrow()),
Some(&self.match_settings),
self.replace_settings,
&mut out,
) {
if !self.repeat || expr_ref == out.as_view() {
break;
}
std::mem::swap(&mut out, &mut out2);
expr_ref = out2.as_view();
}
out.into_inner()
}
pub fn try_with<'c, R: Into<BorrowedOrOwned<'c, Pattern>>>(
&mut self,
rhs: R,
) -> Result<Atom, Symbol> {
let rhs = rhs.into();
if !self.match_settings.allow_new_wildcards_on_rhs {
if let Some(w) = self.pattern.find_new_wildcard(rhs.borrow()) {
return Err(w);
}
}
Ok(self.with(rhs))
}
pub fn with_into<'c, R: Into<BorrowedOrOwned<'c, Pattern>>>(
&self,
rhs: R,
out: &mut Atom,
) -> bool {
let rhs = ReplaceWith::Pattern(rhs.into());
let mut expr_ref = self.target;
let mut out2 = RecycledAtom::new();
let mut replaced = false;
while expr_ref.replace_into(
self.pattern.borrow(),
&rhs,
self.conditions.as_ref().map(|x| x.borrow()),
Some(&self.match_settings),
self.replace_settings,
out,
) {
replaced = true;
if !self.repeat || expr_ref == out.as_view() {
break;
}
std::mem::swap(out, &mut out2);
expr_ref = out2.as_view();
}
if !replaced {
out.set_from_view(&self.target);
}
replaced
}
pub fn with_map<'c, R: MatchMap + 'static>(&self, rhs: R) -> Atom {
let rhs = ReplaceWith::Map(Box::new(rhs));
let mut expr_ref = self.target;
let mut out = RecycledAtom::new();
let mut out2 = RecycledAtom::new();
while expr_ref.replace_into(
self.pattern.borrow(),
&rhs,
self.conditions.as_ref().map(|x| x.borrow()),
Some(&self.match_settings),
self.replace_settings,
&mut out,
) {
if !self.repeat {
break;
}
std::mem::swap(&mut out, &mut out2);
expr_ref = out2.as_view();
}
out.into_inner()
}
pub fn iter<'c, R: Into<BorrowedOrOwned<'a, Pattern>>>(
&'a self,
rhs: R,
) -> ReplaceIterator<'a, 'a> {
ReplaceIterator::new(
self.pattern.borrow(),
self.target,
ReplaceWith::Pattern(rhs.into()),
self.conditions.as_ref().map(|x| x.borrow()),
Some(&self.match_settings),
)
}
pub fn iter_map<R: MatchMap + 'static>(&'a self, rhs: R) -> ReplaceIterator<'a, 'a> {
ReplaceIterator::new(
self.pattern.borrow(),
self.target,
ReplaceWith::Map(Box::new(rhs)),
self.conditions.as_ref().map(|x| x.borrow()),
Some(&self.match_settings),
)
}
pub fn match_iter(&self) -> PatternAtomTreeIterator<'_, '_> {
PatternAtomTreeIterator::new(
&self.pattern,
self.target,
self.conditions.as_ref().map(|x| x.borrow()),
Some(&self.match_settings),
)
}
}
impl<'a: 'b, 'b> IntoIterator for &'a ReplaceBuilder<'a, 'b> {
type Item = HashMap<Symbol, Atom>;
type IntoIter = PatternAtomTreeIterator<'a, 'b>;
fn into_iter(self) -> Self::IntoIter {
PatternAtomTreeIterator::new(
&self.pattern,
self.target,
self.conditions.as_ref().map(|x| x.borrow()),
Some(&self.match_settings),
)
}
}
impl From<Atom> for BorrowedOrOwned<'_, Pattern> {
fn from(atom: Atom) -> Self {
Pattern::from(atom).into()
}
}
impl From<Indeterminate> for BorrowedOrOwned<'_, Pattern> {
fn from(atom: Indeterminate) -> Self {
Pattern::from(atom).into()
}
}
impl From<Symbol> for BorrowedOrOwned<'_, Pattern> {
fn from(atom: Symbol) -> Self {
Pattern::from(atom).into()
}
}
impl<T: Into<Coefficient>> From<T> for BorrowedOrOwned<'_, Pattern> {
fn from(val: T) -> Self {
Atom::num(val.into()).to_pattern().into()
}
}
#[derive(Clone, Copy, Debug)]
pub struct Context {
pub function_level: usize,
pub parent_type: Option<AtomType>,
pub index: usize,
pub child_changed: bool,
}
impl<'a> AtomView<'a> {
pub(crate) fn to_pattern(self) -> Pattern {
Pattern::from_view(self, true)
}
pub(crate) fn is_scalar(&self) -> bool {
match self {
AtomView::Num(_) => true,
AtomView::Var(v) => v.get_symbol().is_scalar(),
AtomView::Fun(f) => f.get_symbol().is_scalar(),
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.is_scalar() && exp.is_scalar()
}
AtomView::Mul(m) => m.iter().all(|child| child.is_scalar()),
AtomView::Add(a) => a.iter().all(|child| child.is_scalar()),
}
}
pub(crate) fn is_integer(&self) -> bool {
match self {
AtomView::Num(n) => n.get_coeff_view().is_integer(),
AtomView::Var(v) => v.get_symbol().is_integer(),
AtomView::Fun(f) => f.get_symbol().is_integer(),
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.is_integer() && exp.is_integer()
}
AtomView::Mul(m) => m.iter().all(|child| child.is_integer()),
AtomView::Add(a) => a.iter().all(|child| child.is_integer()),
}
}
pub(crate) fn is_real(&self) -> bool {
match self {
AtomView::Num(n) => n.get_coeff_view().is_real(),
AtomView::Var(v) => v.get_symbol().is_real(),
AtomView::Fun(f) => {
let s = f.get_symbol();
match s.get_id() {
Symbol::EXP_ID | Symbol::SIN_ID | Symbol::COS_ID => {
f.iter().next().is_some_and(|arg| arg.is_real())
}
Symbol::SQRT_ID | Symbol::LOG_ID => {
f.iter().next().is_some_and(|arg| arg.is_positive())
}
Symbol::IF_ID => {
let mut iter = f.iter();
iter.next().is_some()
&& iter.next().is_some_and(|arg| arg.is_real())
&& iter.next().is_some_and(|arg| arg.is_real())
}
_ => s.is_real(),
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.is_real() && (exp.is_integer() || base.is_positive() && exp.is_real())
}
AtomView::Mul(m) => m.iter().all(|child| child.is_real()),
AtomView::Add(a) => a.iter().all(|child| child.is_real()),
}
}
#[inline]
pub fn has_attributes_of(&self, s: Symbol) -> bool {
if let Some(ss) = self.get_symbol() {
return ss.has_attributes_of(s);
}
!s.is_antisymmetric()
&& !s.is_symmetric()
&& !s.is_cyclesymmetric()
&& !s.is_linear()
&& s.get_tags().is_empty()
&& (!s.is_positive() || self.is_positive())
&& (!s.is_integer() || self.is_integer())
&& (!s.is_real() || self.is_real())
&& (!s.is_scalar() || self.is_scalar())
}
pub(crate) fn is_positive(&self) -> bool {
match self {
AtomView::Num(_) => {
if let Ok(k) = Rational::try_from(*self) {
!k.is_negative()
} else {
false
}
}
AtomView::Var(v) => v.get_symbol().is_positive(),
AtomView::Fun(f) => {
let s = f.get_symbol();
match s.get_id() {
Symbol::EXP_ID => f.iter().next().is_some_and(|arg| arg.is_real()),
Symbol::SQRT_ID => f.iter().next().is_some_and(|arg| arg.is_positive()),
Symbol::IF_ID => {
let mut iter = f.iter();
iter.next().is_some()
&& iter.next().is_some_and(|arg| arg.is_positive())
&& iter.next().is_some_and(|arg| arg.is_positive())
}
_ => s.is_positive(),
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
if let AtomView::Num(_) = exp
&& let Ok(k) = Rational::try_from(exp)
&& k.is_integer()
&& k.numerator_ref() % 2 == 0
{
return base.is_real();
}
base.is_positive() && exp.is_real()
}
AtomView::Mul(m) => m.iter().all(|child| child.is_positive()),
AtomView::Add(a) => a.iter().all(|child| child.is_positive()),
}
}
pub(crate) fn is_finite(&self) -> bool {
match self {
AtomView::Num(n) => !matches!(
n.get_coeff_view(),
CoefficientView::Infinity(_) | CoefficientView::Indeterminate
),
AtomView::Var(_) => true,
AtomView::Fun(f) => f.iter().all(|arg| arg.is_finite()),
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.is_finite() && exp.is_finite()
}
AtomView::Mul(m) => m.iter().all(|child| child.is_finite()),
AtomView::Add(a) => a.iter().all(|child| child.is_finite()),
}
}
pub(crate) fn is_constant(&self) -> bool {
match self {
AtomView::Num(n) => match n.get_coeff_view() {
CoefficientView::RationalPolynomial(r) => r.deserialize().is_constant(),
_ => true,
},
AtomView::Var(v) => match v.get_symbol_id() {
Symbol::PI_ID | Symbol::E_ID => true,
_ => false,
},
AtomView::Fun(f) => match f.get_symbol_id() {
Symbol::EXP_ID
| Symbol::LOG_ID
| Symbol::SQRT_ID
| Symbol::SIN_ID
| Symbol::COS_ID => {
f.get_nargs() == 1 && f.iter().next().is_some_and(|arg| arg.is_constant())
}
_ => false,
},
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.is_constant() && exp.is_constant()
}
AtomView::Mul(m) => m.iter().all(|child| child.is_constant()),
AtomView::Add(a) => a.iter().all(|child| child.is_constant()),
}
}
pub(crate) fn get_all_symbols(&self, include_function_symbols: bool) -> HashSet<Symbol> {
let mut out = HashSet::default();
self.get_all_symbols_impl(include_function_symbols, &mut out);
out
}
pub(crate) fn get_all_symbols_impl(
&self,
include_function_symbols: bool,
out: &mut HashSet<Symbol>,
) {
match self {
AtomView::Num(_) => {}
AtomView::Var(v) => {
out.insert(v.get_symbol());
}
AtomView::Fun(f) => {
if include_function_symbols {
out.insert(f.get_symbol());
}
for arg in f {
arg.get_all_symbols_impl(include_function_symbols, out);
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.get_all_symbols_impl(include_function_symbols, out);
exp.get_all_symbols_impl(include_function_symbols, out);
}
AtomView::Mul(m) => {
for child in m {
child.get_all_symbols_impl(include_function_symbols, out);
}
}
AtomView::Add(a) => {
for child in a {
child.get_all_symbols_impl(include_function_symbols, out);
}
}
}
}
pub(crate) fn get_all_indeterminates(&self, enter_functions: bool) -> HashSet<AtomView<'a>> {
let mut out = HashSet::default();
self.get_all_indeterminates_impl(enter_functions, &mut out);
out
}
fn get_all_indeterminates_impl(&self, enter_functions: bool, out: &mut HashSet<AtomView<'a>>) {
match self {
AtomView::Num(_) => {}
AtomView::Var(_) => {
out.insert(*self);
}
AtomView::Fun(f) => {
out.insert(*self);
if enter_functions {
for arg in f {
arg.get_all_indeterminates_impl(enter_functions, out);
}
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.get_all_indeterminates_impl(enter_functions, out);
exp.get_all_indeterminates_impl(enter_functions, out);
}
AtomView::Mul(m) => {
for child in m {
child.get_all_indeterminates_impl(enter_functions, out);
}
}
AtomView::Add(a) => {
for child in a {
child.get_all_indeterminates_impl(enter_functions, out);
}
}
}
}
pub(crate) fn count_indeterminates(
&self,
enter_functions: bool,
out: &mut HashMap<AtomView<'a>, usize>,
) {
match self {
AtomView::Num(_) => {}
AtomView::Var(_) => {
*out.entry(*self).or_insert(0) += 1;
}
AtomView::Fun(f) => {
*out.entry(*self).or_insert(0) += 1;
if enter_functions {
for arg in f {
arg.count_indeterminates(enter_functions, out);
}
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
base.count_indeterminates(enter_functions, out);
exp.count_indeterminates(enter_functions, out);
}
AtomView::Mul(m) => {
for child in m {
child.count_indeterminates(enter_functions, out);
}
}
AtomView::Add(a) => {
for child in a {
child.count_indeterminates(enter_functions, out);
}
}
}
}
pub fn count_operations_with_subexpressions(
&self,
cses: &mut HashSet<AtomView<'a>>,
) -> OperationCount {
let mut count = OperationCount::default();
let mut counter = |a: AtomView<'a>| {
if !cses.insert(a) {
return false;
}
match a {
AtomView::Mul(m) => {
count.multiplications += m.get_nargs() - 1;
true
}
AtomView::Add(a) => {
count.additions += a.get_nargs() - 1;
true
}
AtomView::Pow(p) => {
if let Ok(i) = isize::try_from(p.get_exp()) {
count.add_integer_power(i as i64);
} else {
count.add_function_call();
}
true
}
_ => true,
}
};
self.visitor(&mut counter);
count
}
pub(crate) fn contains_literally_or_as_symbol(&self, a: AtomView) -> bool {
match a {
AtomView::Var(v) => self.contains_symbol(v.get_symbol()),
_ => self.contains(a),
}
}
pub(crate) fn contains(&self, a: AtomView) -> bool {
let mut stack = Vec::with_capacity(20);
stack.push(*self);
while let Some(c) = stack.pop() {
if a == c {
return true;
}
if a.get_byte_size() > c.get_byte_size() {
continue;
}
match c {
AtomView::Num(_) | AtomView::Var(_) => {}
AtomView::Fun(f) => {
for arg in f {
stack.push(arg);
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
stack.push(base);
stack.push(exp);
}
AtomView::Mul(m) => {
for child in m {
stack.push(child);
}
}
AtomView::Add(a) => {
for child in a {
stack.push(child);
}
}
}
}
false
}
pub(crate) fn contains_indeterminate(&self, s: &Indeterminate) -> bool {
match s {
Indeterminate::Symbol(sym, _) => self.contains_symbol(*sym),
Indeterminate::Function(_, atom) => self.contains(atom.as_view()),
}
}
pub(crate) fn contains_symbol(&self, s: Symbol) -> bool {
let mut stack = Vec::with_capacity(20);
stack.push(*self);
while let Some(c) = stack.pop() {
match c {
AtomView::Num(_) => {}
AtomView::Var(v) => {
if v.get_symbol() == s {
return true;
}
}
AtomView::Fun(f) => {
if f.get_symbol() == s {
return true;
}
for arg in f {
stack.push(arg);
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
stack.push(base);
stack.push(exp);
}
AtomView::Mul(m) => {
for child in m {
stack.push(child);
}
}
AtomView::Add(a) => {
for child in a {
stack.push(child);
}
}
}
}
false
}
pub(crate) fn visitor<F: FnMut(AtomView<'a>) -> bool>(&self, v: &mut F) {
match self {
AtomView::Num(_) | AtomView::Var(_) => {
v(*self);
}
AtomView::Fun(f) => {
if !v(*self) {
return;
}
for arg in f {
arg.visitor(v);
}
}
AtomView::Pow(p) => {
if !v(*self) {
return;
}
let (base, exp) = p.get_base_exp();
base.visitor(v);
exp.visitor(v);
}
AtomView::Mul(m) => {
if !v(*self) {
return;
}
for child in m {
child.visitor(v);
}
}
AtomView::Add(a) => {
if !v(*self) {
return;
}
for child in a {
child.visitor(v);
}
}
}
}
pub(crate) fn alias_subexpressions(
&self,
mut f: impl FnMut(AtomView, usize, usize) -> Option<Atom>,
) -> AliasedAtom {
let mut subexpressions = HashMap::default();
self.count_subexpressions(&mut subexpressions);
let mut subexpr_vec: Vec<_> = subexpressions.into_iter().collect();
subexpr_vec.sort_by(|(k1, _), (k2, _)| k2.get_byte_size().cmp(&k1.get_byte_size()));
let mut subexpr_corrections = HashMap::default();
let mut subs: HashMap<Atom, Atom> = HashMap::default();
let mut inv_subs = HashMap::default();
for (subexpr, mut count) in subexpr_vec.drain(..) {
count += subexpr_corrections.get(&subexpr).cloned().unwrap_or(0);
if count == 1 {
continue;
}
if let Some(replacement) = f(subexpr, count, subs.len()) {
subs.insert(replacement.clone(), subexpr.to_owned());
inv_subs.insert(subexpr, replacement);
} else {
let mut subexpr_correction = HashMap::default();
subexpr.count_subexpressions(&mut subexpr_correction);
for (k, v) in subexpr_correction {
*subexpr_corrections.entry(k).or_insert(0) += v * (count - 1);
}
}
}
let replaced_atom = self.replace_map(|a, _, out| {
if let Some(replacement) = inv_subs.get(&a) {
out.set_from_view(&replacement.as_view());
}
});
for x in subs.values_mut() {
*x = x.replace_map(|a, _, out| {
if a != x.as_view()
&& let Some(replacement) = inv_subs.get(&a)
{
out.set_from_view(&replacement.as_view());
}
});
}
AliasedAtom {
root: replaced_atom,
aliases: subs,
}
}
pub(crate) fn count_subexpressions(&self, subexpressions: &mut HashMap<AtomView<'a>, usize>) {
if subexpressions.contains_key(self) {
*subexpressions.entry(*self).or_insert(0) += 1;
return;
}
match self {
AtomView::Num(_) | AtomView::Var(_) => {}
AtomView::Fun(f) => {
*subexpressions.entry(*self).or_insert(0) += 1;
for arg in f {
arg.count_subexpressions(subexpressions);
}
}
AtomView::Pow(p) => {
*subexpressions.entry(*self).or_insert(0) += 1;
let (base, exp) = p.get_base_exp();
base.count_subexpressions(subexpressions);
exp.count_subexpressions(subexpressions);
}
AtomView::Mul(m) => {
*subexpressions.entry(*self).or_insert(0) += 1;
for child in m {
child.count_subexpressions(subexpressions);
}
}
AtomView::Add(a) => {
*subexpressions.entry(*self).or_insert(0) += 1;
for child in a {
child.count_subexpressions(subexpressions);
}
}
}
}
pub fn has_complex_coefficients(&self) -> bool {
let mut has_complex_coefficient = false;
self.visitor(&mut |a| {
if let AtomView::Num(n) = a
&& !n.get_coeff_view().is_real()
{
has_complex_coefficient = true;
}
!has_complex_coefficient
});
has_complex_coefficient
}
pub fn has_roots(&self) -> bool {
let mut has_roots = false;
self.visitor(&mut |a| {
if let AtomView::Pow(p) = a {
let (_, exp) = p.get_base_exp();
if let AtomView::Num(n) = exp {
if !n.get_coeff_view().is_integer() {
has_roots = true;
}
} else {
has_roots = true;
}
} else if let AtomView::Fun(f) = a
&& f.get_symbol_id() == Symbol::SQRT_ID
{
has_roots = true;
}
!has_roots
});
has_roots
}
pub(crate) fn is_polynomial(
&self,
allow_not_expanded: bool,
allow_negative_powers: bool,
) -> Option<HashSet<AtomView<'a>>> {
let mut vars = HashMap::default();
let mut symbol_cache = HashSet::default();
if self.is_polynomial_impl(
allow_not_expanded,
allow_negative_powers,
&mut vars,
&mut symbol_cache,
) {
symbol_cache.clear();
for (k, v) in vars {
if v {
symbol_cache.insert(k);
}
}
Some(symbol_cache)
} else {
None
}
}
fn is_polynomial_impl(
&self,
allow_not_expanded: bool,
allow_negative_powers: bool,
variables: &mut HashMap<AtomView<'a>, bool>,
symbol_cache: &mut HashSet<AtomView<'a>>,
) -> bool {
if let Some(x) = variables.get(self) {
return *x;
}
macro_rules! block_check {
($e: expr) => {
symbol_cache.clear();
$e.get_all_indeterminates_impl(true, symbol_cache);
for x in symbol_cache.drain() {
if variables.contains_key(&x) {
return false;
} else {
variables.insert(x, false); }
}
variables.insert(*$e, true); };
}
match self {
AtomView::Num(_) => true,
AtomView::Var(_) => {
variables.insert(*self, true);
true
}
AtomView::Fun(_) => {
block_check!(self);
true
}
AtomView::Pow(pow_view) => {
let (base, exp) = pow_view.get_base_exp();
if let AtomView::Num(_) = exp {
let (positive, integer) = if let Ok(k) = i64::try_from(exp) {
(k >= 0, true)
} else {
(false, false)
};
if integer && (allow_negative_powers || positive) {
if variables.get(&base) == Some(&true) {
return true;
}
if allow_not_expanded && positive {
return base.is_polynomial_impl(
allow_not_expanded,
allow_negative_powers,
variables,
symbol_cache,
);
}
block_check!(&base);
return true;
}
}
block_check!(self);
true
}
AtomView::Mul(mul_view) => {
for child in mul_view {
if !allow_not_expanded && let AtomView::Add(_) = child {
if variables.get(&child) == Some(&true) {
continue;
}
block_check!(&child);
continue;
}
if !child.is_polynomial_impl(
allow_not_expanded,
allow_negative_powers,
variables,
symbol_cache,
) {
return false;
}
}
true
}
AtomView::Add(add_view) => {
for child in add_view {
if !child.is_polynomial_impl(
allow_not_expanded,
allow_negative_powers,
variables,
symbol_cache,
) {
return false;
}
}
true
}
}
}
pub(crate) fn replace_map<F: FnMut(AtomView, &Context, &mut Settable<'_, Atom>)>(
&self,
mut m: F,
) -> Atom {
let mut out = Atom::new();
let context = Context {
function_level: 0,
parent_type: None,
index: 0,
child_changed: false,
};
Workspace::get_local().with(|ws| {
let mut set = Settable::from(&mut out);
self.replace_map_no_norm(ws, &mut m, context, &mut set);
if set.is_set() {
let mut a = ws.new_atom();
set.as_view().normalize(ws, &mut a);
std::mem::swap(&mut out, &mut a);
} else {
out.set_from_view(self);
}
});
out
}
pub(crate) fn replace_map_no_norm<F: FnMut(AtomView, &Context, &mut Settable<'_, Atom>)>(
&self,
ws: &Workspace,
m: &mut F,
mut context: Context,
out: &mut Settable<'_, Atom>,
) {
m(*self, &context, out);
if out.is_set() {
return;
}
match self {
AtomView::Num(_) | AtomView::Var(_) => {}
AtomView::Fun(f) => {
let mut fun = None;
context.parent_type = Some(AtomType::Fun);
context.function_level += 1;
let mut arg_h = ws.new_atom();
for (i, arg) in f.iter().enumerate() {
let mut set = Settable::from(arg_h.deref_mut());
context.index = i;
arg.replace_map_no_norm(ws, m, context, &mut set);
if fun.is_none() && set.is_set() {
let fun_o = out.to_fun(f.get_symbol());
for child in f.iter().take(i) {
fun_o.add_arg(child);
}
fun_o.add_arg(set.as_view());
fun = Some(fun_o);
} else if let Some(fun) = &mut fun {
if set.is_set() {
fun.add_arg(set.as_view());
} else {
fun.add_arg(arg);
}
}
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
context.parent_type = Some(AtomType::Pow);
context.index = 0;
let mut base_h = ws.new_atom();
let mut base_set = Settable::from(base_h.deref_mut());
base.replace_map_no_norm(ws, m, context, &mut base_set);
context.index = 1;
let mut exp_h = ws.new_atom();
let mut exp_set = Settable::from(exp_h.deref_mut());
exp.replace_map_no_norm(ws, m, context, &mut exp_set);
if base_set.is_set() && exp_set.is_set() {
out.to_pow(base_set.as_view(), exp_set.as_view());
} else if base_set.is_set() {
out.to_pow(base_set.as_view(), exp);
} else if exp_set.is_set() {
out.to_pow(base, exp_set.as_view());
}
}
AtomView::Mul(mm) => {
let mut mul = None;
context.parent_type = Some(AtomType::Mul);
let mut child_h = ws.new_atom();
for (i, child) in mm.iter().enumerate() {
let mut set = Settable::from(child_h.deref_mut());
context.index = i;
child.replace_map_no_norm(ws, m, context, &mut set);
if mul.is_none() && set.is_set() {
let mul_o = out.to_mul();
for child in mm.iter().take(i) {
mul_o.extend(child);
}
mul_o.extend(set.as_view());
mul = Some(mul_o);
} else if let Some(mul_o) = &mut mul {
if set.is_set() {
mul_o.extend(set.as_view());
} else {
mul_o.extend(child);
}
}
}
}
AtomView::Add(a) => {
let mut add = None;
context.parent_type = Some(AtomType::Add);
let mut child_h = ws.new_atom();
for (i, child) in a.iter().enumerate() {
let mut set = Settable::from(child_h.deref_mut());
context.index = i;
child.replace_map_no_norm(ws, m, context, &mut set);
if add.is_none() && set.is_set() {
let add_o = out.to_add();
for child in a.iter().take(i) {
add_o.extend(child);
}
add_o.extend(set.as_view());
add = Some(add_o);
} else if let Some(mul_o) = &mut add {
if set.is_set() {
mul_o.extend(set.as_view());
} else {
mul_o.extend(child);
}
}
}
}
}
}
pub(crate) fn replace_map_bottom_up<F: FnMut(AtomView, &Context, &mut Settable<'_, Atom>)>(
&self,
mut m: F,
nested: bool, ) -> Atom {
let mut out = Atom::new();
let context = Context {
function_level: 0,
parent_type: None,
index: 0,
child_changed: false,
};
Workspace::get_local().with(|ws| {
let mut set = Settable::from(&mut out);
self.replace_map_bottom_up_impl(ws, &mut m, context, nested, &mut set);
if set.is_set() {
let mut a = ws.new_atom();
set.as_view().normalize(ws, &mut a);
std::mem::swap(&mut out, &mut a);
} else {
out.set_from_view(self);
}
});
out
}
pub(crate) fn replace_map_bottom_up_impl<
F: FnMut(AtomView, &Context, &mut Settable<'_, Atom>),
>(
&self,
ws: &Workspace,
m: &mut F,
mut parent_context: Context,
nested: bool,
out: &mut Settable<'_, Atom>,
) {
let mut context = parent_context;
match self {
AtomView::Num(_) | AtomView::Var(_) => {}
AtomView::Fun(f) => {
let mut fun = None;
context.parent_type = Some(AtomType::Fun);
context.function_level += 1;
let mut arg_h = ws.new_atom();
for (i, arg) in f.iter().enumerate() {
let mut set = Settable::from(arg_h.deref_mut());
context.index = i;
arg.replace_map_bottom_up_impl(ws, m, context, nested, &mut set);
if fun.is_none() && set.is_set() {
parent_context.child_changed = true;
let fun_o = out.to_fun(f.get_symbol());
for child in f.iter().take(i) {
fun_o.add_arg(child);
}
fun_o.add_arg(set.as_view());
fun = Some(fun_o);
} else if let Some(fun) = &mut fun {
if set.is_set() {
fun.add_arg(set.as_view());
} else {
fun.add_arg(arg);
}
}
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
context.parent_type = Some(AtomType::Pow);
context.index = 0;
let mut base_h = ws.new_atom();
let mut base_set = Settable::from(base_h.deref_mut());
base.replace_map_bottom_up_impl(ws, m, context, nested, &mut base_set);
context.index = 1;
let mut exp_h = ws.new_atom();
let mut exp_set = Settable::from(exp_h.deref_mut());
exp.replace_map_bottom_up_impl(ws, m, context, nested, &mut exp_set);
if base_set.is_set() && exp_set.is_set() {
parent_context.child_changed = true;
out.to_pow(base_set.as_view(), exp_set.as_view());
} else if base_set.is_set() {
parent_context.child_changed = true;
out.to_pow(base_set.as_view(), exp);
} else if exp_set.is_set() {
parent_context.child_changed = true;
out.to_pow(base, exp_set.as_view());
}
}
AtomView::Mul(mm) => {
let mut mul = None;
context.parent_type = Some(AtomType::Mul);
let mut child_h = ws.new_atom();
for (i, child) in mm.iter().enumerate() {
let mut set = Settable::from(child_h.deref_mut());
context.index = i;
child.replace_map_bottom_up_impl(ws, m, context, nested, &mut set);
if mul.is_none() && set.is_set() {
parent_context.child_changed = true;
let mul_o = out.to_mul();
for child in mm.iter().take(i) {
mul_o.extend(child);
}
mul_o.extend(set.as_view());
mul = Some(mul_o);
} else if let Some(mul_o) = &mut mul {
if set.is_set() {
mul_o.extend(set.as_view());
} else {
mul_o.extend(child);
}
}
}
}
AtomView::Add(a) => {
let mut add = None;
context.parent_type = Some(AtomType::Add);
let mut child_h = ws.new_atom();
for (i, child) in a.iter().enumerate() {
let mut set = Settable::from(child_h.deref_mut());
context.index = i;
child.replace_map_bottom_up_impl(ws, m, context, nested, &mut set);
if add.is_none() && set.is_set() {
parent_context.child_changed = true;
let add_o = out.to_add();
for child in a.iter().take(i) {
add_o.extend(child);
}
add_o.extend(set.as_view());
add = Some(add_o);
} else if let Some(mul_o) = &mut add {
if set.is_set() {
mul_o.extend(set.as_view());
} else {
mul_o.extend(child);
}
}
}
}
}
if !parent_context.child_changed {
m(*self, &parent_context, out);
} else if nested {
let mut norm = ws.new_atom();
out.as_view().normalize(ws, &mut norm);
std::mem::swap(out.deref_mut(), norm.deref_mut());
let mut child_h = ws.new_atom();
let mut set = Settable::from(child_h.deref_mut());
m(out.as_view(), &parent_context, &mut set);
if set.is_set() {
std::mem::swap(out.deref_mut(), child_h.deref_mut());
}
}
}
pub(crate) fn replace<'b, P: Into<BorrowedOrOwned<'b, Pattern>>>(
&self,
pattern: P,
) -> ReplaceBuilder<'a, 'b> {
ReplaceBuilder::new(*self, pattern)
}
pub(crate) fn replace_into<'b, R: Into<&'b ReplaceWith<'b>>>(
&self,
pattern: &Pattern,
rhs: R,
conditions: Option<&Condition<PatternRestriction>>,
settings: Option<&MatchSettings>,
replace_settings: ReplaceSettings,
out: &mut Atom,
) -> bool {
Workspace::get_local().with(|ws| {
self.replace_with_ws_into(
pattern,
rhs.into(),
ws,
conditions,
settings,
replace_settings,
out,
)
})
}
pub(crate) fn replace_multiple<I, T>(
&self,
replacements: I,
replace_settings: ReplaceSettings,
) -> Atom
where
I: IntoIterator<Item = T>,
T: BorrowReplacement,
{
let mut out = Atom::new();
self.replace_multiple_into(replacements, replace_settings, &mut out);
out
}
pub(crate) fn replace_multiple_into<I, T>(
&self,
replacements: I,
replace_settings: ReplaceSettings,
out: &mut Atom,
) -> bool
where
I: IntoIterator<Item = T>,
T: BorrowReplacement,
{
let replacements = replacements.into_iter().collect::<Vec<_>>();
let mut atom_iter = replacements
.iter()
.map(|r| {
(
AtomMatchIterator::new(r.borrow().pattern),
WrappedMatchStack::from_replacement(r.borrow()),
)
})
.collect::<Vec<_>>();
let max_level = atom_iter.iter().fold(Some((0, true)), |acc, (_, stack)| {
acc.and_then(|(max_level, tree)| {
stack.settings.level_range.1.map(|level| {
(
max_level.max(level),
tree && stack.settings.level_is_tree_depth,
)
})
})
});
Workspace::get_local().with(|ws| {
let mut rhs_cache = HashMap::default();
let mut set = Settable::from(&mut *out);
self.replace_no_norm(
&replacements,
&mut atom_iter,
ws,
0,
0,
max_level,
&mut rhs_cache,
replace_settings,
&mut set,
);
if set.is_set() {
let mut norm = ws.new_atom();
set.as_view().normalize(ws, &mut norm);
std::mem::swap(out, &mut norm);
true
} else {
out.set_from_view(self);
false
}
})
}
fn replace_node<'b, T: BorrowReplacement>(
&self,
replacements: &'b [T],
atom_match_iterators: &mut [(AtomMatchIterator<'a, 'b>, WrappedMatchStack<'a, 'b>)],
workspace: &Workspace,
tree_level: usize,
fn_level: usize,
rhs_cache: &mut HashMap<(usize, Vec<(Symbol, Match<'a>)>), Atom>,
out: &mut Settable<Atom>,
) -> bool {
let mut fits = false;
for (rep_id, r) in replacements.iter().enumerate() {
let r = r.borrow();
if let Pattern::Literal(l) = &r.pattern {
if l.as_view().get_byte_size() <= self.get_byte_size() {
fits = true;
}
} else {
fits = true;
}
let settings = r.settings.unwrap_or(&DEFAULT_MATCH_SETTINGS);
if !settings.partial && tree_level > 0 {
continue;
}
if let Some(max_level) = settings.level_range.1
&& (settings.level_is_tree_depth && tree_level > max_level
|| !settings.level_is_tree_depth && fn_level > max_level)
{
continue;
}
if settings.level_is_tree_depth && tree_level < settings.level_range.0
|| !settings.level_is_tree_depth && fn_level < settings.level_range.0
{
continue;
}
if r.pattern.could_match(*self) {
let (match_iter, match_stack) = &mut atom_match_iterators[rep_id];
match_stack.reset();
match_iter.set_new_target(*self, match_stack);
let it = match_iter;
if let Some(used_flags) = it.next(match_stack) {
let mut rhs_subs = workspace.new_atom();
let key = (rep_id, std::mem::take(&mut match_stack.stack.stack));
if let Some(rhs) = rhs_cache.get(&key) {
match_stack.stack.stack = key.1;
rhs_subs.set_from_view(&rhs.as_view());
} else {
match_stack.stack.stack = key.1;
match r.rhs {
ReplaceWith::Pattern(rhs) => {
rhs.replace_wildcards_with_matches_impl(
workspace,
&mut rhs_subs,
&match_stack.stack,
settings.allow_new_wildcards_on_rhs,
None,
)
.unwrap(); }
ReplaceWith::Map(f) => {
let mut rhs = f(&match_stack.stack);
std::mem::swap(rhs_subs.deref_mut(), &mut rhs);
}
}
if rhs_cache.len() < settings.rhs_cache_size
&& !matches!(
r.rhs,
ReplaceWith::Pattern(BorrowedOrOwned::Owned(Pattern::Literal(_)))
)
&& !matches!(
r.rhs,
ReplaceWith::Pattern(BorrowedOrOwned::Borrowed(Pattern::Literal(
_
)))
)
{
rhs_cache.insert(
(rep_id, match_stack.stack.stack.clone()),
rhs_subs.deref_mut().clone(),
);
}
}
if used_flags.iter().all(|x| *x) {
std::mem::swap(&mut *rhs_subs, out.deref_mut());
return fits;
}
match self {
AtomView::Mul(m) => {
let out = out.to_mul();
for (child, used) in m.iter().zip(used_flags) {
if !used {
out.extend(child);
}
}
out.extend(rhs_subs.as_view());
}
AtomView::Add(a) => {
let out = out.to_add();
for (child, used) in a.iter().zip(used_flags) {
if !used {
out.extend(child);
}
}
out.extend(rhs_subs.as_view());
}
_ => {
std::mem::swap(&mut *rhs_subs, out.deref_mut());
}
}
return fits;
}
}
}
fits
}
fn replace_no_norm<'b, T: BorrowReplacement>(
&self,
replacements: &'b [T],
atom_match_iterators: &mut [(AtomMatchIterator<'a, 'b>, WrappedMatchStack<'a, 'b>)],
workspace: &Workspace,
tree_level: usize,
fn_level: usize,
max_level: Option<(usize, bool)>,
rhs_cache: &mut HashMap<(usize, Vec<(Symbol, Match<'a>)>), Atom>,
replace_settings: ReplaceSettings,
out: &mut Settable<Atom>,
) {
if !replace_settings.bottom_up
&& !replace_settings.nested
&& (!self.replace_node(
replacements,
atom_match_iterators,
workspace,
tree_level,
fn_level,
rhs_cache,
out,
) || out.is_set())
{
return;
}
match self {
AtomView::Fun(f) => {
if let Some((max_level, _)) = max_level
&& fn_level >= max_level
{
return;
}
let mut fun = None;
let mut child_buf = workspace.new_atom();
for (i, arg) in f.iter().enumerate() {
let mut set = Settable::from(child_buf.deref_mut());
arg.replace_no_norm(
replacements,
atom_match_iterators,
workspace,
tree_level + 1,
fn_level + 1,
max_level,
rhs_cache,
replace_settings,
&mut set,
);
if fun.is_none() && set.is_set() {
let fun_o = out.to_fun(f.get_symbol());
if replace_settings.once {
for (index, child) in f.iter().enumerate() {
if index == i {
fun_o.add_arg(set.as_view());
} else {
fun_o.add_arg(child);
}
}
return;
}
for child in f.iter().take(i) {
fun_o.add_arg(child);
}
fun_o.add_arg(set.as_view());
fun = Some(fun_o);
} else if let Some(fun) = &mut fun {
if set.is_set() {
fun.add_arg(set.as_view());
} else {
fun.add_arg(arg);
}
}
}
}
AtomView::Pow(p) => {
if let Some((max_level, true)) = max_level
&& tree_level >= max_level
{
return;
}
let (base, exp) = p.get_base_exp();
let mut base_out = workspace.new_atom();
let mut base_set = Settable::from(base_out.deref_mut());
base.replace_no_norm(
replacements,
atom_match_iterators,
workspace,
tree_level + 1,
fn_level,
max_level,
rhs_cache,
replace_settings,
&mut base_set,
);
if base_set.is_set() && replace_settings.once {
out.to_pow(base_set.as_view(), exp);
return;
}
let mut exp_out = workspace.new_atom();
let mut exp_set = Settable::from(exp_out.deref_mut());
exp.replace_no_norm(
replacements,
atom_match_iterators,
workspace,
tree_level + 1,
fn_level,
max_level,
rhs_cache,
replace_settings,
&mut exp_set,
);
if base_set.is_set() && exp_set.is_set() {
out.to_pow(base_set.as_view(), exp_set.as_view());
} else if base_set.is_set() {
out.to_pow(base_set.as_view(), exp);
} else if exp_set.is_set() {
out.to_pow(base, exp_set.as_view());
}
}
AtomView::Mul(m) => {
if let Some((max_level, true)) = max_level
&& tree_level >= max_level
{
return;
}
let mut mul = None;
let mut child_buf = workspace.new_atom();
for (i, child) in m.iter().enumerate() {
let mut set = Settable::from(child_buf.deref_mut());
child.replace_no_norm(
replacements,
atom_match_iterators,
workspace,
tree_level + 1,
fn_level,
max_level,
rhs_cache,
replace_settings,
&mut set,
);
if mul.is_none() && set.is_set() {
let mul_o = out.to_mul();
if replace_settings.once {
for (index, child) in m.iter().enumerate() {
if index == i {
mul_o.extend(set.as_view());
} else {
mul_o.extend(child);
}
}
return;
}
for child in m.iter().take(i) {
mul_o.extend(child);
}
mul_o.extend(set.as_view());
mul = Some(mul_o);
} else if let Some(mul_o) = &mut mul {
if set.is_set() {
mul_o.extend(set.as_view());
} else {
mul_o.extend(child);
}
}
}
if let Some(mul) = &mut mul {
mul.set_has_coefficient(m.has_coefficient());
}
}
AtomView::Add(a) => {
if let Some((max_level, true)) = max_level
&& tree_level >= max_level
{
return;
}
let mut add = None;
let mut child_buf = workspace.new_atom();
for (i, child) in a.iter().enumerate() {
let mut set = Settable::from(child_buf.deref_mut());
child.replace_no_norm(
replacements,
atom_match_iterators,
workspace,
tree_level + 1,
fn_level,
max_level,
rhs_cache,
replace_settings,
&mut set,
);
if add.is_none() && set.is_set() {
let add_o = out.to_add();
if replace_settings.once {
for (index, child) in a.iter().enumerate() {
if index == i {
add_o.extend(set.as_view());
} else {
add_o.extend(child);
}
}
return;
}
for child in a.iter().take(i) {
add_o.extend(child);
}
add_o.extend(set.as_view());
add = Some(add_o);
} else if let Some(add_o) = &mut add {
if set.is_set() {
add_o.extend(set.as_view());
} else {
add_o.extend(child);
}
}
}
}
_ => {}
}
if replace_settings.bottom_up && !out.is_set() || replace_settings.nested {
if out.is_set() {
let mut buf = workspace.new_atom();
let mut set = Settable::from(buf.deref_mut());
let mut atom_iter = replacements
.iter()
.map(|r| {
(
AtomMatchIterator::new(r.borrow().pattern),
WrappedMatchStack::new(
r.borrow().conditions.unwrap_or(&DEFAULT_PATTERN_CONDITION),
r.borrow().settings.unwrap_or(&DEFAULT_MATCH_SETTINGS),
),
)
})
.collect::<Vec<_>>();
out.as_view().replace_node(
replacements,
&mut atom_iter,
workspace,
tree_level,
fn_level,
&mut HashMap::default(),
&mut set,
);
if set.is_set() {
std::mem::swap(out.deref_mut(), buf.deref_mut());
}
} else {
self.replace_node(
replacements,
atom_match_iterators,
workspace,
tree_level,
fn_level,
rhs_cache,
out,
);
}
}
}
pub(crate) fn replace_with_ws_into(
&self,
pattern: &Pattern,
rhs: &ReplaceWith,
workspace: &Workspace,
conditions: Option<&Condition<PatternRestriction>>,
settings: Option<&MatchSettings>,
replace_settings: ReplaceSettings,
out: &mut Atom,
) -> bool {
let rep = BorrowedReplacement {
pattern,
rhs,
conditions,
settings,
};
let mut atom_iter = std::slice::from_ref(&rep)
.iter()
.map(|r| {
(
AtomMatchIterator::new(r.pattern),
WrappedMatchStack::new(
r.conditions.unwrap_or(&DEFAULT_PATTERN_CONDITION),
r.settings.unwrap_or(&DEFAULT_MATCH_SETTINGS),
),
)
})
.collect::<Vec<_>>();
let max_level = atom_iter.iter().fold(Some((0, true)), |acc, (_, stack)| {
acc.and_then(|(max_level, tree)| {
stack.settings.level_range.1.map(|level| {
(
max_level.max(level),
tree && stack.settings.level_is_tree_depth,
)
})
})
});
let mut rhs_cache = HashMap::default();
let mut set = Settable::from(&mut *out);
self.replace_no_norm(
std::slice::from_ref(&rep),
&mut atom_iter,
workspace,
0,
0,
max_level,
&mut rhs_cache,
replace_settings,
&mut set,
);
if set.is_set() {
let mut norm = workspace.new_atom();
out.as_view().normalize(workspace, &mut norm);
std::mem::swap(out, &mut norm);
true
} else {
out.set_from_view(self);
false
}
}
}
impl Pattern {
pub fn new(atom: Atom) -> Pattern {
atom.to_pattern()
}
#[inline]
fn is_optional_wildcard(&self) -> bool {
matches!(self, Pattern::Wildcard(_, true))
}
pub fn to_atom(&self) -> Result<Atom, &'static str> {
Workspace::get_local().with(|ws| {
let mut out = Atom::new();
self.to_atom_impl(ws, &mut out)?;
Ok(out)
})
}
fn to_atom_impl(&self, ws: &Workspace, out: &mut Atom) -> Result<(), &'static str> {
match self {
Pattern::Literal(a) => {
out.set_from_view(&a.as_view());
}
Pattern::Wildcard(s, optional) => {
if *optional {
*out = s.optional();
} else {
out.to_var(*s);
}
}
Pattern::Fn(s, a) => {
let mut f = ws.new_atom();
let fun = f.to_fun(*s);
let mut arg_h = ws.new_atom();
for arg in a {
arg.to_atom_impl(ws, &mut arg_h)?;
fun.add_arg(arg_h.as_view());
}
f.as_view().normalize(ws, out);
}
Pattern::Pow(p) => {
let mut base = ws.new_atom();
p[0].to_atom_impl(ws, &mut base)?;
let mut exp = ws.new_atom();
p[1].to_atom_impl(ws, &mut exp)?;
let mut pow_h = ws.new_atom();
pow_h.to_pow(base.as_view(), exp.as_view());
pow_h.as_view().normalize(ws, out);
}
Pattern::Mul(m) => {
let mut mul_h = ws.new_atom();
let mul = mul_h.to_mul();
let mut arg_h = ws.new_atom();
for arg in m {
arg.to_atom_impl(ws, &mut arg_h)?;
mul.extend(arg_h.as_view());
}
mul_h.as_view().normalize(ws, out);
}
Pattern::Add(a) => {
let mut add_h = ws.new_atom();
let add = add_h.to_add();
let mut arg_h = ws.new_atom();
for arg in a {
arg.to_atom_impl(ws, &mut arg_h)?;
add.extend(arg_h.as_view());
}
add_h.as_view().normalize(ws, out);
}
Pattern::Alternative(a) => {
let mut f = ws.new_atom();
let fun = f.to_fun(Symbol::ALT);
let mut arg_h = ws.new_atom();
for arg in a {
arg.to_atom_impl(ws, &mut arg_h)?;
fun.add_arg(arg_h.as_view());
}
f.as_view().normalize(ws, out);
}
Pattern::Transformer(_) => Err("Cannot convert transformer to atom")?,
}
Ok(())
}
fn get_all_wildcards(&self) -> HashSet<Symbol> {
let mut wildcards = HashSet::default();
self.get_all_wildcards_impl(&mut wildcards);
wildcards
}
fn get_all_wildcards_impl(&self, wildcards: &mut HashSet<Symbol>) {
match self {
Pattern::Literal(_) => {}
Pattern::Wildcard(s, _) => {
wildcards.insert(*s);
}
Pattern::Fn(s, args) => {
if s.get_wildcard_level() > 0 {
wildcards.insert(*s);
}
for arg in args {
arg.get_all_wildcards_impl(wildcards);
}
}
Pattern::Pow(p) => {
p[0].get_all_wildcards_impl(wildcards);
p[1].get_all_wildcards_impl(wildcards);
}
Pattern::Mul(m) | Pattern::Add(m) => {
for arg in m {
arg.get_all_wildcards_impl(wildcards);
}
}
Pattern::Alternative(alts) => {
for alt in alts {
alt.get_all_wildcards_impl(wildcards);
}
}
Pattern::Transformer(_) => {}
}
}
pub fn alternative(self, rhs: Self) -> Result<Self, String> {
let wc_1 = self.get_all_wildcards();
let wc_2 = rhs.get_all_wildcards();
if wc_1 != wc_2 {
return Err(format!(
"Cannot create alternative pattern with different wildcards: {:?} and {:?}",
wc_1, wc_2
));
}
match (self, rhs) {
(Pattern::Alternative(mut alts1), Pattern::Alternative(alts2)) => {
alts1.extend(alts2);
Ok(Pattern::Alternative(alts1))
}
(Pattern::Alternative(mut alts), p) | (p, Pattern::Alternative(mut alts)) => {
alts.push(p);
Ok(Pattern::Alternative(alts))
}
(p1, p2) => Ok(Pattern::Alternative(vec![p1, p2])),
}
}
pub fn set_optional(self, symbol: Symbol) -> Self {
self.set_optional_impl(symbol)
}
fn set_optional_impl(self, symbol: Symbol) -> Self {
match self {
Pattern::Literal(t) => Pattern::Literal(t),
Pattern::Wildcard(w, mut optional) => {
if w == symbol {
optional = true;
}
Pattern::Wildcard(w, optional)
}
Pattern::Fn(s, args) => {
let new_args = args
.into_iter()
.map(|arg| arg.set_optional_impl(symbol))
.collect::<Vec<_>>();
Pattern::Fn(s, new_args)
}
Pattern::Pow(p) => {
let (base, exp) = (
p[0].clone().set_optional_impl(symbol),
p[1].clone().set_optional_impl(symbol),
);
Pattern::Pow(Box::new([base, exp]))
}
Pattern::Add(a) => {
let mut new_args = vec![];
for x in a {
new_args.push(x.set_optional_impl(symbol));
}
Pattern::Add(new_args)
}
Pattern::Mul(a) => {
let mut new_args = vec![];
for x in a {
new_args.push(x.set_optional_impl(symbol));
}
Pattern::Mul(new_args)
}
Pattern::Alternative(alts) => Pattern::Alternative(
alts.into_iter()
.map(|p| p.set_optional_impl(symbol))
.collect(),
),
Pattern::Transformer(t) => Pattern::Transformer(t),
}
}
pub fn add(&self, rhs: &Self, workspace: &Workspace) -> Self {
if let Pattern::Literal(l1) = self
&& let Pattern::Literal(l2) = rhs
{
let mut e = workspace.new_atom();
let a = e.to_add();
a.extend(l1.as_view());
a.extend(l2.as_view());
let mut b = Atom::default();
e.as_view().normalize(workspace, &mut b);
return Pattern::Literal(b);
}
let mut new_args = vec![];
if let Pattern::Add(l1) = self {
new_args.extend_from_slice(l1);
} else {
new_args.push(self.clone());
}
if let Pattern::Add(l1) = rhs {
new_args.extend_from_slice(l1);
} else {
new_args.push(rhs.clone());
}
Pattern::Add(new_args)
}
pub fn mul(&self, rhs: &Self, workspace: &Workspace) -> Self {
if let Pattern::Literal(l1) = self
&& let Pattern::Literal(l2) = rhs
{
let mut e = workspace.new_atom();
let a = e.to_mul();
a.extend(l1.as_view());
a.extend(l2.as_view());
let mut b = Atom::default();
e.as_view().normalize(workspace, &mut b);
return Pattern::Literal(b);
}
let mut new_args = vec![];
if let Pattern::Mul(l1) = self {
new_args.extend_from_slice(l1);
} else {
new_args.push(self.clone());
}
if let Pattern::Mul(l1) = rhs {
new_args.extend_from_slice(l1);
} else {
new_args.push(rhs.clone());
}
Pattern::Mul(new_args)
}
pub fn div(&self, rhs: &Self, workspace: &Workspace) -> Self {
if let Pattern::Literal(l2) = rhs {
let mut pow = workspace.new_atom();
pow.to_num(-1);
let mut e = workspace.new_atom();
e.to_pow(l2.as_view(), pow.as_view());
let mut b = Atom::default();
e.as_view().normalize(workspace, &mut b);
match self {
Pattern::Mul(m) => {
let mut new_args = m.clone();
new_args.push(Pattern::Literal(b));
Pattern::Mul(new_args)
}
Pattern::Literal(l1) => {
let mut m = workspace.new_atom();
let md = m.to_mul();
md.extend(l1.as_view());
md.extend(b.as_view());
let mut b = Atom::default();
m.as_view().normalize(workspace, &mut b);
Pattern::Literal(b)
}
_ => Pattern::Mul(vec![self.clone(), Pattern::Literal(b)]),
}
} else {
let rhs = Pattern::Pow(Box::new([rhs.clone(), Pattern::Literal(Atom::num(-1))]));
match self {
Pattern::Mul(m) => {
let mut new_args = m.clone();
new_args.push(rhs);
Pattern::Mul(new_args)
}
_ => Pattern::Mul(vec![self.clone(), rhs]),
}
}
}
pub fn pow(&self, rhs: &Self, workspace: &Workspace) -> Self {
if let Pattern::Literal(l1) = self
&& let Pattern::Literal(l2) = rhs
{
let mut e = workspace.new_atom();
e.to_pow(l1.as_view(), l2.as_view());
let mut b = Atom::default();
e.as_view().normalize(workspace, &mut b);
return Pattern::Literal(b);
}
Pattern::Pow(Box::new([self.clone(), rhs.clone()]))
}
pub fn neg(&self, workspace: &Workspace) -> Self {
if let Pattern::Literal(l1) = self {
let mut e = workspace.new_atom();
let a = e.to_mul();
let mut sign = workspace.new_atom();
sign.to_num(-1);
a.extend(l1.as_view());
a.extend(sign.as_view());
let mut b = Atom::default();
e.as_view().normalize(workspace, &mut b);
Pattern::Literal(b)
} else {
Pattern::Mul(vec![self.clone(), Pattern::Literal(Atom::num(-1))])
}
}
pub fn replace_with<'c, R: Into<BorrowedOrOwned<'c, Pattern>>>(self, rhs: R) -> Replacement {
Replacement::new(self, ReplaceWith::Pattern(rhs.into().into_owned()))
}
pub fn replace_with_map<'c, R: MatchMap + 'static>(self, rhs: R) -> Replacement {
Replacement::new(self, ReplaceWith::Map(Box::new(rhs)))
}
pub fn visitor<F: FnMut(&Pattern)>(&self, f: &mut F) {
f(self);
match self {
Pattern::Fn(_, a) => {
for p in a {
p.visitor(f);
}
}
Pattern::Pow(a) => {
a[0].visitor(f);
a[1].visitor(f);
}
Pattern::Mul(a) => {
for p in a {
p.visitor(f);
}
}
Pattern::Add(a) => {
for p in a {
p.visitor(f);
}
}
_ => {}
}
}
pub fn get_wildcard_function_names(&self) -> Vec<Symbol> {
let mut function_names = vec![];
self.visitor(&mut |p| {
if let Pattern::Fn(s, _) = p
&& s.get_wildcard_level() > 0
{
function_names.push(*s);
}
});
function_names
}
pub fn find_new_wildcard(&self, rhs: &Self) -> Option<Symbol> {
let mut wildcards = HashSet::default();
self.visitor(&mut |p| {
if let Pattern::Wildcard(w, _) = p {
wildcards.insert(*w);
}
});
let mut new_found = None;
rhs.visitor(&mut |p| {
if new_found.is_none()
&& let Pattern::Wildcard(w, _) = p
{
if wildcards.insert(*w) {
new_found = Some(*w);
};
}
});
new_found
}
#[inline]
fn could_match(&self, target: AtomView) -> bool {
match (self, target) {
(Pattern::Fn(f1, _), AtomView::Fun(f2)) => {
let s = f2.get_symbol();
f1.get_wildcard_level() > 0 && s.has_attributes_of(*f1) || *f1 == s
}
(Pattern::Fn(f1, args), AtomView::Var(v)) => {
let s = v.get_symbol();
(f1.get_wildcard_level() > 0 && s.has_attributes_of(*f1) || *f1 == s)
&& args.iter().all(Pattern::is_optional_wildcard)
}
(Pattern::Mul(_), AtomView::Mul(_)) => true,
(Pattern::Mul(args), _) => Self::optional_wildcards_can_match_single(args, target),
(Pattern::Add(_), AtomView::Add(_)) => true,
(Pattern::Add(args), _) => Self::optional_wildcards_can_match_single(args, target),
(Pattern::Wildcard(w, _), x) => x.has_attributes_of(*w),
(Pattern::Pow(_), AtomView::Pow(_)) => true,
(Pattern::Pow(args), _) => {
args[1].is_optional_wildcard() && args[0].could_match(target)
}
(Pattern::Literal(p), _) => p.as_view() == target,
(Pattern::Alternative(alternatives), _) => {
alternatives.iter().any(|p| p.could_match(target))
}
(Pattern::Transformer(_), _) => panic!("Pattern is a transformer"),
(_, _) => false,
}
}
fn optional_wildcards_can_match_single(args: &[Pattern], target: AtomView) -> bool {
let mut required_matches = 0;
for arg in args {
if arg.is_optional_wildcard() {
continue;
}
if !arg.could_match(target) {
return false;
}
required_matches += 1;
}
required_matches == 1
}
fn has_wildcard_or_alternative(atom: AtomView<'_>) -> bool {
match atom {
AtomView::Num(_) => false,
AtomView::Var(v) => v.get_wildcard_level() > 0,
AtomView::Fun(f) => {
let s = f.get_symbol();
if s.get_wildcard_level() > 0 || s == Symbol::ALT {
return true;
}
for arg in f {
if Self::has_wildcard_or_alternative(arg) {
return true;
}
}
false
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
Self::has_wildcard_or_alternative(base) || Self::has_wildcard_or_alternative(exp)
}
AtomView::Mul(m) => {
for child in m {
if Self::has_wildcard_or_alternative(child) {
return true;
}
}
false
}
AtomView::Add(a) => {
for child in a {
if Self::has_wildcard_or_alternative(child) {
return true;
}
}
false
}
}
}
pub(crate) fn from_view(atom: AtomView<'_>, is_top_layer: bool) -> Pattern {
fn sort_on_specificity(arg1: &Pattern, arg2: &Pattern) -> std::cmp::Ordering {
match (arg1, arg2) {
(Pattern::Literal(_), Pattern::Literal(_)) => std::cmp::Ordering::Equal,
(Pattern::Literal(_), _) => std::cmp::Ordering::Less,
(_, Pattern::Literal(_)) => std::cmp::Ordering::Greater,
(Pattern::Wildcard(w1, _), Pattern::Wildcard(w2, _)) => w1
.get_wildcard_level()
.cmp(&w2.get_wildcard_level())
.then_with(|| w2.has_attributes().cmp(&w1.has_attributes())), (Pattern::Wildcard(..), _) => std::cmp::Ordering::Greater, (_, Pattern::Wildcard(..)) => std::cmp::Ordering::Less,
(Pattern::Pow(p1), Pattern::Pow(p2)) => sort_on_specificity(&p1[0], &p2[0])
.then_with(|| sort_on_specificity(&p1[1], &p2[1])),
(Pattern::Pow(_), _) => std::cmp::Ordering::Less,
(_, Pattern::Pow(_)) => std::cmp::Ordering::Greater,
(Pattern::Fn(n1, arg1), Pattern::Fn(n2, arg2)) => n1
.get_wildcard_level()
.cmp(&n2.get_wildcard_level())
.then_with(|| arg1.len().cmp(&arg2.len()))
.then_with(|| {
arg1.iter()
.zip(arg2)
.fold(std::cmp::Ordering::Equal, |acc, (a1, a2)| {
acc.then_with(|| sort_on_specificity(a1, a2))
})
}),
(Pattern::Fn(_, _), _) => std::cmp::Ordering::Less,
(_, Pattern::Fn(_, _)) => std::cmp::Ordering::Greater,
(Pattern::Mul(m1), Pattern::Mul(m2)) => m1.len().cmp(&m2.len()).then_with(|| {
m1.iter()
.zip(m2)
.fold(std::cmp::Ordering::Equal, |acc, (a1, a2)| {
acc.then_with(|| sort_on_specificity(a1, a2))
})
.then_with(|| m1.len().cmp(&m2.len()))
}),
(Pattern::Mul(_), _) => std::cmp::Ordering::Less,
(_, Pattern::Mul(_)) => std::cmp::Ordering::Greater,
(Pattern::Add(a1), Pattern::Add(a2)) => a1.len().cmp(&a2.len()).then_with(|| {
a1.iter()
.zip(a2)
.fold(std::cmp::Ordering::Equal, |acc, (a1, a2)| {
acc.then_with(|| sort_on_specificity(a1, a2))
})
.then_with(|| a1.len().cmp(&a2.len()))
}),
(Pattern::Add(_), _) => std::cmp::Ordering::Less,
(_, Pattern::Add(_)) => std::cmp::Ordering::Greater,
(Pattern::Alternative(_), Pattern::Alternative(_)) => std::cmp::Ordering::Equal,
(Pattern::Alternative(_), _) => std::cmp::Ordering::Less,
(_, Pattern::Alternative(_)) => std::cmp::Ordering::Greater,
(Pattern::Transformer(_), Pattern::Transformer(_)) => std::cmp::Ordering::Equal,
}
}
if Self::has_wildcard_or_alternative(atom)
|| is_top_layer && matches!(atom, AtomView::Mul(_) | AtomView::Add(_))
{
match atom {
AtomView::Var(v) => Pattern::Wildcard(v.get_symbol(), false),
AtomView::Fun(f) => {
let name = f.get_symbol();
let mut args = Vec::with_capacity(f.get_nargs());
for arg in f {
args.push(Self::from_view(arg, false));
}
if name.is_symmetric() {
args.sort_unstable_by(sort_on_specificity);
}
if name == Symbol::ALT {
Pattern::Alternative(args)
} else if name == Symbol::OPT
&& args.len() == 1
&& let Pattern::Wildcard(w, _) = &args[0]
{
Pattern::Wildcard(w.clone(), true)
} else {
Pattern::Fn(name, args)
}
}
AtomView::Pow(p) => {
let (base, exp) = p.get_base_exp();
Pattern::Pow(Box::new([
Self::from_view(base, false),
Self::from_view(exp, false),
]))
}
AtomView::Mul(m) => {
let mut args = Vec::with_capacity(m.get_nargs());
for child in m {
args.push(Self::from_view(child, false));
}
args.sort_unstable_by(sort_on_specificity);
Pattern::Mul(args)
}
AtomView::Add(a) => {
let mut args = Vec::with_capacity(a.get_nargs());
for child in a {
args.push(Self::from_view(child, false));
}
args.sort_unstable_by(sort_on_specificity);
Pattern::Add(args)
}
AtomView::Num(_) => unreachable!("Number cannot have wildcard"),
}
} else {
let mut oa = Atom::default();
oa.set_from_view(&atom);
Pattern::Literal(oa)
}
}
pub fn replace_wildcards(&self, matches: &HashMap<Symbol, Atom>) -> Result<Atom, String> {
let mut out = Atom::new();
Workspace::get_local().with(|ws| self.replace_wildcards_impl(matches, ws, &mut out))?;
Ok(out)
}
fn replace_wildcards_impl(
&self,
matches: &HashMap<Symbol, Atom>,
ws: &Workspace,
out: &mut Atom,
) -> Result<(), String> {
match self {
Pattern::Literal(atom) => out.set_from_view(&atom.as_view()),
Pattern::Wildcard(symbol, _) => {
if let Some(a) = matches.get(symbol) {
out.set_from_view(&a.as_view());
} else {
out.to_var(*symbol);
}
}
Pattern::Fn(symbol, args) => {
let symbol = if let Some(a) = matches.get(symbol) {
if let Some(s) = a.as_view().get_symbol() {
s
} else {
return Err(format!(
"Wildcard function name expected for {}, got {}",
symbol.get_name(),
a
));
}
} else {
*symbol
};
let mut fun = ws.new_atom();
let f = fun.to_fun(symbol);
let mut arg = ws.new_atom();
for a in args {
a.replace_wildcards_impl(matches, ws, &mut arg)?;
f.add_arg(arg.as_view());
}
fun.as_view().normalize(ws, out);
}
Pattern::Pow(args) => {
let mut pow = ws.new_atom();
let mut base = ws.new_atom();
args[0].replace_wildcards_impl(matches, ws, &mut base)?;
let mut exp = ws.new_atom();
args[1].replace_wildcards_impl(matches, ws, &mut exp)?;
pow.to_pow(base.as_view(), exp.as_view());
pow.as_view().normalize(ws, out);
}
Pattern::Mul(args) => {
let mut mul = ws.new_atom();
let m = mul.to_mul();
let mut arg = ws.new_atom();
for a in args {
a.replace_wildcards_impl(matches, ws, &mut arg)?;
m.extend(arg.as_view());
}
mul.as_view().normalize(ws, out);
}
Pattern::Add(args) => {
let mut add = ws.new_atom();
let aa = add.to_add();
let mut arg = ws.new_atom();
for a in args {
a.replace_wildcards_impl(matches, ws, &mut arg)?;
aa.extend(arg.as_view());
}
add.as_view().normalize(ws, out);
}
Pattern::Alternative(_) => {
return Err(format!(
"Encountered alternative during substitution of wildcards from a map",
));
}
Pattern::Transformer(_) => {
return Err(format!(
"Encountered transformer during substitution of wildcards from a map",
));
}
}
Ok(())
}
pub fn replace_wildcards_with_matches(&self, match_stack: &MatchStack<'_>) -> Atom {
Workspace::get_local().with(|ws| {
let mut out = Atom::new();
self.replace_wildcards_with_matches_impl(ws, &mut out, match_stack, true, None)
.unwrap();
out
})
}
pub fn replace_wildcards_with_matches_impl(
&self,
workspace: &Workspace,
out: &mut Atom,
match_stack: &MatchStack<'_>,
allow_new_wildcards_on_rhs: bool,
transformer_input: Option<&Pattern>,
) -> Result<(), TransformerError> {
match self {
Pattern::Wildcard(name, _) => {
if let Some(w) = match_stack.get(*name) {
w.to_atom_into(out);
} else if allow_new_wildcards_on_rhs {
out.to_var(*name);
} else {
Err(TransformerError::ValueError(format!(
"Unsubstituted wildcard {name:?}",
)))?;
}
}
&Pattern::Fn(mut name, ref args) => {
if name.get_wildcard_level() > 0 {
if let Some(w) = match_stack.get(name) {
if let Match::FunctionName(fname) = w {
name = *fname;
} else if let Match::Single(a) = w {
if let AtomView::Var(v) = a {
name = v.get_symbol();
} else {
Err(TransformerError::ValueError(format!(
"Wildcard must be a function name instead of {}",
w.to_atom()
)))?;
}
} else {
Err(TransformerError::ValueError(format!(
"Wildcard must be a function name instead of {}",
w.to_atom()
)))?;
}
} else if !allow_new_wildcards_on_rhs {
Err(TransformerError::ValueError(format!(
"Unsubstituted wildcard {name:?}",
)))?;
}
}
let mut func_h = workspace.new_atom();
let func = func_h.to_fun(name);
for arg in args {
if let Pattern::Wildcard(w, _) = arg {
if let Some(w) = match_stack.get(*w) {
match w {
Match::Single(s) => func.add_arg(*s),
Match::Multiple(t, wargs) => match t {
SliceType::Arg | SliceType::Empty | SliceType::One => {
func.add_args(wargs);
}
_ => {
let mut handle = workspace.new_atom();
w.to_atom_into(&mut handle);
func.add_arg(handle.as_view())
}
},
Match::FunctionName(s) => {
func.add_arg(InlineVar::new(*s).as_view())
}
}
} else if allow_new_wildcards_on_rhs {
func.add_arg(workspace.new_var(*w).as_view())
} else {
Err(TransformerError::ValueError(format!(
"Unsubstituted wildcard {w:?}",
)))?;
}
continue;
}
let mut handle = workspace.new_atom();
arg.replace_wildcards_with_matches_impl(
workspace,
&mut handle,
match_stack,
allow_new_wildcards_on_rhs,
transformer_input,
)?;
func.add_arg(handle.as_view());
}
func_h.as_view().normalize(workspace, out);
}
Pattern::Pow(base_and_exp) => {
let mut base = workspace.new_atom();
let mut exp = workspace.new_atom();
let mut oas = [&mut base, &mut exp];
for (out, arg) in oas.iter_mut().zip(base_and_exp.iter()) {
if let Pattern::Wildcard(w, _) = arg {
if let Some(w) = match_stack.get(*w) {
match w {
Match::Single(s) => out.set_from_view(s),
Match::Multiple(_, _) => {
let mut handle = workspace.new_atom();
w.to_atom_into(&mut handle);
out.set_from_view(&handle.as_view())
}
Match::FunctionName(s) => {
out.set_from_view(&InlineVar::new(*s).as_view())
}
}
} else if allow_new_wildcards_on_rhs {
out.set_from_view(&workspace.new_var(*w).as_view());
} else {
Err(TransformerError::ValueError(format!(
"Unsubstituted wildcard {w:?}",
)))?;
}
continue;
}
let mut handle = workspace.new_atom();
arg.replace_wildcards_with_matches_impl(
workspace,
&mut handle,
match_stack,
allow_new_wildcards_on_rhs,
transformer_input,
)?;
out.set_from_view(&handle.as_view());
}
let mut pow_h = workspace.new_atom();
pow_h.to_pow(oas[0].as_view(), oas[1].as_view());
pow_h.as_view().normalize(workspace, out);
}
Pattern::Mul(args) => {
let mut mul_h = workspace.new_atom();
let mul = mul_h.to_mul();
for arg in args {
if let Pattern::Wildcard(w, _) = arg {
if let Some(w) = match_stack.get(*w) {
match w {
Match::Single(s) => mul.extend(*s),
Match::Multiple(t, wargs) => match t {
SliceType::Mul | SliceType::Empty | SliceType::One => {
for arg in wargs {
mul.extend(*arg);
}
}
_ => {
let mut handle = workspace.new_atom();
w.to_atom_into(&mut handle);
mul.extend(handle.as_view())
}
},
Match::FunctionName(s) => mul.extend(InlineVar::new(*s).as_view()),
}
} else if allow_new_wildcards_on_rhs {
mul.extend(workspace.new_var(*w).as_view());
} else {
Err(TransformerError::ValueError(format!(
"Unsubstituted wildcard {w:?}"
)))?;
}
continue;
}
let mut handle = workspace.new_atom();
arg.replace_wildcards_with_matches_impl(
workspace,
&mut handle,
match_stack,
allow_new_wildcards_on_rhs,
transformer_input,
)?;
mul.extend(handle.as_view());
}
mul_h.as_view().normalize(workspace, out);
}
Pattern::Add(args) => {
let mut add_h = workspace.new_atom();
let add = add_h.to_add();
for arg in args {
if let Pattern::Wildcard(w, _) = arg {
if let Some(w) = match_stack.get(*w) {
match w {
Match::Single(s) => add.extend(*s),
Match::Multiple(t, wargs) => match t {
SliceType::Add | SliceType::Empty | SliceType::One => {
for arg in wargs {
add.extend(*arg);
}
}
_ => {
let mut handle = workspace.new_atom();
w.to_atom_into(&mut handle);
add.extend(handle.as_view())
}
},
Match::FunctionName(s) => add.extend(InlineVar::new(*s).as_view()),
}
} else if allow_new_wildcards_on_rhs {
add.extend(workspace.new_var(*w).as_view());
} else {
Err(TransformerError::ValueError(format!(
"Unsubstituted wildcard {w:?}"
)))?;
}
continue;
}
let mut handle = workspace.new_atom();
arg.replace_wildcards_with_matches_impl(
workspace,
&mut handle,
match_stack,
allow_new_wildcards_on_rhs,
transformer_input,
)?;
add.extend(handle.as_view());
}
add_h.as_view().normalize(workspace, out);
}
Pattern::Literal(oa) => {
out.set_from_view(&oa.as_view());
}
Pattern::Alternative(_) => Err(TransformerError::ValueError(
"Cannot replace wildcards in an alternative pattern.".to_owned(),
))?,
Pattern::Transformer(p) => {
let (pat, ts) = &**p;
let pat = if let Some(p) = pat.as_ref() {
p
} else if let Some(input_p) = transformer_input {
input_p
} else {
Err(TransformerError::ValueError(
"Transformer is missing an expression to act on.".to_owned(),
))?
};
let mut handle = workspace.new_atom();
pat.replace_wildcards_with_matches_impl(
workspace,
&mut handle,
match_stack,
allow_new_wildcards_on_rhs,
transformer_input,
)?;
let _ = Transformer::execute_chain(
handle.as_view(),
ts,
workspace,
&Default::default(),
out,
)?;
}
}
Ok(())
}
}
impl std::fmt::Debug for Pattern {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Wildcard(arg0, arg1) => {
f.debug_tuple("Wildcard").field(arg0).field(arg1).finish()
}
Self::Fn(arg0, arg1) => f.debug_tuple("Fn").field(arg0).field(arg1).finish(),
Self::Pow(arg0) => f.debug_tuple("Pow").field(arg0).finish(),
Self::Mul(arg0) => f.debug_tuple("Mul").field(arg0).finish(),
Self::Add(arg0) => f.debug_tuple("Add").field(arg0).finish(),
Self::Literal(arg0) => f.debug_tuple("Literal").field(arg0).finish(),
Self::Alternative(arg0) => f.debug_tuple("Alternative").field(arg0).finish(),
Self::Transformer(arg0) => f.debug_tuple("Transformer").field(arg0).finish(),
}
}
}
pub trait FilterFn: Fn(&Match) -> bool + DynClone + Send + Sync {}
dyn_clone::clone_trait_object!(FilterFn);
impl<T: Clone + Send + Sync + Fn(&Match) -> bool> FilterFn for T {}
pub trait FilterSingleFn: Fn(AtomView<'_>) -> bool + DynClone + Send + Sync {}
dyn_clone::clone_trait_object!(FilterSingleFn);
impl<T: Clone + Send + Sync + Fn(AtomView<'_>) -> bool> FilterSingleFn for T {}
pub trait CmpFn: Fn(&Match, &Match) -> bool + DynClone + Send + Sync {}
dyn_clone::clone_trait_object!(CmpFn);
impl<T: Clone + Send + Sync + Fn(&Match, &Match) -> bool> CmpFn for T {}
pub trait MatchStackFn: Fn(&MatchStack) -> ConditionResult + DynClone + Send + Sync {}
dyn_clone::clone_trait_object!(MatchStackFn);
impl<T: Clone + Send + Sync + Fn(&MatchStack) -> ConditionResult> MatchStackFn for T {}
pub enum WildcardRestriction {
Length(usize, Option<usize>), IsAtomType(AtomType),
HasTag(String),
IsLiteralWildcard(Symbol),
Filter(Box<dyn FilterFn>),
Cmp(Symbol, Box<dyn CmpFn>),
NotGreedy,
}
impl WildcardRestriction {
pub fn filter(f: impl FilterFn + 'static) -> Self {
WildcardRestriction::Filter(Box::new(f))
}
pub fn cmp(s: Symbol, f: impl CmpFn + 'static) -> Self {
WildcardRestriction::Cmp(s, Box::new(f))
}
}
impl std::fmt::Display for WildcardRestriction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
WildcardRestriction::Length(min, Some(max)) => write!(f, "length={min}-{max}"),
WildcardRestriction::Length(min, None) => write!(f, "length > {min}"),
WildcardRestriction::IsAtomType(t) => write!(f, "type = {t}"),
WildcardRestriction::IsLiteralWildcard(s) => write!(f, "= {s}"),
WildcardRestriction::Filter(_) => write!(f, "filter"),
WildcardRestriction::Cmp(s, _) => write!(f, "cmp with {s}"),
WildcardRestriction::NotGreedy => write!(f, "not greedy"),
WildcardRestriction::HasTag(tag) => write!(f, "has tag {tag}"),
}
}
}
pub type WildcardAndRestriction = (Symbol, WildcardRestriction);
pub enum PatternRestriction {
Wildcard(WildcardAndRestriction),
MatchStack(Box<dyn MatchStackFn>),
}
impl Condition<PatternRestriction> {
pub fn match_stack(f: impl MatchStackFn + 'static) -> Self {
PatternRestriction::MatchStack(Box::new(f)).into()
}
}
impl std::fmt::Display for PatternRestriction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PatternRestriction::Wildcard((s, r)) => write!(f, "{s}: {r}"),
PatternRestriction::MatchStack(_) => write!(f, "match_function"),
}
}
}
impl Clone for PatternRestriction {
fn clone(&self) -> Self {
match self {
PatternRestriction::Wildcard(w) => PatternRestriction::Wildcard(w.clone()),
PatternRestriction::MatchStack(f) => {
PatternRestriction::MatchStack(dyn_clone::clone_box(f))
}
}
}
}
impl From<WildcardAndRestriction> for PatternRestriction {
fn from(value: WildcardAndRestriction) -> Self {
PatternRestriction::Wildcard(value)
}
}
impl From<WildcardAndRestriction> for Condition<PatternRestriction> {
fn from(value: WildcardAndRestriction) -> Self {
PatternRestriction::Wildcard(value).into()
}
}
impl From<PatternRestriction> for Option<Condition<PatternRestriction>> {
fn from(value: PatternRestriction) -> Self {
Some(Condition::from(value))
}
}
impl Symbol {
pub fn restrict(&self, restriction: WildcardRestriction) -> Condition<PatternRestriction> {
Condition::from((*self, restriction))
}
pub fn filter_tag(&self, tag: String) -> Condition<PatternRestriction> {
if tag.contains("::") {
self.restrict(WildcardRestriction::HasTag(tag))
} else {
panic!("Tag {} must contain namespace", tag);
}
}
pub fn filter(
&self,
f: impl FilterSingleFn + 'static + Clone,
) -> Condition<PatternRestriction> {
if self.get_wildcard_level() != 1 {
panic!(
"filter can only be used on single wildcards (with one underscore), but {} has level {}",
self,
self.get_wildcard_level()
);
}
self.restrict(WildcardRestriction::filter(move |m| match m {
Match::Single(a) => f(*a),
_ => unreachable!("Expected single match for filter, but got {m:?}"),
}))
}
pub fn filter_match(&self, f: impl FilterFn + 'static) -> Condition<PatternRestriction> {
self.restrict(WildcardRestriction::filter(f))
}
pub fn filter_cmp(&self, s: Symbol, f: impl CmpFn + 'static) -> Condition<PatternRestriction> {
self.restrict(WildcardRestriction::Cmp(s, Box::new(f)))
}
}
static DEFAULT_PATTERN_CONDITION: Condition<PatternRestriction> = Condition::True;
#[derive(Clone, Debug, Default)]
pub enum Condition<T> {
And(Box<(Condition<T>, Condition<T>)>),
Or(Box<(Condition<T>, Condition<T>)>),
Not(Box<Condition<T>>),
Yield(T),
#[default]
True,
False,
}
impl<T: std::fmt::Display> std::fmt::Display for Condition<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Condition::And(a) => write!(f, "({}) & ({})", a.0, a.1),
Condition::Or(o) => write!(f, "{} | {}", o.0, o.1),
Condition::Not(n) => write!(f, "!({n})"),
Condition::True => write!(f, "True"),
Condition::False => write!(f, "False"),
Condition::Yield(t) => write!(f, "{t}"),
}
}
}
pub trait Evaluate {
type State<'a>;
fn evaluate(&self, state: &Self::State<'_>) -> Result<ConditionResult, String>;
}
impl<T: Evaluate> Evaluate for Condition<T> {
type State<'a> = T::State<'a>;
fn evaluate(&self, state: &T::State<'_>) -> Result<ConditionResult, String> {
Ok(match self {
Condition::And(a) => a.0.evaluate(state)? & a.1.evaluate(state)?,
Condition::Or(o) => o.0.evaluate(state)? | o.1.evaluate(state)?,
Condition::Not(n) => !n.evaluate(state)?,
Condition::True => ConditionResult::True,
Condition::False => ConditionResult::False,
Condition::Yield(t) => t.evaluate(state)?,
})
}
}
impl<T> From<T> for Condition<T> {
fn from(value: T) -> Self {
Condition::Yield(value)
}
}
impl<T, R: Into<Condition<T>>> std::ops::BitOr<R> for Condition<T> {
type Output = Condition<T>;
fn bitor(self, rhs: R) -> Self::Output {
Condition::Or(Box::new((self, rhs.into())))
}
}
impl<T, R: Into<Condition<T>>> std::ops::BitAnd<R> for Condition<T> {
type Output = Condition<T>;
fn bitand(self, rhs: R) -> Self::Output {
Condition::And(Box::new((self, rhs.into())))
}
}
impl<T> std::ops::Not for Condition<T> {
type Output = Condition<T>;
fn not(self) -> Self::Output {
Condition::Not(Box::new(self))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConditionResult {
True,
False,
Inconclusive,
}
impl std::ops::BitOr<ConditionResult> for ConditionResult {
type Output = ConditionResult;
fn bitor(self, rhs: ConditionResult) -> Self::Output {
match (self, rhs) {
(ConditionResult::True, _) => ConditionResult::True,
(_, ConditionResult::True) => ConditionResult::True,
(ConditionResult::False, ConditionResult::False) => ConditionResult::False,
_ => ConditionResult::Inconclusive,
}
}
}
impl std::ops::BitAnd<ConditionResult> for ConditionResult {
type Output = ConditionResult;
fn bitand(self, rhs: ConditionResult) -> Self::Output {
match (self, rhs) {
(ConditionResult::False, _) => ConditionResult::False,
(_, ConditionResult::False) => ConditionResult::False,
(ConditionResult::True, ConditionResult::True) => ConditionResult::True,
_ => ConditionResult::Inconclusive,
}
}
}
impl std::ops::Not for ConditionResult {
type Output = ConditionResult;
fn not(self) -> Self::Output {
match self {
ConditionResult::True => ConditionResult::False,
ConditionResult::False => ConditionResult::True,
ConditionResult::Inconclusive => ConditionResult::Inconclusive,
}
}
}
impl From<bool> for ConditionResult {
fn from(value: bool) -> Self {
if value {
ConditionResult::True
} else {
ConditionResult::False
}
}
}
impl ConditionResult {
pub fn is_true(&self) -> bool {
matches!(self, ConditionResult::True)
}
pub fn is_false(&self) -> bool {
matches!(self, ConditionResult::False)
}
pub fn is_inconclusive(&self) -> bool {
matches!(self, ConditionResult::Inconclusive)
}
}
#[derive(Clone, Debug)]
pub enum Relation {
Eq(Pattern, Pattern),
Ne(Pattern, Pattern),
Gt(Pattern, Pattern),
Ge(Pattern, Pattern),
Lt(Pattern, Pattern),
Le(Pattern, Pattern),
Contains(Pattern, Pattern),
IsType(Pattern, AtomType),
Matches(
Pattern,
Pattern,
Condition<PatternRestriction>,
MatchSettings,
),
}
impl std::fmt::Display for Relation {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Relation::Eq(a, b) => write!(f, "{a} == {b}"),
Relation::Ne(a, b) => write!(f, "{a} != {b}"),
Relation::Gt(a, b) => write!(f, "{a} > {b}"),
Relation::Ge(a, b) => write!(f, "{a} >= {b}"),
Relation::Lt(a, b) => write!(f, "{a} < {b}"),
Relation::Le(a, b) => write!(f, "{a} <= {b}"),
Relation::Contains(a, b) => write!(f, "{a} contains {b}"),
Relation::IsType(a, b) => write!(f, "{a} is type {b:?}"),
Relation::Matches(a, b, _, _) => write!(f, "{a} matches {b}"),
}
}
}
impl Evaluate for Relation {
type State<'a> = Option<AtomView<'a>>;
fn evaluate(&self, state: &Option<AtomView>) -> Result<ConditionResult, String> {
Workspace::get_local().with(|ws| {
let mut out1 = ws.new_atom();
let mut out2 = ws.new_atom();
let m = MatchStack::new();
let pat = state.map(|x| x.to_pattern());
Ok(match self {
Relation::Eq(a, b)
| Relation::Ne(a, b)
| Relation::Gt(a, b)
| Relation::Ge(a, b)
| Relation::Lt(a, b)
| Relation::Le(a, b)
| Relation::Contains(a, b) => {
a.replace_wildcards_with_matches_impl(ws, &mut out1, &m, true, pat.as_ref())
.map_err(|e| match e {
TransformerError::Interrupt => "Interrupted by user".into(),
TransformerError::ValueError(v) => v,
})?;
b.replace_wildcards_with_matches_impl(ws, &mut out2, &m, true, pat.as_ref())
.map_err(|e| match e {
TransformerError::Interrupt => "Interrupted by user".into(),
TransformerError::ValueError(v) => v,
})?;
match self {
Relation::Eq(_, _) => out1 == out2,
Relation::Ne(_, _) => out1 != out2,
Relation::Gt(_, _) => out1.as_view() > out2.as_view(),
Relation::Ge(_, _) => out1.as_view() >= out2.as_view(),
Relation::Lt(_, _) => out1.as_view() < out2.as_view(),
Relation::Le(_, _) => out1.as_view() <= out2.as_view(),
Relation::Contains(_, _) => out1.contains(out2.as_view()),
_ => unreachable!(),
}
}
Relation::Matches(a, pattern, cond, settings) => {
a.replace_wildcards_with_matches_impl(ws, &mut out1, &m, true, pat.as_ref())
.map_err(|e| match e {
TransformerError::Interrupt => "Interrupted by user".into(),
TransformerError::ValueError(v) => v,
})?;
out1.pattern_match(pattern, Some(cond), Some(settings))
.next()
.is_some()
}
Relation::IsType(a, b) => {
a.replace_wildcards_with_matches_impl(ws, &mut out1, &m, true, pat.as_ref())
.map_err(|e| match e {
TransformerError::Interrupt => "Interrupted by user".into(),
TransformerError::ValueError(v) => v,
})?;
match out1.as_ref() {
Atom::Var(_) => *b == AtomType::Var,
Atom::Fun(_) => *b == AtomType::Fun,
Atom::Num(_) => *b == AtomType::Num,
Atom::Add(_) => *b == AtomType::Add,
Atom::Mul(_) => *b == AtomType::Mul,
Atom::Pow(_) => *b == AtomType::Pow,
Atom::Zero => *b == AtomType::Num,
}
}
}
.into())
})
}
}
impl Evaluate for Condition<PatternRestriction> {
type State<'a> = MatchStack<'a>;
fn evaluate(&self, state: &MatchStack) -> Result<ConditionResult, String> {
Ok(match self {
Condition::And(a) => a.0.evaluate(state)? & a.1.evaluate(state)?,
Condition::Or(o) => o.0.evaluate(state)? | o.1.evaluate(state)?,
Condition::Not(n) => !n.evaluate(state)?,
Condition::True => ConditionResult::True,
Condition::False => ConditionResult::False,
Condition::Yield(t) => match t {
PatternRestriction::Wildcard((v, r)) => {
if let Some((_, value)) = state.stack.iter().find(|(k, _)| k == v) {
match r {
WildcardRestriction::IsAtomType(t) => match value {
Match::Single(AtomView::Num(_)) => *t == AtomType::Num,
Match::Single(AtomView::Var(_)) => *t == AtomType::Var,
Match::Single(AtomView::Add(_)) => *t == AtomType::Add,
Match::Single(AtomView::Mul(_)) => *t == AtomType::Mul,
Match::Single(AtomView::Pow(_)) => *t == AtomType::Pow,
Match::Single(AtomView::Fun(_)) => *t == AtomType::Fun,
_ => false,
},
WildcardRestriction::IsLiteralWildcard(wc) => match value {
Match::Single(AtomView::Var(v)) => wc == &v.get_symbol(),
Match::FunctionName(s) => wc == s,
_ => false,
},
WildcardRestriction::Length(min, max) => match value {
Match::Single(_) | Match::FunctionName(_) => {
*min <= 1 && max.map(|m| m >= 1).unwrap_or(true)
}
Match::Multiple(_, slice) => {
*min <= slice.len()
&& max.map(|m| m >= slice.len()).unwrap_or(true)
}
},
WildcardRestriction::Filter(f) => f(value),
WildcardRestriction::Cmp(v2, f) => {
if let Some((_, value2)) = state.stack.iter().find(|(k, _)| k == v2)
{
f(value, value2)
} else {
return Ok(ConditionResult::Inconclusive);
}
}
WildcardRestriction::NotGreedy => true,
WildcardRestriction::HasTag(tag) => match value {
Match::Single(AtomView::Var(v)) => v.get_symbol().has_tag(tag),
Match::Single(AtomView::Fun(f)) => f.get_symbol().has_tag(tag),
Match::FunctionName(s) => s.has_tag(tag),
_ => false,
},
}
.into()
} else {
ConditionResult::Inconclusive
}
}
PatternRestriction::MatchStack(mf) => mf(state),
},
})
}
}
impl Condition<PatternRestriction> {
fn check_possible(&self, var: Symbol, value: &Match, stack: &MatchStack) -> ConditionResult {
match self {
Condition::And(a) => {
a.0.check_possible(var, value, stack) & a.1.check_possible(var, value, stack)
}
Condition::Or(o) => {
o.0.check_possible(var, value, stack) | o.1.check_possible(var, value, stack)
}
Condition::Not(n) => !n.check_possible(var, value, stack),
Condition::True => ConditionResult::True,
Condition::False => ConditionResult::False,
Condition::Yield(restriction) => {
let (v, r) = match restriction {
PatternRestriction::Wildcard((v, r)) => (v, r),
PatternRestriction::MatchStack(mf) => {
return mf(stack);
}
};
if *v != var {
match r {
WildcardRestriction::Cmp(v, _) if *v == var => {}
_ => {
return ConditionResult::Inconclusive;
}
}
}
match r {
WildcardRestriction::IsAtomType(t) => {
let is_type = match t {
AtomType::Num => matches!(value, Match::Single(AtomView::Num(_))),
AtomType::Var => matches!(value, Match::Single(AtomView::Var(_))),
AtomType::Add => matches!(
value,
Match::Single(AtomView::Add(_))
| Match::Multiple(SliceType::Add, _)
),
AtomType::Mul => matches!(
value,
Match::Single(AtomView::Mul(_))
| Match::Multiple(SliceType::Mul, _)
),
AtomType::Pow => matches!(
value,
Match::Single(AtomView::Pow(_))
| Match::Multiple(SliceType::Pow, _)
),
AtomType::Fun => matches!(value, Match::Single(AtomView::Fun(_))),
};
(is_type == matches!(r, WildcardRestriction::IsAtomType(_))).into()
}
WildcardRestriction::IsLiteralWildcard(wc) => match value {
Match::Single(AtomView::Var(v)) => (wc == &v.get_symbol()).into(),
Match::FunctionName(s) => (wc == s).into(),
_ => false.into(),
},
WildcardRestriction::Length(min, max) => match &value {
Match::Single(_) | Match::FunctionName(_) => {
(*min <= 1 && max.map(|m| m >= 1).unwrap_or(true)).into()
}
Match::Multiple(_, slice) => (*min <= slice.len()
&& max.map(|m| m >= slice.len()).unwrap_or(true))
.into(),
},
WildcardRestriction::Filter(f) => f(value).into(),
WildcardRestriction::Cmp(v2, f) => {
if *v == var {
if let Some((_, value2)) = stack.stack.iter().find(|(k, _)| k == v2) {
f(value, value2).into()
} else {
ConditionResult::Inconclusive
}
} else if let Some((_, value2)) = stack.stack.iter().find(|(k, _)| k == v) {
f(value2, value).into()
} else {
ConditionResult::Inconclusive
}
}
WildcardRestriction::NotGreedy => true.into(),
WildcardRestriction::HasTag(tag) => match value {
Match::Single(AtomView::Var(v)) => v.get_symbol().has_tag(tag).into(),
Match::Single(AtomView::Fun(f)) => f.get_symbol().has_tag(tag).into(),
Match::FunctionName(s) => s.has_tag(tag).into(),
_ => false.into(),
},
}
}
}
}
fn get_range_hint(&self, var: Symbol) -> (Option<usize>, Option<usize>) {
match self {
Condition::And(a) => {
let (min1, max1) = a.0.get_range_hint(var);
let (min2, max2) = a.1.get_range_hint(var);
(
match (min1, min2) {
(None, None) => None,
(None, Some(m)) => Some(m),
(Some(m), None) => Some(m),
(Some(m1), Some(m2)) => Some(m1.max(m2)),
},
match (max1, max2) {
(None, None) => None,
(None, Some(m)) => Some(m),
(Some(m), None) => Some(m),
(Some(m1), Some(m2)) => Some(m1.min(m2)),
},
)
}
Condition::Or(o) => {
let (min1, max1) = o.0.get_range_hint(var);
let (min2, max2) = o.1.get_range_hint(var);
(
if let (Some(m1), Some(m2)) = (min1, min2) {
Some(m1.min(m2))
} else {
None
},
if let (Some(m1), Some(m2)) = (max1, max2) {
Some(m1.max(m2))
} else {
None
},
)
}
Condition::Not(_) => {
(None, None)
}
Condition::True | Condition::False => (None, None),
Condition::Yield(restriction) => {
let (v, r) = match restriction {
PatternRestriction::Wildcard((v, r)) => (v, r),
PatternRestriction::MatchStack(_) => {
return (None, None);
}
};
if *v != var {
return (None, None);
}
match r {
WildcardRestriction::Length(min, max) => (Some(*min), *max),
WildcardRestriction::IsAtomType(
AtomType::Var | AtomType::Num | AtomType::Fun,
)
| WildcardRestriction::IsLiteralWildcard(_) => (Some(1), Some(1)),
_ => (None, None),
}
}
}
}
}
impl Clone for WildcardRestriction {
fn clone(&self) -> Self {
match self {
Self::Length(min, max) => Self::Length(*min, *max),
Self::IsAtomType(t) => Self::IsAtomType(*t),
Self::IsLiteralWildcard(w) => Self::IsLiteralWildcard(*w),
Self::Filter(f) => Self::Filter(dyn_clone::clone_box(f)),
Self::Cmp(i, f) => Self::Cmp(*i, dyn_clone::clone_box(f)),
Self::NotGreedy => Self::NotGreedy,
Self::HasTag(tag) => Self::HasTag(tag.clone()),
}
}
}
impl std::fmt::Debug for WildcardRestriction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Length(arg0, arg1) => f.debug_tuple("Length").field(arg0).field(arg1).finish(),
Self::IsAtomType(t) => write!(f, "Is{t:?}"),
Self::IsLiteralWildcard(arg0) => {
f.debug_tuple("IsLiteralWildcard").field(arg0).finish()
}
Self::Filter(_) => f.debug_tuple("Filter").finish(),
Self::Cmp(arg0, _) => f.debug_tuple("Cmp").field(arg0).finish(),
Self::NotGreedy => write!(f, "NotGreedy"),
Self::HasTag(tag) => f.debug_tuple("HasTag").field(tag).finish(),
}
}
}
impl std::fmt::Debug for PatternRestriction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PatternRestriction::Wildcard(arg0) => f.debug_tuple("Wildcard").field(arg0).finish(),
PatternRestriction::MatchStack(_) => f.debug_tuple("Match").finish(),
}
}
}
#[derive(Clone, PartialEq, Eq, Hash)]
pub enum Match<'a> {
Single(AtomView<'a>),
Multiple(SliceType, Vec<AtomView<'a>>),
FunctionName(Symbol),
}
impl std::fmt::Display for Match<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Single(a) => a.fmt(f),
Self::Multiple(t, list) => match t {
SliceType::Add | SliceType::Mul | SliceType::Arg | SliceType::Pow => {
f.write_str("(")?;
for (i, a) in list.iter().enumerate() {
if i > 0 {
match t {
SliceType::Add => {
f.write_str("+")?;
}
SliceType::Mul => {
f.write_str("*")?;
}
SliceType::Arg => {
f.write_str(",")?;
}
SliceType::Pow => {
f.write_str("^")?;
}
_ => unreachable!(),
}
}
a.fmt(f)?;
}
f.write_str(")")
}
SliceType::One => list[0].fmt(f),
SliceType::Empty => f.write_str("()"),
},
Self::FunctionName(name) => name.fmt(f),
}
}
}
impl std::fmt::Debug for Match<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Single(a) => f.debug_tuple("").field(a).finish(),
Self::Multiple(t, list) => f.debug_tuple("").field(t).field(list).finish(),
Self::FunctionName(name) => f.debug_tuple("Fn").field(name).finish(),
}
}
}
impl<'a> Match<'a> {
pub fn to_atom(&self) -> Atom {
let mut out = Atom::default();
self.to_atom_into(&mut out);
out
}
pub fn as_single(&self) -> Option<AtomView<'a>> {
match self {
Self::Single(v) => Some(*v),
_ => None,
}
}
pub fn matches(&self, other: &Match<'a>) -> bool {
match (self, other) {
(Self::Single(a), Self::Single(b)) => a == b,
(Self::Multiple(t1, list1), Self::Multiple(t2, list2)) => t1 == t2 && list1 == list2,
(Self::FunctionName(n1), Self::FunctionName(n2)) => n1 == n2,
(Self::Single(a), Self::FunctionName(n)) => a == n,
(Self::FunctionName(n), Self::Single(b)) => b == n,
_ => false,
}
}
pub fn to_atom_into(&self, out: &mut Atom) {
match self {
Self::Single(v) => {
out.set_from_view(v);
}
Self::Multiple(t, args) => match t {
SliceType::Add => {
let add = out.to_add();
for arg in args {
add.extend(*arg);
}
add.set_normalized(true);
}
SliceType::Mul => {
let mut has_coefficient = false;
let mul = out.to_mul();
for arg in args {
has_coefficient |= matches!(arg, AtomView::Num(_));
mul.extend(*arg);
}
mul.set_has_coefficient(has_coefficient);
mul.set_normalized(true);
}
SliceType::Arg => {
let fun = out.to_fun(Symbol::ARG);
fun.add_args(args);
fun.set_normalized(true);
}
SliceType::Pow => {
let p = out.to_pow(args[0], args[1]);
p.set_normalized(true);
}
SliceType::One => {
out.set_from_view(&args[0]);
}
SliceType::Empty => {
let f = out.to_fun(Symbol::ARG);
f.set_normalized(true);
}
},
Self::FunctionName(n) => {
out.to_var(*n);
}
}
}
}
#[derive(Debug, Clone)]
pub struct MatchSettings {
pub(crate) non_greedy_wildcards: Vec<Symbol>,
pub(crate) level_range: (usize, Option<usize>),
pub(crate) level_is_tree_depth: bool,
pub(crate) partial: bool,
pub(crate) allow_new_wildcards_on_rhs: bool,
pub(crate) rhs_cache_size: usize,
}
static DEFAULT_MATCH_SETTINGS: MatchSettings = MatchSettings::new();
impl MatchSettings {
pub const fn new() -> Self {
Self {
non_greedy_wildcards: Vec::new(),
level_range: (0, None),
level_is_tree_depth: false,
partial: true,
allow_new_wildcards_on_rhs: false,
rhs_cache_size: 0,
}
}
pub fn cached() -> Self {
Self {
non_greedy_wildcards: Vec::new(),
level_range: (0, None),
level_is_tree_depth: false,
partial: true,
allow_new_wildcards_on_rhs: false,
rhs_cache_size: 100,
}
}
pub fn non_greedy_wildcards(mut self, non_greedy_wildcards: Vec<Symbol>) -> Self {
self.non_greedy_wildcards = non_greedy_wildcards;
self
}
pub fn level_range(mut self, level_range: (usize, Option<usize>)) -> Self {
self.level_range = level_range;
self
}
pub fn level_is_tree_depth(mut self, level_is_tree_depth: bool) -> Self {
self.level_is_tree_depth = level_is_tree_depth;
self
}
pub fn partial(mut self, partial: bool) -> Self {
self.partial = partial;
self
}
pub fn allow_new_wildcards_on_rhs(mut self, allow_new_wildcards_on_rhs: bool) -> Self {
self.allow_new_wildcards_on_rhs = allow_new_wildcards_on_rhs;
self
}
pub fn rhs_cache_size(mut self, rhs_cache_size: usize) -> Self {
self.rhs_cache_size = rhs_cache_size;
self
}
}
impl Default for MatchSettings {
fn default() -> Self {
MatchSettings::new()
}
}
#[derive(Debug, Clone)]
pub struct MatchStack<'a> {
stack: Vec<(Symbol, Match<'a>)>,
}
impl<'a> From<Vec<(Symbol, Match<'a>)>> for MatchStack<'a> {
fn from(value: Vec<(Symbol, Match<'a>)>) -> Self {
MatchStack { stack: value }
}
}
impl Default for MatchStack<'_> {
fn default() -> Self {
Self::new()
}
}
impl<'a> MatchStack<'a> {
pub fn new() -> Self {
MatchStack { stack: Vec::new() }
}
pub fn get_atom(&self, key: Symbol) -> Option<AtomView<'a>> {
self.get(key)?.as_single()
}
pub fn get(&self, key: Symbol) -> Option<&Match<'a>> {
if key.get_wildcard_level() == 0 {
panic!(
"Cannot get match for a non-wildcard symbol: {}",
key.get_name()
);
}
for (rk, rv) in self.stack.iter() {
if rk == &key {
return Some(rv);
}
}
None
}
pub fn get_matches(&self) -> &[(Symbol, Match<'a>)] {
&self.stack
}
pub fn into_matches(self) -> Vec<(Symbol, Match<'a>)> {
self.stack
}
}
pub struct WrappedMatchStack<'a, 'b> {
stack: MatchStack<'a>,
conditions: &'b Condition<PatternRestriction>,
settings: &'b MatchSettings,
}
impl std::fmt::Display for MatchStack<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("[")?;
for (i, (k, v)) in self.stack.iter().enumerate() {
if i > 0 {
f.write_str(", ")?;
}
f.write_fmt(format_args!("{k}: {v}"))?;
}
f.write_str("]")
}
}
impl<'a, 'b> IntoIterator for &'b MatchStack<'a> {
type Item = &'b (Symbol, Match<'a>);
type IntoIter = std::slice::Iter<'b, (Symbol, Match<'a>)>;
fn into_iter(self) -> Self::IntoIter {
self.stack.iter()
}
}
impl std::fmt::Display for WrappedMatchStack<'_, '_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.stack.fmt(f)
}
}
impl std::fmt::Debug for WrappedMatchStack<'_, '_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MatchStack")
.field("stack", &self.stack)
.finish()
}
}
impl<'a, 'b> WrappedMatchStack<'a, 'b> {
pub fn from_replacement(replacement: BorrowedReplacement<'b>) -> WrappedMatchStack<'a, 'b> {
WrappedMatchStack {
stack: MatchStack::new(),
conditions: replacement.conditions.unwrap_or(&DEFAULT_PATTERN_CONDITION),
settings: replacement.settings.unwrap_or(&DEFAULT_MATCH_SETTINGS),
}
}
pub fn reset(&mut self) {
self.stack.stack.clear();
}
pub fn new(
conditions: &'b Condition<PatternRestriction>,
settings: &'b MatchSettings,
) -> WrappedMatchStack<'a, 'b> {
WrappedMatchStack {
stack: MatchStack::new(),
conditions,
settings,
}
}
pub fn insert(&mut self, key: Symbol, value: Match<'a>) -> Result<usize, MatchError> {
if key.has_attributes() {
match &value {
Match::Single(s) => {
if !s.has_attributes_of(key) {
return Err(MatchError::StructurallyImpossible);
}
}
Match::Multiple(_, list) => {
for s in list {
if !s.has_attributes_of(key) {
return Err(MatchError::StructurallyImpossible);
}
}
}
Match::FunctionName(n) => {
if !n.has_attributes_of(key) {
return Err(MatchError::StructurallyImpossible);
}
}
}
}
for (rk, rv) in self.stack.stack.iter() {
if rk == &key {
if rv.matches(&value) {
return Ok(self.stack.stack.len());
} else {
return Err(MatchError::ImpossibleDueToConstraints);
}
}
}
self.stack.stack.push((key, value));
if self
.conditions
.check_possible(key, &self.stack.stack.last().unwrap().1, &self.stack)
== ConditionResult::False
{
self.stack.stack.pop();
Err(MatchError::ImpossibleDueToConstraints)
} else {
Ok(self.stack.stack.len() - 1)
}
}
pub fn get_match_stack(&self) -> &MatchStack<'a> {
&self.stack
}
pub fn get_matches(&self) -> &[(Symbol, Match<'a>)] {
&self.stack.stack
}
#[inline]
pub fn len(&self) -> usize {
self.stack.stack.len()
}
#[inline]
pub fn truncate(&mut self, len: usize) {
self.stack.stack.truncate(len)
}
pub fn get_range(&self, identifier: Symbol) -> (usize, Option<usize>) {
self.get_range_impl(identifier, false)
}
pub fn get_range_with_optional(
&self,
identifier: Symbol,
optional: bool,
) -> (usize, Option<usize>) {
self.get_range_impl(identifier, optional)
}
fn get_range_impl(&self, identifier: Symbol, optional: bool) -> (usize, Option<usize>) {
if identifier.get_wildcard_level() == 0 {
return (1, Some(1));
}
for (rk, rv) in self.stack.stack.iter() {
if *rk == identifier {
return match rv {
Match::Single(a) => (
if optional && (*a == ZERO.as_view() || *a == ONE.as_view()) {
0
} else {
1
},
Some(1),
),
Match::Multiple(slice_type, slice) => {
match slice_type {
SliceType::Empty => (0, Some(0)),
SliceType::Arg => (slice.len(), Some(slice.len())),
_ => {
(1, Some(slice.len()))
}
}
}
Match::FunctionName(_) => (1, Some(1)),
};
}
}
let (minimal, maximal) = self.conditions.get_range_hint(identifier);
let range = match identifier.get_wildcard_level() {
1 => (minimal.unwrap_or(1), Some(maximal.unwrap_or(1))), 2 => (minimal.unwrap_or(1), maximal), _ => (minimal.unwrap_or(0), maximal), };
if optional { (0, range.1) } else { range }
}
}
#[derive(Debug)]
struct WildcardIter {
initialized: bool,
name: Symbol,
indices: Vec<u32>,
size_target: u32,
min_size: u32,
max_size: u32,
greedy: bool,
resume_point: Option<Checkpoint>,
}
#[derive(Clone, Copy, Debug)]
struct Checkpoint(usize);
impl Checkpoint {
fn new(stack_len: usize) -> Self {
Self(stack_len)
}
fn restore(self, match_stack: &mut WrappedMatchStack<'_, '_>) {
debug_assert!(
match_stack.len() >= self.0,
"match stack was cleaned past its frame"
);
match_stack.truncate(self.0);
}
}
#[derive(Debug)]
enum PatternIter<'a, 'b> {
Wildcard(WildcardIter),
Node(Option<usize>, Box<AtomMatchIterator<'a, 'b>>),
}
#[derive(Debug)]
struct FnMatchIterator<'a, 'b> {
initialized: bool,
name: Symbol,
args: SubSliceIterator<'a, 'b>,
target: AtomView<'a>,
}
impl<'a, 'b> FnMatchIterator<'a, 'b> {
fn new(name: Symbol, args: &'b [Pattern]) -> Self {
FnMatchIterator {
initialized: false,
name,
args: SubSliceIterator::new(args, SliceType::Arg),
target: AtomView::ZERO,
}
}
fn set_target(&mut self, target: AtomView<'a>) {
self.initialized = false;
self.target = target;
}
fn next(&mut self, match_stack: &mut WrappedMatchStack<'a, 'b>) -> Result<(), MatchError> {
if !self.initialized {
self.initialized = true;
match self.target {
AtomView::Fun(f) => {
let target_name = f.get_symbol();
if self.name.get_wildcard_level() > 0 {
match_stack.insert(self.name, Match::FunctionName(target_name))?;
} else if target_name != self.name {
return Err(MatchError::StructurallyImpossible);
}
self.args.set_list_target(
self.target,
match_stack,
true,
!target_name.is_antisymmetric() && !target_name.is_symmetric(),
target_name.is_cyclesymmetric(),
);
}
AtomView::Var(v) => {
if !self.args.pattern.iter().all(Pattern::is_optional_wildcard) {
return Err(MatchError::StructurallyImpossible);
}
let target_name = v.get_symbol();
if self.name.get_wildcard_level() > 0 {
match_stack.insert(self.name, Match::FunctionName(target_name))?;
} else if target_name != self.name {
return Err(MatchError::StructurallyImpossible);
}
self.args
.set_empty_list_target(match_stack, true, true, false);
}
_ => return Err(MatchError::StructurallyImpossible),
}
}
self.args.next(match_stack).map(|_| ())
}
}
#[derive(Debug)]
struct AlternativeIter<'a, 'b> {
variants: Vec<AtomMatchIterator<'a, 'b>>,
variant_index: usize,
variant_initialized: bool,
target: AtomView<'a>,
structural_impossible: bool,
}
impl<'a, 'b> AlternativeIter<'a, 'b> {
fn new(alternatives: &'b [Pattern]) -> Self {
AlternativeIter {
variants: alternatives
.iter()
.map(|pattern| AtomMatchIterator::new(pattern))
.collect(),
variant_index: 0,
variant_initialized: false,
target: AtomView::ZERO,
structural_impossible: true,
}
}
fn set_target(&mut self, target: AtomView<'a>) {
self.variant_index = 0;
self.variant_initialized = false;
self.target = target;
self.structural_impossible = true;
}
fn next(&mut self, match_stack: &mut WrappedMatchStack<'a, 'b>) -> Result<(), MatchError> {
while self.variant_index < self.variants.len() {
let variant_index = self.variant_index;
if !self.variant_initialized {
self.variants[variant_index].set_new_target_complete(self.target, match_stack);
self.variant_initialized = true;
}
match self.variants[variant_index].next_result(match_stack) {
Ok(_) => {
self.structural_impossible = false;
return Ok(());
}
Err(e) => {
if !matches!(e, MatchError::StructurallyImpossible) {
self.structural_impossible = false;
}
}
}
self.variant_index += 1;
self.variant_initialized = false;
}
if self.structural_impossible {
Err(MatchError::StructurallyImpossible)
} else {
Err(MatchError::NoMoreMatches)
}
}
}
#[derive(Debug)]
struct LiteralAtomIter<'b> {
literal: &'b Atom,
try_match_atom: bool,
}
impl<'b> LiteralAtomIter<'b> {
fn new(literal: &'b Atom) -> Self {
LiteralAtomIter {
literal,
try_match_atom: true,
}
}
fn set_target(&mut self) {
self.try_match_atom = true;
}
fn next<'a>(&mut self, target: AtomView<'a>) -> Result<(), MatchError> {
if self.try_match_atom {
self.try_match_atom = false;
if self.literal.as_view() == target {
return Ok(());
}
}
Err(MatchError::StructurallyImpossible)
}
}
#[derive(Debug)]
struct SliceAtomIter<'a, 'b> {
pattern: &'b Pattern,
target: AtomView<'a>,
iter: SubSliceIterator<'a, 'b>,
single_atom_fallback: Option<AtomView<'a>>,
used_flags: Vec<bool>,
}
impl<'a, 'b> SliceAtomIter<'a, 'b> {
fn new(pattern: &'b Pattern, pat_list: &'b [Pattern], slice_type: SliceType) -> Self {
SliceAtomIter {
pattern,
target: AtomView::ZERO,
iter: SubSliceIterator::new(pat_list, slice_type),
single_atom_fallback: None,
used_flags: Vec::new(),
}
}
fn set_target(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
force_complete: bool,
) {
self.target = target;
self.single_atom_fallback = self
.iter
.can_match_list_atom_as_single_with_optional(self.target)
.then_some(self.target);
self.iter.set_target(
self.target,
match_stack,
true,
matches!(self.pattern, Pattern::Wildcard(..) | Pattern::Literal(_)),
);
if force_complete {
self.iter.complete = true;
}
}
fn next(
&mut self,
match_stack: &mut WrappedMatchStack<'a, 'b>,
) -> Result<Option<&[bool]>, MatchError> {
let primary_error = match self.iter.next(match_stack) {
Ok(used_flags) => {
self.used_flags.clear();
self.used_flags.extend_from_slice(used_flags);
return Ok(Some(&self.used_flags));
}
Err(e) => e,
};
if let Some(target) = self.single_atom_fallback.take() {
let complete = self.iter.complete;
self.iter.set_single_atom_target(target, complete);
match self.iter.next(match_stack) {
Ok(used_flags) => {
self.used_flags.clear();
self.used_flags.extend_from_slice(used_flags);
Ok(Some(&self.used_flags))
}
Err(fallback_error) => {
if matches!(primary_error, MatchError::StructurallyImpossible)
&& matches!(fallback_error, MatchError::StructurallyImpossible)
{
Err(MatchError::StructurallyImpossible)
} else {
Err(fallback_error)
}
}
}
} else {
Err(primary_error)
}
}
}
#[derive(Debug)]
struct WildcardAtomIter<'a, 'b> {
name: Symbol,
optional: bool,
try_match_atom: bool,
direct_resume_point: Option<Checkpoint>,
slice: SliceAtomIter<'a, 'b>,
}
impl<'a, 'b> WildcardAtomIter<'a, 'b> {
fn new(pattern: &'b Pattern, name: Symbol, optional: bool) -> Self {
WildcardAtomIter {
name,
optional,
try_match_atom: true,
direct_resume_point: None,
slice: SliceAtomIter::new(pattern, std::slice::from_ref(pattern), SliceType::One),
}
}
fn set_target(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
force_complete: bool,
) {
self.try_match_atom = true;
self.direct_resume_point = None;
self.slice.set_target(target, match_stack, force_complete);
}
fn next(
&mut self,
target: AtomView<'a>,
match_stack: &mut WrappedMatchStack<'a, 'b>,
) -> Result<Option<&[bool]>, MatchError> {
if self.try_match_atom {
self.try_match_atom = false;
let range = match_stack.get_range_with_optional(self.name, self.optional);
if range.0 <= 1
&& range.1.map(|w| w >= 1).unwrap_or(true)
&& let Ok(new_stack_len) = match_stack.insert(self.name, Match::Single(target))
{
self.direct_resume_point = Some(Checkpoint::new(new_stack_len));
return Ok(None);
}
}
if let Some(resume_point) = self.direct_resume_point.take() {
resume_point.restore(match_stack);
}
self.slice.next(match_stack)
}
}
#[derive(Debug)]
enum AtomMatcher<'a, 'b> {
Literal(LiteralAtomIter<'b>),
Wildcard(WildcardAtomIter<'a, 'b>),
Slice(SliceAtomIter<'a, 'b>),
Function(FnMatchIterator<'a, 'b>),
Alternative(AlternativeIter<'a, 'b>),
}
impl<'a, 'b> AtomMatcher<'a, 'b> {
fn new(pattern: &'b Pattern) -> Self {
match pattern {
Pattern::Literal(literal) => AtomMatcher::Literal(LiteralAtomIter::new(literal)),
Pattern::Wildcard(name, optional) => {
AtomMatcher::Wildcard(WildcardAtomIter::new(pattern, *name, *optional))
}
Pattern::Pow(p) => {
AtomMatcher::Slice(SliceAtomIter::new(pattern, p.as_slice(), SliceType::Pow))
}
Pattern::Mul(m) => {
AtomMatcher::Slice(SliceAtomIter::new(pattern, m.as_slice(), SliceType::Mul))
}
Pattern::Add(a) => {
AtomMatcher::Slice(SliceAtomIter::new(pattern, a.as_slice(), SliceType::Add))
}
Pattern::Fn(name, args) => AtomMatcher::Function(FnMatchIterator::new(*name, args)),
Pattern::Alternative(alternatives) => {
AtomMatcher::Alternative(AlternativeIter::new(alternatives))
}
Pattern::Transformer(_) => panic!("Transformer is not allowed on lhs"),
}
}
fn set_target(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
force_complete: bool,
) {
match self {
AtomMatcher::Literal(iter) => iter.set_target(),
AtomMatcher::Wildcard(iter) => iter.set_target(target, match_stack, force_complete),
AtomMatcher::Slice(iter) => iter.set_target(target, match_stack, force_complete),
AtomMatcher::Function(iter) => iter.set_target(target),
AtomMatcher::Alternative(iter) => iter.set_target(target),
}
}
fn next(
&mut self,
target: AtomView<'a>,
match_stack: &mut WrappedMatchStack<'a, 'b>,
) -> Result<Option<&[bool]>, MatchError> {
match self {
AtomMatcher::Literal(iter) => iter.next(target).map(|_| None),
AtomMatcher::Wildcard(iter) => iter.next(target, match_stack),
AtomMatcher::Slice(iter) => iter.next(match_stack),
AtomMatcher::Function(iter) => iter.next(match_stack).map(|_| None),
AtomMatcher::Alternative(iter) => iter.next(match_stack).map(|_| None),
}
}
}
#[derive(Debug)]
pub struct AtomMatchIterator<'a, 'b> {
matcher: AtomMatcher<'a, 'b>,
target: AtomView<'a>,
used_flags: Vec<bool>,
frame: Option<Checkpoint>,
}
impl<'a, 'b> AtomMatchIterator<'a, 'b> {
pub fn new(pattern: &'b Pattern) -> AtomMatchIterator<'a, 'b> {
AtomMatchIterator {
matcher: AtomMatcher::new(pattern),
target: AtomView::ZERO,
used_flags: Vec::new(),
frame: None,
}
}
#[inline]
pub fn set_new_target(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
) {
self.set_new_target_impl(target, match_stack, false);
}
#[inline]
fn set_new_target_complete(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
) {
self.set_new_target_impl(target, match_stack, true);
}
#[inline]
fn set_new_target_impl(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
force_complete: bool,
) {
self.target = target;
self.frame = Some(Checkpoint::new(match_stack.len()));
self.matcher.set_target(target, match_stack, force_complete);
}
pub fn next(&mut self, match_stack: &mut WrappedMatchStack<'a, 'b>) -> Option<&[bool]> {
self.next_result(match_stack).ok()
}
pub fn next_result(
&mut self,
match_stack: &mut WrappedMatchStack<'a, 'b>,
) -> Result<&[bool], MatchError> {
match self.matcher.next(self.target, match_stack) {
Ok(used_flags) => {
self.used_flags.clear();
if let Some(used_flags) = used_flags {
self.used_flags.extend_from_slice(used_flags);
}
Ok(&self.used_flags)
}
Err(e) => {
self.frame.unwrap().restore(match_stack);
Err(e)
}
}
}
fn discard_current(&mut self, match_stack: &mut WrappedMatchStack<'a, 'b>) {
self.frame.unwrap().restore(match_stack);
}
}
#[derive(Debug)]
struct TypedSlice<'a> {
data: Vec<AtomView<'a>>,
slice_type: SliceType,
}
impl<'a> TypedSlice<'a> {
fn empty() -> Self {
TypedSlice {
data: Vec::new(),
slice_type: SliceType::Empty,
}
}
fn set_one(&mut self, a: AtomView<'a>) {
self.data.clear();
self.data.push(a);
self.slice_type = SliceType::One;
}
fn set_empty(&mut self, slice_type: SliceType) {
self.data.clear();
self.slice_type = slice_type;
}
fn set_list(&mut self, a: AtomView<'a>) {
match a {
AtomView::Mul(m) => {
self.data.clear();
self.data.extend(m.iter());
self.slice_type = SliceType::Mul;
}
AtomView::Add(a) => {
self.data.clear();
self.data.extend(a.iter());
self.slice_type = SliceType::Add;
}
AtomView::Pow(p) => {
self.data.clear();
self.data.extend(p.iter());
self.slice_type = SliceType::Pow;
}
AtomView::Fun(f) => {
self.data.clear();
self.data.extend(f.iter());
self.slice_type = SliceType::Arg;
}
AtomView::Var(_) | AtomView::Num(_) => {
self.data.clear();
self.data.push(a);
self.slice_type = SliceType::One;
}
}
}
fn len(&self) -> usize {
self.data.len()
}
fn get_type(&self) -> SliceType {
self.slice_type
}
fn get(&self, index: usize) -> AtomView<'a> {
self.data[index]
}
}
#[derive(Debug)]
pub struct SubSliceIterator<'a, 'b> {
pattern: &'b [Pattern], target: TypedSlice<'a>,
iterators: Vec<PatternIter<'a, 'b>>,
used_flag: Vec<bool>,
compatibility_flag: Vec<u64>, initialized: bool,
processed_iterators: usize,
frame: Option<Checkpoint>,
complete: bool, ordered_gapless: bool, cyclic: bool, do_not_match_to_single_atom_in_list: bool,
do_not_match_entire_slice: bool,
slice_type: SliceType,
}
pub enum MatchError {
StructurallyImpossible,
ImpossibleDueToConstraints,
NoMoreMatches,
}
impl<'a, 'b> SubSliceIterator<'a, 'b> {
pub fn new(pattern: &'b [Pattern], slice_type: SliceType) -> SubSliceIterator<'a, 'b> {
let iterators = pattern
.iter()
.map(|p| match &p {
Pattern::Wildcard(name, _) => PatternIter::Wildcard(WildcardIter {
initialized: false,
name: *name,
indices: Vec::new(),
size_target: 0,
min_size: 0,
max_size: 0,
greedy: false,
resume_point: None,
}),
Pattern::Transformer(_) => panic!("Transformer is not allowed on lhs"),
_ => PatternIter::Node(None, Box::new(AtomMatchIterator::new(p))),
})
.collect();
SubSliceIterator {
pattern,
iterators,
used_flag: Vec::with_capacity(pattern.len()),
compatibility_flag: Vec::with_capacity(pattern.len()),
target: TypedSlice::empty(),
initialized: false,
processed_iterators: 0,
frame: None,
complete: false,
ordered_gapless: false,
cyclic: false,
do_not_match_to_single_atom_in_list: false,
do_not_match_entire_slice: false,
slice_type,
}
}
fn optional_wildcard_can_be_empty(&self, pattern_index: usize) -> bool {
match self.slice_type {
SliceType::Add | SliceType::Mul | SliceType::Arg => true,
SliceType::Pow => pattern_index == 1,
_ => false,
}
}
fn pattern_length_range(
&self,
pattern_index: usize,
match_stack: &WrappedMatchStack<'a, 'b>,
) -> (usize, Option<usize>) {
if let Pattern::Wildcard(id, optional) = self.pattern[pattern_index] {
match_stack.get_range_with_optional(
id,
optional && self.optional_wildcard_can_be_empty(pattern_index),
)
} else {
(1, Some(1))
}
}
fn can_match_single_with_optional(&self, target: AtomView<'a>) -> bool {
if !matches!(self.slice_type, SliceType::Add | SliceType::Mul) {
return false;
}
Pattern::optional_wildcards_can_match_single(self.pattern, target)
}
fn can_match_pow_as_single(&self, target: AtomView<'a>) -> bool {
self.slice_type == SliceType::Pow
&& self.pattern.len() == 2
&& self.pattern[1].is_optional_wildcard()
&& self.pattern[0].could_match(target)
}
fn can_match_as_single_with_optional(&self, target: AtomView<'a>) -> bool {
match self.slice_type {
SliceType::Add | SliceType::Mul => self.can_match_single_with_optional(target),
SliceType::Pow => self.can_match_pow_as_single(target),
_ => false,
}
}
fn can_match_list_atom_as_single_with_optional(&self, target: AtomView<'a>) -> bool {
matches!(
(target, self.slice_type),
(AtomView::Mul(_), SliceType::Mul)
| (AtomView::Add(_), SliceType::Add)
| (AtomView::Pow(_), SliceType::Pow)
) && self.can_match_as_single_with_optional(target)
}
fn set_single_atom_target(&mut self, target: AtomView<'a>, complete: bool) {
self.target.set_one(target);
self.used_flag.clear();
self.compatibility_flag.clear();
self.used_flag.resize(self.target.len(), false);
self.compatibility_flag.resize(self.target.len(), 0);
self.initialized = false;
self.processed_iterators = 0;
self.complete = complete;
self.ordered_gapless = self.slice_type == SliceType::Pow;
self.cyclic = false;
self.do_not_match_to_single_atom_in_list = false;
self.do_not_match_entire_slice = false;
}
fn optional_default_match_for(slice_type: SliceType) -> Match<'a> {
match slice_type {
SliceType::Add => Match::Single(ZERO.as_view()),
SliceType::Mul | SliceType::Pow => Match::Single(ONE.as_view()),
SliceType::Arg => Match::Multiple(SliceType::Arg, Vec::new()),
_ => Match::Multiple(SliceType::Empty, Vec::new()),
}
}
pub fn set_target(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
do_not_match_to_single_atom_in_list: bool,
do_not_match_entire_slice: bool,
) {
let mut shortcut_done = false;
match (self.slice_type, target) {
(SliceType::Mul, AtomView::Mul(_))
| (SliceType::Add, AtomView::Add(_))
| (SliceType::Pow, AtomView::Pow(_)) => {
self.target.set_list(target);
}
(SliceType::Mul | SliceType::Add | SliceType::Pow, _) => {
self.target.set_one(target);
if !self.can_match_as_single_with_optional(target) {
shortcut_done = true; }
}
(SliceType::One, AtomView::Mul(_) | AtomView::Add(_)) => {
if matches!(self.pattern[0], Pattern::Wildcard(..)) {
self.target.set_list(target);
} else {
if do_not_match_to_single_atom_in_list
&& !matches!(self.pattern[0], Pattern::Alternative(_))
&& !self.pattern[0].could_match(target)
{
shortcut_done = true; }
self.target.set_one(target);
}
}
(_, _) => {
self.target.set_one(target);
}
};
let min_length: usize = self
.pattern
.iter()
.enumerate()
.map(|(i, _)| self.pattern_length_range(i, match_stack).0)
.sum();
let mut target_len = self.target.len();
if do_not_match_entire_slice {
target_len -= 1;
}
if min_length > target_len {
shortcut_done = true;
};
self.used_flag.clear();
self.compatibility_flag.clear();
if !shortcut_done {
self.used_flag.resize(self.target.len(), false);
self.compatibility_flag.resize(self.target.len(), 0);
}
self.initialized = shortcut_done;
self.processed_iterators = 0;
self.complete = !match_stack.settings.partial;
self.ordered_gapless = self.slice_type == SliceType::Pow;
self.cyclic = false;
self.do_not_match_to_single_atom_in_list = do_not_match_to_single_atom_in_list;
self.do_not_match_entire_slice = do_not_match_entire_slice;
self.frame = Some(Checkpoint::new(match_stack.len()));
}
pub fn set_list_target(
&mut self,
target: AtomView<'a>,
match_stack: &WrappedMatchStack<'a, 'b>,
complete: bool,
ordered: bool,
cyclic: bool,
) {
let mut shortcut_done = false;
self.target.set_list(target);
let min_length: usize = self
.pattern
.iter()
.enumerate()
.map(|(i, _)| self.pattern_length_range(i, match_stack).0)
.sum();
if min_length > self.target.len() {
shortcut_done = true;
};
let max_length: usize = self
.pattern
.iter()
.enumerate()
.map(|(i, _)| {
self.pattern_length_range(i, match_stack)
.1
.unwrap_or(self.target.len())
})
.sum();
if complete && max_length < self.target.len() {
shortcut_done = true;
};
self.used_flag.clear();
self.compatibility_flag.clear();
self.used_flag.resize(self.target.len(), false);
self.compatibility_flag.resize(self.target.len(), 0);
self.initialized = shortcut_done;
self.processed_iterators = 0;
self.complete = complete;
self.ordered_gapless = ordered;
self.cyclic = cyclic;
self.do_not_match_to_single_atom_in_list = false;
self.do_not_match_entire_slice = false;
self.frame = Some(Checkpoint::new(match_stack.len()));
}
pub fn set_empty_list_target(
&mut self,
match_stack: &WrappedMatchStack<'a, 'b>,
complete: bool,
ordered: bool,
cyclic: bool,
) {
let mut shortcut_done = false;
self.target.set_empty(self.slice_type);
let min_length: usize = self
.pattern
.iter()
.enumerate()
.map(|(i, _)| self.pattern_length_range(i, match_stack).0)
.sum();
if min_length > 0 {
shortcut_done = true;
};
self.used_flag.clear();
self.compatibility_flag.clear();
self.initialized = shortcut_done;
self.processed_iterators = 0;
self.complete = complete;
self.ordered_gapless = ordered;
self.cyclic = cyclic;
self.do_not_match_to_single_atom_in_list = false;
self.do_not_match_entire_slice = false;
self.frame = Some(Checkpoint::new(match_stack.len()));
}
pub fn next(
&mut self,
match_stack: &mut WrappedMatchStack<'a, 'b>,
) -> Result<&[bool], MatchError> {
let mut forward_pass = !self.initialized;
self.initialized = true;
let mut structural_mismatch = forward_pass;
'next_match: loop {
if !forward_pass && self.processed_iterators == 0 {
let error = if structural_mismatch {
MatchError::StructurallyImpossible
} else {
MatchError::NoMoreMatches
};
self.frame.unwrap().restore(match_stack);
return Err(error);
}
if forward_pass && self.processed_iterators == self.pattern.len() {
if self.complete && self.used_flag.iter().any(|x| !*x)
|| self.do_not_match_to_single_atom_in_list && self.used_flag.len() > 1
&& self.used_flag.iter().map(|x| *x as usize).sum::<usize>() == 1
{
forward_pass = false;
} else {
return Ok(&self.used_flag);
}
}
if forward_pass {
let wildcard_ranges = if matches!(
self.pattern[self.processed_iterators],
Pattern::Wildcard(..)
) {
Some((
self.pattern_length_range(self.processed_iterators, match_stack),
if self.complete {
self.pattern[self.processed_iterators + 1..]
.iter()
.enumerate()
.map(|(offset, _)| {
self.pattern_length_range(
self.processed_iterators + 1 + offset,
match_stack,
)
})
.collect::<Vec<_>>()
} else {
Vec::new()
},
))
} else {
None
};
let it = &mut self.iterators[self.processed_iterators];
match (&self.pattern[self.processed_iterators], it) {
(Pattern::Wildcard(name, _), PatternIter::Wildcard(w)) => {
let mut size_left = self.used_flag.iter().filter(|x| !*x).count();
let (range, future_ranges) = wildcard_ranges.as_ref().unwrap();
let range = *range;
if name.get_wildcard_level() > 1 && match_stack.stack.get(*name).is_some() {
structural_mismatch = false;
}
if self.do_not_match_entire_slice {
size_left -= 1;
if size_left < range.0 {
forward_pass = false;
continue 'next_match;
}
}
let mut range = (
range.0,
range.1.map(|m| m.min(size_left)).unwrap_or(size_left),
);
if self.complete {
let mut new_min = size_left;
let mut new_max = size_left;
for p_range in future_ranges {
if new_min > 0 {
if let Some(m) = p_range.1 {
new_min -= m.min(new_min);
} else {
new_min = 0;
}
}
if new_max < p_range.0 {
forward_pass = false;
continue 'next_match;
}
new_max -= p_range.0;
}
range.0 = range.0.max(new_min);
range.1 = range.1.min(new_max);
if range.0 > range.1 {
forward_pass = false;
continue 'next_match;
}
}
let greedy = !match_stack.settings.non_greedy_wildcards.contains(name);
w.initialized = false;
w.name = *name;
w.indices.clear();
w.size_target = if greedy {
range.1 as u32
} else {
range.0 as u32
};
w.min_size = range.0 as u32;
w.max_size = range.1 as u32;
w.greedy = greedy;
w.resume_point = None;
}
(Pattern::Transformer(_), _) => panic!("Transformer is not allowed on lhs"),
(_, PatternIter::Node(index, _)) => *index = None,
(p, i) => panic!("Pattern and iterator type mismatch: {:?} vs {:?}", p, i),
}
self.processed_iterators += 1;
}
forward_pass = true;
let optional_default_match = Self::optional_default_match_for(self.slice_type);
match &mut self.iterators[self.processed_iterators - 1] {
PatternIter::Wildcard(w) => {
let mut wildcard_forward_pass = !w.initialized;
w.initialized = true;
if !wildcard_forward_pass {
w.resume_point.take().unwrap().restore(match_stack);
}
'next_wildcard_match: loop {
let start_index =
w.indices
.last()
.map(|x| *x as usize + 1)
.unwrap_or_else(|| {
if self.cyclic {
let mut pos =
self.used_flag.iter().position(|x| *x).unwrap_or(0);
while self.used_flag[pos] {
pos = (pos + 1) % self.used_flag.len();
}
pos
} else {
0
}
});
if !wildcard_forward_pass {
let last_iterator_empty = w.indices.is_empty();
if let Some(last_index) = w.indices.pop() {
self.used_flag[last_index as usize] = false;
}
if last_iterator_empty {
if w.greedy {
if w.size_target > w.min_size {
w.size_target -= 1;
} else {
break;
}
} else if w.size_target < w.max_size {
w.size_target += 1;
} else {
break;
}
} else if self.ordered_gapless {
if !self.cyclic || self.used_flag.iter().any(|x| *x) {
continue 'next_wildcard_match;
}
}
}
if !self.cyclic {
let remaining_available = self.used_flag[start_index..]
.iter()
.filter(|used| !**used)
.count();
if w.indices.len() + remaining_available < w.size_target as usize {
wildcard_forward_pass = false;
continue 'next_wildcard_match;
}
}
if w.size_target == 0 && w.indices.is_empty() {
match match_stack.insert(w.name, optional_default_match.clone()) {
Ok(new_stack_len) => {
w.resume_point = Some(Checkpoint::new(new_stack_len));
continue 'next_match;
}
Err(MatchError::StructurallyImpossible) => {}
Err(_) => {
structural_mismatch = false;
}
}
wildcard_forward_pass = false;
continue 'next_wildcard_match;
}
let mut tried_first_option = false;
let mut k = start_index;
loop {
if k == self.target.len() {
if self.cyclic && !w.indices.is_empty() {
k = 0;
} else {
break;
}
}
if self.ordered_gapless && tried_first_option {
if !self.cyclic || self.used_flag.iter().any(|x| *x) {
break;
}
}
if self.used_flag[k]
|| w.name.get_wildcard_level() == 1
&& self.processed_iterators < 64
&& self.compatibility_flag[k]
& (1 << (self.processed_iterators - 1))
!= 0
{
if self.cyclic {
break;
}
k += 1;
continue;
}
self.used_flag[k] = true;
w.indices.push(k as u32);
if w.indices.len() == w.size_target as usize {
tried_first_option = true;
let matched = if w.indices.len() == 1 {
match self.target.get(k) {
AtomView::Mul(m) => Match::Multiple(SliceType::Mul, {
let mut v = Vec::new();
for x in m {
v.push(x);
}
v
}),
AtomView::Add(a) => Match::Multiple(SliceType::Add, {
let mut v = Vec::new();
for x in a {
v.push(x);
}
v
}),
x => Match::Single(x),
}
} else {
let mut atoms = Vec::with_capacity(w.indices.len());
for i in &w.indices {
atoms.push(self.target.get(*i as usize));
}
Match::Multiple(self.target.get_type(), atoms)
};
match match_stack.insert(w.name, matched) {
Ok(new_stack_len) => {
w.resume_point = Some(Checkpoint::new(new_stack_len));
continue 'next_match;
}
Err(MatchError::StructurallyImpossible) => {
if self.processed_iterators < 64
&& w.name.get_wildcard_level() == 1
{
self.compatibility_flag[k] |=
1 << (self.processed_iterators - 1) as u64;
}
}
Err(_) => {
structural_mismatch = false;
}
}
w.indices.pop();
self.used_flag[k] = false;
}
k += 1;
}
wildcard_forward_pass = false;
}
}
PatternIter::Node(index, s) => {
let mut tried_first_option = false;
let mut ii = match index {
Some(jj) => {
if !structural_mismatch {
match s.next_result(match_stack) {
Ok(_) => {
continue 'next_match;
}
Err(_) => {}
}
} else {
s.discard_current(match_stack);
}
self.used_flag[*jj] = false;
tried_first_option = true;
*jj + 1
}
None => {
if self.cyclic && !self.used_flag.iter().all(|u| *u) {
let mut pos = self.used_flag.iter().position(|x| *x).unwrap_or(0);
while self.used_flag[pos] {
pos = (pos + 1) % self.used_flag.len();
}
pos
} else {
0
}
}
};
while ii < self.target.len() {
if self.used_flag[ii]
|| self.processed_iterators < 64
&& self.compatibility_flag[ii]
& (1 << (self.processed_iterators - 1))
!= 0
{
if self.cyclic {
break;
}
ii += 1;
continue;
}
if self.ordered_gapless && tried_first_option {
if !self.cyclic || self.used_flag.iter().any(|x| *x) {
break;
}
}
tried_first_option = true;
let new_target = self.target.get(ii);
s.set_new_target_complete(new_target, match_stack);
match s.next_result(match_stack) {
Ok(_) => {
*index = Some(ii);
self.used_flag[ii] = true;
continue 'next_match;
}
Err(e) => {
if matches!(e, MatchError::StructurallyImpossible) {
if self.processed_iterators < 64 {
self.compatibility_flag[ii] |=
1 << (self.processed_iterators - 1) as u64;
};
} else {
structural_mismatch = false;
}
}
}
ii += 1;
}
}
}
forward_pass = false;
self.processed_iterators -= 1;
}
}
}
pub struct AtomTreeIterator<'a> {
stack: Vec<(Option<usize>, usize, ListIterator<'a>)>,
settings: MatchSettings,
}
impl<'a> AtomTreeIterator<'a> {
pub fn new(target: AtomView<'a>, settings: MatchSettings) -> AtomTreeIterator<'a> {
AtomTreeIterator {
stack: vec![(None, 0, ListIterator::from_one(target))],
settings,
}
}
pub fn reset(&mut self, target: AtomView<'a>) {
self.stack.clear();
self.stack.push((None, 0, ListIterator::from_one(target)));
}
pub fn next_into(&mut self, mut position: Option<&mut Vec<usize>>) -> Option<AtomView<'a>> {
while let Some((ind, level, mut slice)) = self.stack.pop() {
if let Some(max_level) = self.settings.level_range.1
&& level > max_level
{
continue;
}
if let Some(ind) = ind {
if let Some(sub_atom) = slice.next() {
self.stack.push((Some(ind + 1), level, slice)); self.stack
.push((None, level, ListIterator::from_one(sub_atom))); }
} else {
if let Some(position) = position.as_mut() {
position.clear();
for s in &self.stack {
position.push(s.0.unwrap() - 1);
}
}
let atom = slice.next().unwrap();
let new_level = if self.settings.level_is_tree_depth {
level + 1
} else {
level
};
match atom {
AtomView::Fun(f) => self.stack.push((Some(0), level + 1, f.iter())),
AtomView::Pow(p) => self.stack.push((Some(0), new_level, p.iter())),
AtomView::Mul(m) => self.stack.push((Some(0), new_level, m.iter())),
AtomView::Add(a) => self.stack.push((Some(0), new_level, a.iter())),
_ => {}
}
if level >= self.settings.level_range.0 {
return Some(atom);
}
}
}
None
}
}
impl<'a> Iterator for AtomTreeIterator<'a> {
type Item = (Vec<usize>, AtomView<'a>);
fn next(&mut self) -> Option<Self::Item> {
let mut location = Vec::new();
self.next_into(Some(&mut location))
.map(|atom| (location, atom))
}
}
pub struct PatternAtomTreeIterator<'a, 'b> {
atom_tree_iterator: AtomTreeIterator<'a>,
pattern_iter: AtomMatchIterator<'a, 'b>,
match_stack: WrappedMatchStack<'a, 'b>,
tree_pos: Vec<usize>,
used_flags: Vec<bool>,
first_match: bool,
}
pub struct PatternMatch<'a, 'b> {
pub position: &'b [usize],
pub used_flags: &'b [bool],
pub target: AtomView<'a>,
pub match_stack: &'b MatchStack<'a>,
}
impl<'a: 'b, 'b> PatternAtomTreeIterator<'a, 'b> {
pub fn new(
pattern: &'b Pattern,
target: AtomView<'a>,
conditions: Option<&'b Condition<PatternRestriction>>,
settings: Option<&'b MatchSettings>,
) -> PatternAtomTreeIterator<'a, 'b> {
let mut it =
AtomTreeIterator::new(target, settings.unwrap_or(&DEFAULT_MATCH_SETTINGS).clone());
it.next();
let match_stack = WrappedMatchStack::new(
conditions.unwrap_or(&DEFAULT_PATTERN_CONDITION),
settings.unwrap_or(&DEFAULT_MATCH_SETTINGS),
);
let mut pattern_iter = AtomMatchIterator::new(pattern);
pattern_iter.set_new_target(target, &match_stack);
PatternAtomTreeIterator {
atom_tree_iterator: it,
pattern_iter,
match_stack,
tree_pos: Vec::new(),
used_flags: Vec::new(),
first_match: false,
}
}
pub fn next_detailed(&mut self) -> Option<PatternMatch<'a, '_>> {
loop {
if let Some(used_flags) = self.pattern_iter.next(&mut self.match_stack) {
self.used_flags.clear();
self.used_flags.extend_from_slice(used_flags);
self.first_match = true;
return Some(PatternMatch {
position: &self.tree_pos,
used_flags: &self.used_flags,
target: self.pattern_iter.target,
match_stack: &self.match_stack.stack,
});
}
if !self.match_stack.settings.partial {
return None;
}
if let Some(cur_target) = self.atom_tree_iterator.next_into(Some(&mut self.tree_pos)) {
self.pattern_iter
.set_new_target(cur_target, &self.match_stack);
} else {
return None;
}
}
}
}
impl<'a: 'b, 'b> Iterator for PatternAtomTreeIterator<'a, 'b> {
type Item = HashMap<Symbol, Atom>;
fn next(&mut self) -> Option<HashMap<Symbol, Atom>> {
if self.next_detailed().is_some() {
Some(
self.match_stack
.get_matches()
.iter()
.map(|(key, m)| (*key, m.to_atom()))
.collect(),
)
} else {
None
}
}
}
pub struct ReplaceIterator<'a, 'b> {
rhs: ReplaceWith<'b>,
pattern_tree_iterator: PatternAtomTreeIterator<'a, 'b>,
target: AtomView<'a>,
}
impl<'a: 'b, 'b> ReplaceIterator<'a, 'b> {
pub fn new(
pattern: &'b Pattern,
target: AtomView<'a>,
rhs: ReplaceWith<'b>,
conditions: Option<&'a Condition<PatternRestriction>>,
settings: Option<&'a MatchSettings>,
) -> ReplaceIterator<'a, 'b> {
ReplaceIterator {
pattern_tree_iterator: PatternAtomTreeIterator::new(
pattern, target, conditions, settings,
),
rhs,
target,
}
}
fn copy_and_replace(
out: &mut Atom,
position: &[usize],
used_flags: &[bool],
target: AtomView<'a>,
rhs: AtomView<'_>,
workspace: &Workspace,
) {
if let Some((first, rest)) = position.split_first() {
match target {
AtomView::Fun(f) => {
let slice = f.to_slice();
let out = out.to_fun(f.get_symbol());
let mut oa = workspace.new_atom();
for (index, arg) in slice.iter().enumerate() {
if index == *first {
Self::copy_and_replace(&mut oa, rest, used_flags, arg, rhs, workspace);
out.add_arg(oa.as_view());
} else {
out.add_arg(arg);
}
}
}
AtomView::Pow(p) => {
let slice = p.to_slice();
if *first == 0 {
let mut oa = workspace.new_atom();
Self::copy_and_replace(
&mut oa,
rest,
used_flags,
slice.get(0),
rhs,
workspace,
);
out.to_pow(oa.as_view(), slice.get(1));
} else {
let mut oa = workspace.new_atom();
Self::copy_and_replace(
&mut oa,
rest,
used_flags,
slice.get(1),
rhs,
workspace,
);
out.to_pow(slice.get(0), oa.as_view());
}
}
AtomView::Mul(m) => {
let slice = m.to_slice();
let out = out.to_mul();
let mut oa = workspace.new_atom();
for (index, arg) in slice.iter().enumerate() {
if index == *first {
Self::copy_and_replace(&mut oa, rest, used_flags, arg, rhs, workspace);
out.extend(oa.as_view());
} else {
out.extend(arg);
}
}
}
AtomView::Add(a) => {
let slice = a.to_slice();
let out = out.to_add();
let mut oa = workspace.new_atom();
for (index, arg) in slice.iter().enumerate() {
if index == *first {
Self::copy_and_replace(&mut oa, rest, used_flags, arg, rhs, workspace);
out.extend(oa.as_view());
} else {
out.extend(arg);
}
}
}
_ => unreachable!("Atom does not have children"),
}
} else {
match target {
AtomView::Mul(m) => {
let out = out.to_mul();
for (child, used) in m.iter().zip(used_flags) {
if !used {
out.extend(child);
}
}
out.extend(rhs);
}
AtomView::Add(a) => {
let out = out.to_add();
for (child, used) in a.iter().zip(used_flags) {
if !used {
out.extend(child);
}
}
out.extend(rhs);
}
_ => {
out.set_from_view(&rhs);
}
}
}
}
pub fn next_into(&mut self, out: &mut Atom) -> Option<()> {
let allow = self
.pattern_tree_iterator
.atom_tree_iterator
.settings
.allow_new_wildcards_on_rhs;
if let Some(pattern_match) = self.pattern_tree_iterator.next_detailed() {
Workspace::get_local().with(|ws| {
let mut new_rhs = ws.new_atom();
match &self.rhs {
ReplaceWith::Pattern(p) => {
p.replace_wildcards_with_matches_impl(
ws,
&mut new_rhs,
pattern_match.match_stack,
allow,
None,
)
.unwrap(); }
ReplaceWith::Map(f) => {
let mut new_atom = f(pattern_match.match_stack);
std::mem::swap(&mut new_atom, &mut new_rhs);
}
}
let mut h = ws.new_atom();
ReplaceIterator::copy_and_replace(
&mut h,
pattern_match.position,
pattern_match.used_flags,
self.target,
new_rhs.as_view(),
ws,
);
h.as_view().normalize(ws, out);
});
Some(())
} else {
None
}
}
}
impl<'a: 'b, 'b> Iterator for ReplaceIterator<'a, 'b> {
type Item = Atom;
fn next(&mut self) -> Option<Self::Item> {
let mut out = Atom::new();
self.next_into(&mut out).map(|_| out)
}
}
#[cfg(test)]
mod test {
use super::{
AtomMatchIterator, DEFAULT_MATCH_SETTINGS, DEFAULT_PATTERN_CONDITION, WrappedMatchStack,
};
use crate::{
atom::{Atom, AtomCore, AtomType},
id::{AtomTreeIterator, Condition, ConditionResult, Match, MatchSettings, Replacement},
parse,
printer::PrintOptions,
symbol,
};
fn assert_contains_all(got: &[Atom], expected: &[Atom]) {
for expected in expected {
assert!(
got.contains(expected),
"missing expected replacement {expected}; got {got:?}"
);
}
}
#[test]
fn repeated_optional() {
let x = parse!("x*y");
let pattern = parse!("(a_+x)*(a_+y)")
.to_pattern()
.set_optional(symbol!("a_"));
let r = x.replace(pattern).with(1);
assert_eq!(r, 1);
}
#[test]
fn complete_match() {
let input = parse!("f(1)*f(2)");
let pat = input.replace(parse!("f(x_)")).partial(false);
let mut it = pat.match_iter();
assert_eq!(it.next(), None);
let pat = input.replace(parse!("f(x_)*f(y_)")).partial(false);
let mut it = pat.iter(parse!("g(x_,y_)"));
assert_eq!(it.next(), Some(parse!("g(1,2)")));
assert_eq!(it.next(), Some(parse!("g(2,1)")));
}
#[test]
fn exhaust_large_fixed_size_wildcard() {
let input = parse!(
"f(a0+a1+a2+a3+a4+a5+a6+a7+a8+a9+
a10+a11+a12+a13+a14+a15+a16+a17+a18+a19+
a20+a21+a22+a23+a24+a25+a26+a27+a28+x*y)"
);
let pattern = parse!("f(g__+x*h__)").to_pattern();
let settings = MatchSettings::default();
assert_eq!(
input.pattern_match(&pattern, None, Some(&settings)).count(),
1
);
}
#[test]
fn atom_tree_iterator() {
let a = parse!("v1*f1(v2 + v3, v3^2, f1(v1))");
let mut it = AtomTreeIterator::new(a.as_view(), MatchSettings::default());
assert_eq!(it.next().unwrap(), (vec![], a.as_view()));
assert_eq!(it.next().unwrap(), (vec![0], parse!("v1").as_view()));
assert_eq!(
it.next().unwrap(),
(vec![1], parse!("f1(v2+v3,v3^2,f1(v1))").as_view()),
);
assert_eq!(it.next().unwrap(), (vec![1, 0], parse!("v2+v3").as_view()),);
assert_eq!(it.next().unwrap(), (vec![1, 0, 0], parse!("v2").as_view()));
assert_eq!(it.next().unwrap(), (vec![1, 0, 1], parse!("v3").as_view()));
assert_eq!(it.next().unwrap(), (vec![1, 1], parse!("v3^2").as_view()));
assert_eq!(it.next().unwrap(), (vec![1, 1, 0], parse!("v3").as_view()));
assert_eq!(it.next().unwrap(), (vec![1, 1, 1], parse!("2").as_view()));
assert_eq!(it.next().unwrap(), (vec![1, 2], parse!("f1(v1)").as_view()));
assert_eq!(it.next().unwrap(), (vec![1, 2, 0], parse!("v1").as_view()));
assert!(it.next().is_none());
}
#[test]
fn replace_iter() {
let a = parse!("v1*v2*v3*f1(v4)");
let pat = a.replace(parse!("x_"));
let mut it = pat.iter(parse!("v5"));
assert_eq!(it.next().unwrap(), parse!("v5"));
assert_eq!(it.next().unwrap(), parse!("v2*v3*v5*f1(v4)"));
assert_eq!(it.next().unwrap(), parse!("v1*v3*v5*f1(v4)"));
assert_eq!(it.next().unwrap(), parse!("v1*v2*v5*f1(v4)"));
assert_eq!(it.next().unwrap(), parse!("v1*v2*v3*v5"));
assert_eq!(it.next().unwrap(), parse!("v1*v2*v3*f1(v5)"));
assert!(it.next().is_none());
}
#[test]
fn replace_wildcards_with_map() {
let a = parse!("f1(v1__, 5) + v1*v2_ + v3^v3_").to_pattern();
let r = a
.replace_wildcards(
&[
(symbol!("v1__"), parse!("arg(v4, v5)")),
(symbol!("v2_"), Atom::num(4)),
(symbol!("v3_"), Atom::num(5)),
]
.into_iter()
.collect(),
)
.unwrap();
let res = parse!("f1(v4, v5, 5) + v1*4 + v3^5");
assert_eq!(r, res);
}
#[test]
fn replace_wildcards() {
let a = parse!("f1(v1__, 5) + v1*v2_ + v3^v3_").to_pattern();
let r11 = Atom::var(symbol!("v4"));
let r12 = Atom::var(symbol!("v5"));
let r2 = Atom::num(4);
let r3 = Atom::num(5);
let r = a.replace_wildcards_with_matches(
&vec![
(
symbol!("v1__"),
Match::Multiple(
crate::atom::SliceType::Arg,
vec![r11.as_view(), r12.as_view()],
),
),
(symbol!("v2_"), Match::Single(r2.as_view())),
(symbol!("v3_"), Match::Single(r3.as_view())),
]
.into(),
);
let res = parse!("f1(v4, v5, 5) + v1*4 + v3^5");
assert_eq!(r, res);
}
#[test]
fn replace_map() {
let a = parse!("v1 + f1(1,2, f1((1+v1)^2), (v1+v2)^2)");
let mut tmp = Atom::new();
let r = a.replace_map(move |arg, context, out| {
if context.function_level > 0 {
if arg.expand_into(None, &mut tmp) {
out.set_from_view(&tmp.as_view());
}
}
});
let res = parse!("v1+f1(1,2,f1(2*v1+v1^2+1),v1^2+v2^2+2*v1*v2)");
assert_eq!(r, res);
}
#[test]
fn overlap() {
let a = parse!("(v1*(v2+v2^2+1)+v2^2 + v2)");
let p = parse!("v2+v2^v1_");
let rhs = parse!("v2*(1+v2^(v1_-1))");
let r = a.replace(p).with(rhs);
let res = parse!("v1*(v2+v2^2+1)+v2*(v2+1)");
assert_eq!(r, res);
}
#[test]
fn level_restriction() {
let a = parse!("v1*f1(v1,f1(v1))");
let p = parse!("v1");
let rhs = parse!("1");
let r = a.replace(p).level_range((1, Some(1))).with(rhs);
let res = parse!("v1*f1(1,f1(v1))");
assert_eq!(r, res);
}
#[test]
fn multiple() {
let a = parse!("f(v1,v2)");
let r = a.replace_multiple([
Replacement::new(parse!("v1"), parse!("v2")),
Replacement::new(parse!("v2"), parse!("v1")),
]);
let res = parse!("f(v2,v1)");
assert_eq!(r, res);
}
#[test]
fn map_rhs() {
let (v1, v2, v4, v5) = symbol!("v1_", "v2_", "v4_", "v5_");
let a = parse!("v1(2,1)*v2(3,1)");
let p = parse!("v1_(v2_,v3_)*v4_(v5_,v3_)");
let r = a.replace(p).with_map(move |m| {
let s = format!(
"{}(mu{})*{}(mu{})",
m.get(v1).unwrap().to_atom().printer(PrintOptions::file()),
m.get_atom(v2).unwrap().printer(PrintOptions::file()),
m.get(v4).unwrap().to_atom().printer(PrintOptions::file()),
m.get_atom(v5).unwrap().printer(PrintOptions::file())
);
parse!(&s)
});
let res = parse!("v1(mu2)*v2(mu3)");
assert_eq!(r, res);
}
#[test]
fn repeat_replace() {
let mut a = parse!("f(10)");
let p1 = parse!("f(v1_)").to_pattern();
let rhs1 = parse!("f(v1_ - 1)").to_pattern();
let rest = symbol!("v1_").filter_match(|x| {
let n: Result<i64, _> = x.to_atom().try_into();
if let Ok(y) = n { y > 0i64 } else { false }
});
a = a.replace(p1).when(rest).repeat().with(rhs1);
let res = parse!("f(0)");
assert_eq!(a, res);
}
#[test]
fn repeat_replace_same_input_output() {
let mut a = parse!("2");
let p1 = parse!("x_").to_pattern();
let rhs1 = parse!("x_").to_pattern();
a = a.replace(p1).repeat().with(rhs1);
let res = parse!("2");
assert_eq!(a, res);
}
#[test]
fn match_stack_filter() {
let a = parse!("f(1,2,3,4)");
let p1 = parse!("f(v1_,v2_,v3_,v4_)").to_pattern();
let rhs1 = parse!("f(v4_,v3_,v2_,v1_)").to_pattern();
let rest = Condition::match_stack(|m| {
for x in m.stack.windows(2) {
if x[0].1.to_atom() >= x[1].1.to_atom() {
return false.into();
}
}
if m.stack.len() == 4 {
true.into()
} else {
ConditionResult::Inconclusive
}
});
let r = a.replace(&p1).when(&rest).with(&rhs1);
let res = parse!("f(4,3,2,1)");
assert_eq!(r, res);
let b = parse!("f(1,2,4,3)");
let r = b.replace(p1).when(rest).with(rhs1);
assert_eq!(r, b);
}
#[test]
fn match_cache() {
let expr = parse!("f1(1)*f1(2)+f1(1)*f1(2)*f2");
let pat = parse!("v1_(id1_)*v2_(id2_)");
let replacements = expr
.replace(pat.clone())
.iter(parse!("f1(id1_)"))
.collect::<Vec<_>>();
let expected_replacements = [
parse!("f2*f1(1)+f1(1)*f1(2)"),
parse!("f2*f1(2)+f1(1)*f1(2)"),
parse!("f2*f1(1)*f1(2)+f1(1)"),
parse!("f2*f1(1)*f1(2)+f1(2)"),
];
assert_eq!(replacements.len(), expected_replacements.len());
for expected in expected_replacements {
assert!(
replacements.contains(&expected),
"missing expected replacement {expected}; got {replacements:?}"
);
}
let expr = expr.replace(pat).with(parse!("f1(id1_)"));
assert!(
[parse!("f1(1)+f2*f1(1)"), parse!("f1(2)+f2*f1(2)")].contains(&expr),
"unexpected cached replacement result {expr}"
);
}
#[test]
fn match_cyclic() {
let rhs = parse!("1").to_pattern();
let expr = parse!("fc1(1,2,3)");
let p = parse!("fc1(v1__,v1_,1)");
let expr = expr.replace(p).with(&rhs);
assert_eq!(expr, 1);
let expr = parse!("fc1(1,2,3)");
let p = parse!("fc1(v1__,2)");
let expr = expr.replace(p).with(&rhs);
assert_eq!(expr, 1);
let expr = parse!("fc1(1,2,3)");
let p = parse!("fc1(v1__,v1_,2)");
let expr = expr.replace(p).with(&rhs);
assert_eq!(expr, 1);
let expr = parse!("fc1(v1,4,3,5,4)");
let p = parse!("fc1(v1__,v1_,v2_,v1_)");
let expr = expr.replace(p).with(&rhs);
assert_eq!(expr, 1);
let expr = parse!("fc1(f1(1),f1(2),f1(3))");
let p = parse!("fc1(f1(v1_),f1(2),f1(3))");
let expr = expr.replace(p).with(&rhs);
assert_eq!(expr, 1);
let expr = parse!("fc1(1,2,3)*f(2)");
let p = parse!("fc1(v1_,v2___)*f(v1_)");
let expr = expr.replace(p).with(&rhs);
assert_eq!(expr, 1);
}
#[test]
fn is_polynomial() {
let e = parse!("v1^2 + (1+v5)^3 / v1 + (1+v3)*(1+v4)^v7 + v1^2 + (v1+v2)^3");
let vars = e.as_view().is_polynomial(true, true).unwrap();
assert_eq!(vars.len(), 5);
let e = parse!("(1+v5)^(3/2) / v6 + (1+v3)*(1+v4)^v7 + (v1+v2)^3");
let vars = e.as_view().is_polynomial(false, false).unwrap();
assert_eq!(vars.len(), 5);
}
#[test]
fn symbol_attribute_filter() {
let _ = symbol!("symbolica::symbol_attribute_filter::fsym"; Symmetric);
let _ = symbol!("symbolica::symbol_attribute_filter::fsym_"; Symmetric);
let _ = symbol!("symbolica::symbol_attribute_filter::xscal"; Scalar);
let _ = symbol!("symbolica::symbol_attribute_filter::xscal__"; Scalar);
let r = parse!("f(1)")
.replace(parse!("symbolica::symbol_attribute_filter::fsym_(x_)"))
.with(1);
assert_ne!(r, 1);
let r = parse!("symbolica::symbol_attribute_filter::fsym(1,symbolica::symbol_attribute_filter::xscal^2 + 2,3)")
.replace(parse!("symbolica::symbol_attribute_filter::fsym_(symbolica::symbol_attribute_filter::xscal__)"))
.with(1);
assert_eq!(r, 1);
let r = parse!("f(1,x,2)")
.replace(parse!("f(symbolica::symbol_attribute_filter::xscal__)"))
.with(1);
assert_eq!(r, parse!("f(1,x,2)"));
}
#[test]
fn nested() {
let res = parse!("f(x+x*y+f(x+x^2),x,f(x+x^2))").replace_map_bottom_up(|a, c, o| {
if c.parent_type == Some(AtomType::Fun) {
let r = a.horner_scheme(None, false, false);
if r.as_view() != a {
**o = r;
}
}
});
assert_eq!(res, parse!("f(x*(1+y)+f(x*(1+x)),x,f(x*(1+x)))"));
}
#[test]
fn alternative() {
let pat = (parse!("x").to_pattern() | parse!("y").to_pattern()) * parse!("x_").to_pattern();
assert_eq!(parse!("alt(x,y)*x_"), pat.to_atom().unwrap());
let rhs = parse!("f(x_)").to_pattern();
let e = parse!("x*z").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(z)"));
let e = parse!("y*z").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(z)"));
let e = parse!("a*z").replace(&pat).with(&rhs);
assert_eq!(e, parse!("a*z"));
}
#[test]
fn match_stack_reset_during_backtracking() {
let replacements = parse!("f(1)*g(2)*g(3)")
.replace(parse!("h_(x_)*h_(y_)"))
.iter(parse!("hit(h_,x_,y_)"))
.collect::<Vec<_>>();
let expected = [parse!("f(1)*hit(g,2,3)"), parse!("f(1)*hit(g,3,2)")];
assert_eq!(replacements.len(), expected.len());
assert_contains_all(&replacements, &expected);
let alt = parse!("f(x_,0)").to_pattern() | parse!("f(0,x_)").to_pattern();
let pat = alt * parse!("f(x_,1)").to_pattern();
let replacements = parse!("f(1,0)*f(0,2)*f(2,1)")
.replace(&pat)
.iter(parse!("hit(x_)"))
.collect::<Vec<_>>();
let expected = [parse!("f(1,0)*hit(2)")];
assert_eq!(replacements.len(), expected.len());
assert_contains_all(&replacements, &expected);
let pat = parse!("g(x_*o_)*f(x_)")
.to_pattern()
.set_optional(symbol!("o_"));
let replacements = parse!("g(a*b)*f(a*b)")
.replace(&pat)
.iter(parse!("hit(x_,o_)"))
.collect::<Vec<_>>();
let expected = [parse!("hit(a*b,1)")];
assert_eq!(replacements.len(), expected.len());
assert_contains_all(&replacements, &expected);
}
#[test]
fn match_iterator_owns_cleanup_boundary() {
let outer_value = parse!("z");
let target = parse!("f(a,g(b))");
let pattern = parse!("f(x_,g(y_))").to_pattern();
let outer = symbol!("outer_");
let mut match_stack =
WrappedMatchStack::new(&DEFAULT_PATTERN_CONDITION, &DEFAULT_MATCH_SETTINGS);
assert!(
match_stack
.insert(outer, Match::Single(outer_value.as_view()))
.is_ok()
);
let mut iter = AtomMatchIterator::new(&pattern);
iter.set_new_target(target.as_view(), &match_stack);
assert!(iter.next_result(&mut match_stack).is_ok());
assert!(match_stack.len() > 1);
iter.discard_current(&mut match_stack);
assert_eq!(match_stack.len(), 1);
assert_eq!(
match_stack.stack.get(outer),
Some(&Match::Single(outer_value.as_view()))
);
iter.set_new_target(target.as_view(), &match_stack);
while iter.next_result(&mut match_stack).is_ok() {}
assert_eq!(match_stack.len(), 1);
assert!(iter.next_result(&mut match_stack).is_err());
assert_eq!(match_stack.len(), 1);
}
#[test]
fn discard_match_that_adds_no_binding() {
let target = parse!("a");
let pattern = parse!("x_").to_pattern();
let wildcard = symbol!("x_");
let mut match_stack =
WrappedMatchStack::new(&DEFAULT_PATTERN_CONDITION, &DEFAULT_MATCH_SETTINGS);
assert!(
match_stack
.insert(wildcard, Match::Single(target.as_view()))
.is_ok()
);
let mut iter = AtomMatchIterator::new(&pattern);
iter.set_new_target(target.as_view(), &match_stack);
assert!(iter.next_result(&mut match_stack).is_ok());
assert_eq!(match_stack.len(), 1);
iter.discard_current(&mut match_stack);
assert_eq!(match_stack.len(), 1);
}
#[test]
fn nested_function_bindings_restore_to_the_outer_frame() {
let outer_value = parse!("sentinel");
let target = parse!("g(a(b),c(d))");
let pattern = parse!("g(f_(x_),h_(y_))").to_pattern();
let outer = symbol!("outer_");
let f = symbol!("f_");
let h = symbol!("h_");
let x = symbol!("x_");
let y = symbol!("y_");
let mut match_stack =
WrappedMatchStack::new(&DEFAULT_PATTERN_CONDITION, &DEFAULT_MATCH_SETTINGS);
assert!(
match_stack
.insert(outer, Match::Single(outer_value.as_view()))
.is_ok()
);
let mut iter = AtomMatchIterator::new(&pattern);
iter.set_new_target_complete(target.as_view(), &match_stack);
assert!(iter.next_result(&mut match_stack).is_ok());
assert_eq!(match_stack.len(), 5);
assert_eq!(
match_stack.stack.get(f),
Some(&Match::FunctionName(symbol!("a")))
);
assert_eq!(
match_stack.stack.get(h),
Some(&Match::FunctionName(symbol!("c")))
);
assert_eq!(
match_stack.stack.get(x),
Some(&Match::Single(parse!("b").as_view()))
);
assert_eq!(
match_stack.stack.get(y),
Some(&Match::Single(parse!("d").as_view()))
);
assert!(iter.next_result(&mut match_stack).is_err());
assert_eq!(match_stack.len(), 1);
assert_eq!(
match_stack.stack.get(outer),
Some(&Match::Single(outer_value.as_view()))
);
}
#[test]
fn failed_alternative_variant_does_not_leak_bindings() {
let outer_value = parse!("sentinel");
let target = parse!("f(1,2)");
let pattern = parse!("f(x_,0)").to_pattern() | parse!("f(1,x_)").to_pattern();
let outer = symbol!("outer_");
let x = symbol!("x_");
let two = parse!("2");
let mut match_stack =
WrappedMatchStack::new(&DEFAULT_PATTERN_CONDITION, &DEFAULT_MATCH_SETTINGS);
assert!(
match_stack
.insert(outer, Match::Single(outer_value.as_view()))
.is_ok()
);
let mut iter = AtomMatchIterator::new(&pattern);
iter.set_new_target_complete(target.as_view(), &match_stack);
assert!(iter.next_result(&mut match_stack).is_ok());
assert_eq!(match_stack.len(), 2);
assert_eq!(
match_stack.stack.get(x),
Some(&Match::Single(two.as_view()))
);
assert!(iter.next_result(&mut match_stack).is_err());
assert_eq!(match_stack.len(), 1);
assert_eq!(
match_stack.stack.get(outer),
Some(&Match::Single(outer_value.as_view()))
);
}
#[test]
fn direct_wildcard_and_slice_fallback_share_one_cleanup_boundary() {
let outer_value = parse!("sentinel");
let target = parse!("a*b*c");
let pattern = parse!("x__").to_pattern();
let outer = symbol!("outer_");
let x = symbol!("x__");
let mut match_stack =
WrappedMatchStack::new(&DEFAULT_PATTERN_CONDITION, &DEFAULT_MATCH_SETTINGS);
assert!(
match_stack
.insert(outer, Match::Single(outer_value.as_view()))
.is_ok()
);
let mut iter = AtomMatchIterator::new(&pattern);
iter.set_new_target(target.as_view(), &match_stack);
let mut match_count = 0;
while iter.next_result(&mut match_stack).is_ok() {
match_count += 1;
assert_eq!(match_stack.len(), 2);
let matched = match_stack.stack.get(x).unwrap();
if match_count == 1 {
assert_eq!(matched, &Match::Single(target.as_view()));
} else {
assert_ne!(matched, &Match::Single(target.as_view()));
}
}
assert_eq!(match_count, 4);
assert_eq!(match_stack.len(), 1);
assert_eq!(
match_stack.stack.get(outer),
Some(&Match::Single(outer_value.as_view()))
);
}
#[test]
fn sibling_backtracking_discards_nested_function_bindings() {
let outer_value = parse!("sentinel");
let target = parse!("f(a)*f(b)*g(b)");
let pattern = parse!("h_(x_)*g(x_)").to_pattern();
let outer = symbol!("outer_");
let h = symbol!("h_");
let x = symbol!("x_");
let b = parse!("b");
let mut match_stack =
WrappedMatchStack::new(&DEFAULT_PATTERN_CONDITION, &DEFAULT_MATCH_SETTINGS);
assert!(
match_stack
.insert(outer, Match::Single(outer_value.as_view()))
.is_ok()
);
let mut iter = AtomMatchIterator::new(&pattern);
iter.set_new_target(target.as_view(), &match_stack);
assert!(iter.next_result(&mut match_stack).is_ok());
assert_eq!(match_stack.len(), 3);
assert_eq!(
match_stack.stack.get(h),
Some(&Match::FunctionName(symbol!("f")))
);
assert_eq!(match_stack.stack.get(x), Some(&Match::Single(b.as_view())));
while iter.next_result(&mut match_stack).is_ok() {
assert_eq!(match_stack.len(), 3);
}
assert_eq!(match_stack.len(), 1);
assert_eq!(
match_stack.stack.get(outer),
Some(&Match::Single(outer_value.as_view()))
);
}
#[test]
fn ranged_wildcard_backtracking_keeps_a_stable_stack_shape() {
let _ = symbol!("symbolica::id_test::stack_sym"; Symmetric);
let outer_value = parse!("sentinel");
let target = parse!("symbolica::id_test::stack_sym(a,b,g(c),d,e)");
let pattern = parse!("symbolica::id_test::stack_sym(x__,g(y_),z__)").to_pattern();
let outer = symbol!("outer_");
let x = symbol!("x__");
let y = symbol!("y_");
let z = symbol!("z__");
let mut match_stack =
WrappedMatchStack::new(&DEFAULT_PATTERN_CONDITION, &DEFAULT_MATCH_SETTINGS);
assert!(
match_stack
.insert(outer, Match::Single(outer_value.as_view()))
.is_ok()
);
let mut iter = AtomMatchIterator::new(&pattern);
iter.set_new_target_complete(target.as_view(), &match_stack);
let mut match_count = 0;
while iter.next_result(&mut match_stack).is_ok() {
match_count += 1;
assert_eq!(match_stack.len(), 4);
assert!(match_stack.stack.get(x).is_some());
assert!(match_stack.stack.get(y).is_some());
assert!(match_stack.stack.get(z).is_some());
}
assert!(match_count > 1);
assert_eq!(match_stack.len(), 1);
assert_eq!(
match_stack.stack.get(outer),
Some(&Match::Single(outer_value.as_view()))
);
}
#[test]
fn match_stack_reset_nested_slice_atom_fallback() {
let target = parse!("rubi_int(log(a/(a+b*x))*log(c*x/(a+b*x))^2/(x*(a+b*x)),x)");
let pattern =
parse!("rubi_int(u__*log(v_)*log(e__*(f__*(a__+b__*w_)^p_*(c__+d__*w_)^q_)^r_)^s_,w_)")
.to_pattern()
.set_optional(symbol!("e__"))
.set_optional(symbol!("f__"))
.set_optional(symbol!("a__"))
.set_optional(symbol!("b__"))
.set_optional(symbol!("c__"))
.set_optional(symbol!("d__"))
.set_optional(symbol!("p_"))
.set_optional(symbol!("q_"))
.set_optional(symbol!("r_"))
.set_optional(symbol!("s_"));
let bindings =
parse!("bindings(u__,v_,e__,f__,a__,b__,w_,p_,c__,d__,q_,r_,s_)").to_pattern();
let matched = target.replace(&pattern).iter(&bindings).next();
assert_eq!(
matched,
Some(parse!(
"bindings(1/(x*(a+b*x)),a/(a+b*x),1,c,0,1,x,1,a,b,-1,1,2)"
))
);
}
#[test]
fn optional() {
let pat = parse!("(a_+b_*x)^p_")
.to_pattern()
.set_optional(symbol!("a_"))
.set_optional(symbol!("b_"))
.set_optional(symbol!("p_"));
let rhs = parse!("f(a_,b_,p_)").to_pattern();
let e = parse!("(1+2*x)^3").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(1,2,3)"));
let e = parse!("(1+x)^3").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(1,1,3)"));
let e = parse!("1+2*x").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(1,2,1)"));
let replacements = parse!("1+2*x").replace(&pat).iter(&rhs).collect::<Vec<_>>();
let non_default = parse!("f(1,2,1)");
let default = parse!("1+2*f(0,1,1)");
assert_contains_all(&replacements, &[non_default.clone(), default.clone()]);
assert!(
replacements.iter().position(|r| r == &non_default).unwrap()
< replacements.iter().position(|r| r == &default).unwrap(),
"non-default match should be yielded before optional-default fallback"
);
let e = parse!("1+x").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(1,1,1)"));
let e = parse!("x").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(0,1,1)"));
let pat = parse!("x_*o_").to_pattern().set_optional(symbol!("o_"));
let rhs = parse!("f(x_,o_)").to_pattern();
let e = parse!("x").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(x,1)"));
let pat = parse!("x_^o_").to_pattern().set_optional(symbol!("o_"));
let rhs = parse!("f(x_,o_)").to_pattern();
let e = parse!("x").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(x,1)"));
let pat = parse!("x_+o_").to_pattern().set_optional(symbol!("o_"));
let rhs = parse!("f(x_,o_)").to_pattern();
let e = parse!("x").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(x,0)"));
let pat = parse!("x*y*o_").to_pattern().set_optional(symbol!("o_"));
let rhs = parse!("f(o_)").to_pattern();
let e = parse!("x*y").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(1)"));
let replacements = parse!("x*y*z").replace(&pat).iter(&rhs).collect::<Vec<_>>();
let non_default = parse!("f(z)");
let default = parse!("z*f(1)");
assert_contains_all(&replacements, &[non_default.clone(), default.clone()]);
assert!(
replacements.iter().position(|r| r == &non_default).unwrap()
< replacements.iter().position(|r| r == &default).unwrap(),
"non-default match should be yielded before optional-default fallback"
);
let pat = parse!("(b_+2*x)*n_")
.to_pattern()
.set_optional(symbol!("b_"))
.set_optional(symbol!("n_"));
let rhs = parse!("f(b_,n_)").to_pattern();
let e = parse!("2*x").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f(0,1)"));
let pat = parse!("x*o1_*o2_")
.to_pattern()
.set_optional(symbol!("o1_"))
.set_optional(symbol!("o2_"));
let rhs = parse!("f(o1_,o2_)").to_pattern();
let replacements = parse!("x*y").replace(&pat).iter(&rhs).collect::<Vec<_>>();
assert_contains_all(&replacements, &[parse!("f(y,1)"), parse!("f(1,y)")]);
for replacement in replacements {
let s = replacement.to_string();
assert!(
!s.contains("arg") && !s.contains("()"),
"nonsensical optional wildcard replacement: {s}"
);
}
let replacements = parse!("x").replace(&pat).iter(&rhs).collect::<Vec<_>>();
assert_contains_all(&replacements, &[parse!("f(1,1)")]);
for replacement in replacements {
let s = replacement.to_string();
assert!(
!s.contains("arg") && !s.contains("()"),
"nonsensical both-default optional wildcard replacement: {s}"
);
}
let pat = parse!("f(x___)").to_pattern();
let rhs = parse!("g(x___)").to_pattern();
let e = parse!("f").replace(&pat).with(&rhs);
assert_eq!(e, parse!("f"));
assert!(parse!("f").replace(&pat).iter(&rhs).next().is_none());
let pat = pat.set_optional(symbol!("x___"));
let e = parse!("f").replace(&pat).with(&rhs);
assert_eq!(e, parse!("g()"));
}
}