use smallvec::{SmallVec, smallvec};
use std::fmt;
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ExprId(pub u32);
impl fmt::Debug for ExprId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "e{}", self.0)
}
}
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NumId(pub u32);
impl fmt::Debug for NumId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "n{}", self.0)
}
}
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct SymbolId(pub u32);
impl fmt::Debug for SymbolId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "s{}", self.0)
}
}
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct CtxId(pub u32);
impl fmt::Debug for CtxId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "ctx{}", self.0)
}
}
#[derive(Clone, PartialEq, Eq, Hash)]
pub enum ExprNode {
Num(NumId),
Symbol(SymbolId),
Pi,
E,
ImaginaryUnit,
EulerGamma,
Catalan,
GoldenRatio,
PhysicalConstant(SymbolId, ExprId),
Infinity,
NegInfinity,
ComplexInfinity,
NaN,
Add(SmallVec<[ExprId; 6]>),
Mul(SmallVec<[ExprId; 6]>),
Min(SmallVec<[ExprId; 4]>),
Max(SmallVec<[ExprId; 4]>),
Pow(ExprId, ExprId),
Neg(ExprId),
Floor(ExprId),
Ceiling(ExprId),
Sin(ExprId),
Cos(ExprId),
Tan(ExprId),
Exp(ExprId),
Ln(ExprId),
Abs(ExprId),
Asin(ExprId),
Acos(ExprId),
Atan(ExprId),
Atan2(ExprId, ExprId),
Sinh(ExprId),
Cosh(ExprId),
Tanh(ExprId),
Asinh(ExprId),
Acosh(ExprId),
Atanh(ExprId),
Sign(ExprId),
Heaviside(ExprId),
DiracDelta(ExprId),
Re(ExprId),
Im(ExprId),
Conjugate(ExprId),
Arg(ExprId),
Gamma(ExprId),
LogGamma(ExprId),
Digamma(ExprId),
Erf(ExprId),
Erfc(ExprId),
LambertW(ExprId),
Beta(ExprId, ExprId),
Si(ExprId),
Ci(ExprId),
Ei(ExprId),
Li(ExprId),
Zeta(ExprId),
Polygamma(ExprId, ExprId),
KroneckerDelta(ExprId, ExprId),
Factorial(ExprId),
Binomial(ExprId, ExprId),
BoolTrue,
BoolFalse,
Gt(ExprId, ExprId),
Ge(ExprId, ExprId),
Eq_(ExprId, ExprId),
Ne(ExprId, ExprId),
And(SmallVec<[ExprId; 6]>),
Or(SmallVec<[ExprId; 6]>),
Not(ExprId),
Piecewise(SmallVec<[(ExprId, ExprId); 3]>),
Apply(SymbolId, SmallVec<[ExprId; 2]>),
Derivative(ExprId, ExprId),
Integral(ExprId, ExprId),
DefiniteIntegral(ExprId, ExprId, ExprId, ExprId),
Sum(ExprId, ExprId, ExprId, ExprId),
Product_(ExprId, ExprId, ExprId, ExprId),
Limit(ExprId, ExprId, ExprId),
Series(ExprId, ExprId, ExprId, ExprId),
LaplaceTransform(ExprId, ExprId, ExprId),
InverseLaplaceTransform(ExprId, ExprId, ExprId),
Residue(ExprId, ExprId, ExprId),
RootOf(ExprId, ExprId),
RootSum(ExprId, ExprId, ExprId),
DSolve(ExprId, ExprId, ExprId),
ConditionSet(ExprId, ExprId),
EmptySet,
UniversalSet,
Interval(ExprId, ExprId, u8),
FiniteSet(SmallVec<[ExprId; 4]>),
SetUnion(SmallVec<[ExprId; 4]>),
SetIntersection(SmallVec<[ExprId; 4]>),
SetComplement(ExprId, ExprId),
}
pub const INTERVAL_LEFT_OPEN: u8 = 0x01;
pub const INTERVAL_RIGHT_OPEN: u8 = 0x02;
pub const INTERVAL_BOTH_OPEN: u8 = 0x03;
pub const INTERVAL_BOTH_CLOSED: u8 = 0x00;
impl ExprNode {
pub fn children(&self) -> SmallVec<[ExprId; 6]> {
match self {
ExprNode::Num(_)
| ExprNode::Symbol(_)
| ExprNode::Pi
| ExprNode::E
| ExprNode::ImaginaryUnit
| ExprNode::EulerGamma
| ExprNode::Catalan
| ExprNode::GoldenRatio
| ExprNode::PhysicalConstant(_, _)
| ExprNode::Infinity
| ExprNode::NegInfinity
| ExprNode::ComplexInfinity
| ExprNode::NaN
| ExprNode::BoolTrue
| ExprNode::BoolFalse
| ExprNode::EmptySet
| ExprNode::UniversalSet => smallvec![],
ExprNode::Add(ids) | ExprNode::Mul(ids) | ExprNode::And(ids) | ExprNode::Or(ids) => {
ids.clone()
}
ExprNode::Min(ids) | ExprNode::Max(ids) => ids.iter().copied().collect(),
ExprNode::FiniteSet(ids) | ExprNode::SetUnion(ids) | ExprNode::SetIntersection(ids) => {
ids.iter().copied().collect()
}
ExprNode::Piecewise(pairs) => {
let mut result = SmallVec::new();
for &(val, cond) in pairs {
result.push(val);
result.push(cond);
}
result
}
ExprNode::Pow(a, b)
| ExprNode::Atan2(a, b)
| ExprNode::Binomial(a, b)
| ExprNode::Beta(a, b)
| ExprNode::Polygamma(a, b)
| ExprNode::KroneckerDelta(a, b)
| ExprNode::Gt(a, b)
| ExprNode::Ge(a, b)
| ExprNode::Eq_(a, b)
| ExprNode::Ne(a, b)
| ExprNode::Derivative(a, b)
| ExprNode::Integral(a, b)
| ExprNode::SetComplement(a, b)
| ExprNode::RootOf(a, b)
| ExprNode::ConditionSet(a, b) => {
smallvec![*a, *b]
}
ExprNode::Interval(a, b, _) => {
smallvec![*a, *b]
}
ExprNode::Limit(a, b, c)
| ExprNode::LaplaceTransform(a, b, c)
| ExprNode::InverseLaplaceTransform(a, b, c)
| ExprNode::Residue(a, b, c)
| ExprNode::DSolve(a, b, c)
| ExprNode::RootSum(a, b, c) => smallvec![*a, *b, *c],
ExprNode::Sum(a, b, c, d)
| ExprNode::Product_(a, b, c, d)
| ExprNode::Series(a, b, c, d)
| ExprNode::DefiniteIntegral(a, b, c, d) => {
smallvec![*a, *b, *c, *d]
}
ExprNode::Neg(x)
| ExprNode::Floor(x)
| ExprNode::Ceiling(x)
| ExprNode::Sin(x)
| ExprNode::Cos(x)
| ExprNode::Tan(x)
| ExprNode::Exp(x)
| ExprNode::Ln(x)
| ExprNode::Abs(x)
| ExprNode::Asin(x)
| ExprNode::Acos(x)
| ExprNode::Atan(x)
| ExprNode::Sinh(x)
| ExprNode::Cosh(x)
| ExprNode::Tanh(x)
| ExprNode::Asinh(x)
| ExprNode::Acosh(x)
| ExprNode::Atanh(x)
| ExprNode::Sign(x)
| ExprNode::Heaviside(x)
| ExprNode::DiracDelta(x)
| ExprNode::Re(x)
| ExprNode::Im(x)
| ExprNode::Conjugate(x)
| ExprNode::Arg(x)
| ExprNode::Gamma(x)
| ExprNode::LogGamma(x)
| ExprNode::Digamma(x)
| ExprNode::Erf(x)
| ExprNode::Erfc(x)
| ExprNode::LambertW(x)
| ExprNode::Si(x)
| ExprNode::Ci(x)
| ExprNode::Ei(x)
| ExprNode::Li(x)
| ExprNode::Zeta(x)
| ExprNode::Factorial(x)
| ExprNode::Not(x) => smallvec![*x],
ExprNode::Apply(_, args) => {
args.iter().copied().collect()
}
}
}
#[inline]
pub fn for_each_child(&self, mut f: impl FnMut(ExprId)) {
match self {
ExprNode::Num(_)
| ExprNode::Symbol(_)
| ExprNode::Pi
| ExprNode::E
| ExprNode::ImaginaryUnit
| ExprNode::EulerGamma
| ExprNode::Catalan
| ExprNode::GoldenRatio
| ExprNode::PhysicalConstant(_, _)
| ExprNode::Infinity
| ExprNode::NegInfinity
| ExprNode::ComplexInfinity
| ExprNode::NaN
| ExprNode::BoolTrue
| ExprNode::BoolFalse
| ExprNode::EmptySet
| ExprNode::UniversalSet => {}
ExprNode::Add(ids) | ExprNode::Mul(ids) | ExprNode::And(ids) | ExprNode::Or(ids) => {
for &id in ids {
f(id);
}
}
ExprNode::Min(ids) | ExprNode::Max(ids) => {
for &id in ids {
f(id);
}
}
ExprNode::FiniteSet(ids) | ExprNode::SetUnion(ids) | ExprNode::SetIntersection(ids) => {
for &id in ids {
f(id);
}
}
ExprNode::Piecewise(pairs) => {
for &(val, cond) in pairs {
f(val);
f(cond);
}
}
ExprNode::Pow(a, b)
| ExprNode::Atan2(a, b)
| ExprNode::Binomial(a, b)
| ExprNode::Beta(a, b)
| ExprNode::Polygamma(a, b)
| ExprNode::KroneckerDelta(a, b)
| ExprNode::Gt(a, b)
| ExprNode::Ge(a, b)
| ExprNode::Eq_(a, b)
| ExprNode::Ne(a, b)
| ExprNode::Derivative(a, b)
| ExprNode::Integral(a, b)
| ExprNode::SetComplement(a, b)
| ExprNode::RootOf(a, b)
| ExprNode::ConditionSet(a, b) => {
f(*a);
f(*b);
}
ExprNode::Interval(a, b, _) => {
f(*a);
f(*b);
}
ExprNode::Limit(a, b, c)
| ExprNode::LaplaceTransform(a, b, c)
| ExprNode::InverseLaplaceTransform(a, b, c)
| ExprNode::Residue(a, b, c)
| ExprNode::DSolve(a, b, c)
| ExprNode::RootSum(a, b, c) => {
f(*a);
f(*b);
f(*c);
}
ExprNode::Sum(a, b, c, d)
| ExprNode::Product_(a, b, c, d)
| ExprNode::Series(a, b, c, d)
| ExprNode::DefiniteIntegral(a, b, c, d) => {
f(*a);
f(*b);
f(*c);
f(*d);
}
ExprNode::Neg(x)
| ExprNode::Floor(x)
| ExprNode::Ceiling(x)
| ExprNode::Sin(x)
| ExprNode::Cos(x)
| ExprNode::Tan(x)
| ExprNode::Exp(x)
| ExprNode::Ln(x)
| ExprNode::Abs(x)
| ExprNode::Asin(x)
| ExprNode::Acos(x)
| ExprNode::Atan(x)
| ExprNode::Sinh(x)
| ExprNode::Cosh(x)
| ExprNode::Tanh(x)
| ExprNode::Asinh(x)
| ExprNode::Acosh(x)
| ExprNode::Atanh(x)
| ExprNode::Sign(x)
| ExprNode::Heaviside(x)
| ExprNode::DiracDelta(x)
| ExprNode::Re(x)
| ExprNode::Im(x)
| ExprNode::Conjugate(x)
| ExprNode::Arg(x)
| ExprNode::Gamma(x)
| ExprNode::LogGamma(x)
| ExprNode::Digamma(x)
| ExprNode::Erf(x)
| ExprNode::Erfc(x)
| ExprNode::LambertW(x)
| ExprNode::Si(x)
| ExprNode::Ci(x)
| ExprNode::Ei(x)
| ExprNode::Li(x)
| ExprNode::Zeta(x)
| ExprNode::Factorial(x)
| ExprNode::Not(x) => f(*x),
ExprNode::Apply(_, args) => {
for &id in args {
f(id);
}
}
}
}
#[inline]
pub fn child_count(&self) -> usize {
match self {
ExprNode::Num(_)
| ExprNode::Symbol(_)
| ExprNode::Pi
| ExprNode::E
| ExprNode::ImaginaryUnit
| ExprNode::EulerGamma
| ExprNode::Catalan
| ExprNode::GoldenRatio
| ExprNode::PhysicalConstant(_, _)
| ExprNode::Infinity
| ExprNode::NegInfinity
| ExprNode::ComplexInfinity
| ExprNode::NaN
| ExprNode::BoolTrue
| ExprNode::BoolFalse
| ExprNode::EmptySet
| ExprNode::UniversalSet => 0,
ExprNode::Add(ids) | ExprNode::Mul(ids) | ExprNode::And(ids) | ExprNode::Or(ids) => {
ids.len()
}
ExprNode::Min(ids) | ExprNode::Max(ids) => ids.len(),
ExprNode::FiniteSet(ids) | ExprNode::SetUnion(ids) | ExprNode::SetIntersection(ids) => {
ids.len()
}
ExprNode::Piecewise(pairs) => pairs.len() * 2,
ExprNode::Pow(..)
| ExprNode::Atan2(..)
| ExprNode::Binomial(..)
| ExprNode::Beta(..)
| ExprNode::Polygamma(..)
| ExprNode::KroneckerDelta(..)
| ExprNode::Gt(..)
| ExprNode::Ge(..)
| ExprNode::Eq_(..)
| ExprNode::Ne(..)
| ExprNode::Derivative(..)
| ExprNode::Integral(..)
| ExprNode::SetComplement(..)
| ExprNode::Interval(..)
| ExprNode::RootOf(..)
| ExprNode::ConditionSet(..) => 2,
ExprNode::Limit(..)
| ExprNode::LaplaceTransform(..)
| ExprNode::InverseLaplaceTransform(..)
| ExprNode::Residue(..)
| ExprNode::DSolve(..)
| ExprNode::RootSum(..) => 3,
ExprNode::Sum(..)
| ExprNode::Product_(..)
| ExprNode::Series(..)
| ExprNode::DefiniteIntegral(..) => 4,
ExprNode::Neg(_)
| ExprNode::Floor(_)
| ExprNode::Ceiling(_)
| ExprNode::Sin(_)
| ExprNode::Cos(_)
| ExprNode::Tan(_)
| ExprNode::Exp(_)
| ExprNode::Ln(_)
| ExprNode::Abs(_)
| ExprNode::Asin(_)
| ExprNode::Acos(_)
| ExprNode::Atan(_)
| ExprNode::Sinh(_)
| ExprNode::Cosh(_)
| ExprNode::Tanh(_)
| ExprNode::Asinh(_)
| ExprNode::Acosh(_)
| ExprNode::Atanh(_)
| ExprNode::Sign(_)
| ExprNode::Heaviside(_)
| ExprNode::DiracDelta(_)
| ExprNode::Re(_)
| ExprNode::Im(_)
| ExprNode::Conjugate(_)
| ExprNode::Arg(_)
| ExprNode::Gamma(_)
| ExprNode::LogGamma(_)
| ExprNode::Digamma(_)
| ExprNode::Erf(_)
| ExprNode::Erfc(_)
| ExprNode::LambertW(_)
| ExprNode::Si(_)
| ExprNode::Ci(_)
| ExprNode::Ei(_)
| ExprNode::Li(_)
| ExprNode::Zeta(_)
| ExprNode::Factorial(_)
| ExprNode::Not(_) => 1,
ExprNode::Apply(_, args) => args.len(),
}
}
pub fn is_atom(&self) -> bool {
matches!(
self,
ExprNode::Num(_)
| ExprNode::Symbol(_)
| ExprNode::Pi
| ExprNode::E
| ExprNode::ImaginaryUnit
| ExprNode::EulerGamma
| ExprNode::Catalan
| ExprNode::GoldenRatio
| ExprNode::PhysicalConstant(_, _)
| ExprNode::Infinity
| ExprNode::NegInfinity
| ExprNode::ComplexInfinity
| ExprNode::NaN
| ExprNode::BoolTrue
| ExprNode::BoolFalse
| ExprNode::EmptySet
| ExprNode::UniversalSet
)
}
pub fn is_set_node(&self) -> bool {
matches!(
self,
ExprNode::EmptySet
| ExprNode::UniversalSet
| ExprNode::Interval(..)
| ExprNode::FiniteSet(_)
| ExprNode::SetUnion(_)
| ExprNode::SetIntersection(_)
| ExprNode::SetComplement(..)
)
}
}
impl fmt::Debug for ExprNode {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ExprNode::Num(id) => write!(f, "Num({id:?})"),
ExprNode::Symbol(id) => write!(f, "Symbol({id:?})"),
ExprNode::Pi => write!(f, "Pi"),
ExprNode::E => write!(f, "E"),
ExprNode::ImaginaryUnit => write!(f, "ImaginaryUnit"),
ExprNode::EulerGamma => write!(f, "EulerGamma"),
ExprNode::Catalan => write!(f, "Catalan"),
ExprNode::GoldenRatio => write!(f, "GoldenRatio"),
ExprNode::PhysicalConstant(name, val) => {
write!(f, "PhysicalConstant({name:?}, {val:?})")
}
ExprNode::Infinity => write!(f, "Infinity"),
ExprNode::NegInfinity => write!(f, "NegInfinity"),
ExprNode::ComplexInfinity => write!(f, "ComplexInfinity"),
ExprNode::NaN => write!(f, "NaN"),
ExprNode::Add(ids) => f.debug_tuple("Add").field(ids).finish(),
ExprNode::Mul(ids) => f.debug_tuple("Mul").field(ids).finish(),
ExprNode::Pow(base, exp) => f.debug_tuple("Pow").field(base).field(exp).finish(),
ExprNode::Neg(x) => f.debug_tuple("Neg").field(x).finish(),
ExprNode::Sin(x) => f.debug_tuple("Sin").field(x).finish(),
ExprNode::Cos(x) => f.debug_tuple("Cos").field(x).finish(),
ExprNode::Tan(x) => f.debug_tuple("Tan").field(x).finish(),
ExprNode::Exp(x) => f.debug_tuple("Exp").field(x).finish(),
ExprNode::Ln(x) => f.debug_tuple("Ln").field(x).finish(),
ExprNode::Abs(x) => f.debug_tuple("Abs").field(x).finish(),
ExprNode::Asin(x) => f.debug_tuple("Asin").field(x).finish(),
ExprNode::Acos(x) => f.debug_tuple("Acos").field(x).finish(),
ExprNode::Atan(x) => f.debug_tuple("Atan").field(x).finish(),
ExprNode::Atan2(y, x) => f.debug_tuple("Atan2").field(y).field(x).finish(),
ExprNode::Sinh(x) => f.debug_tuple("Sinh").field(x).finish(),
ExprNode::Cosh(x) => f.debug_tuple("Cosh").field(x).finish(),
ExprNode::Tanh(x) => f.debug_tuple("Tanh").field(x).finish(),
ExprNode::Asinh(x) => f.debug_tuple("Asinh").field(x).finish(),
ExprNode::Acosh(x) => f.debug_tuple("Acosh").field(x).finish(),
ExprNode::Atanh(x) => f.debug_tuple("Atanh").field(x).finish(),
ExprNode::Floor(x) => f.debug_tuple("Floor").field(x).finish(),
ExprNode::Ceiling(x) => f.debug_tuple("Ceiling").field(x).finish(),
ExprNode::Min(ids) => f.debug_tuple("Min").field(ids).finish(),
ExprNode::Max(ids) => f.debug_tuple("Max").field(ids).finish(),
ExprNode::Sign(id) => write!(f, "Sign({id:?})"),
ExprNode::Heaviside(x) => f.debug_tuple("Heaviside").field(x).finish(),
ExprNode::DiracDelta(x) => f.debug_tuple("DiracDelta").field(x).finish(),
ExprNode::Re(x) => f.debug_tuple("Re").field(x).finish(),
ExprNode::Im(x) => f.debug_tuple("Im").field(x).finish(),
ExprNode::Conjugate(x) => f.debug_tuple("Conjugate").field(x).finish(),
ExprNode::Arg(x) => f.debug_tuple("Arg").field(x).finish(),
ExprNode::Gamma(x) => f.debug_tuple("Gamma").field(x).finish(),
ExprNode::LogGamma(x) => f.debug_tuple("LogGamma").field(x).finish(),
ExprNode::Digamma(x) => f.debug_tuple("Digamma").field(x).finish(),
ExprNode::Erf(x) => f.debug_tuple("Erf").field(x).finish(),
ExprNode::Erfc(x) => f.debug_tuple("Erfc").field(x).finish(),
ExprNode::LambertW(x) => f.debug_tuple("LambertW").field(x).finish(),
ExprNode::Beta(a, b) => f.debug_tuple("Beta").field(a).field(b).finish(),
ExprNode::Si(x) => f.debug_tuple("Si").field(x).finish(),
ExprNode::Ci(x) => f.debug_tuple("Ci").field(x).finish(),
ExprNode::Ei(x) => f.debug_tuple("Ei").field(x).finish(),
ExprNode::Li(x) => f.debug_tuple("Li").field(x).finish(),
ExprNode::Zeta(x) => f.debug_tuple("Zeta").field(x).finish(),
ExprNode::Polygamma(n, x) => f.debug_tuple("Polygamma").field(n).field(x).finish(),
ExprNode::KroneckerDelta(i, j) => {
f.debug_tuple("KroneckerDelta").field(i).field(j).finish()
}
ExprNode::Factorial(id) => write!(f, "Factorial({id:?})"),
ExprNode::Binomial(n, k) => write!(f, "Binomial({n:?}, {k:?})"),
ExprNode::BoolTrue => write!(f, "BoolTrue"),
ExprNode::BoolFalse => write!(f, "BoolFalse"),
ExprNode::Gt(a, b) => f.debug_tuple("Gt").field(a).field(b).finish(),
ExprNode::Ge(a, b) => f.debug_tuple("Ge").field(a).field(b).finish(),
ExprNode::Eq_(a, b) => f.debug_tuple("Eq_").field(a).field(b).finish(),
ExprNode::Ne(a, b) => f.debug_tuple("Ne").field(a).field(b).finish(),
ExprNode::And(ids) => f.debug_tuple("And").field(ids).finish(),
ExprNode::Or(ids) => f.debug_tuple("Or").field(ids).finish(),
ExprNode::Not(x) => f.debug_tuple("Not").field(x).finish(),
ExprNode::Piecewise(pairs) => {
let mut d = f.debug_tuple("Piecewise");
for pair in pairs {
d.field(pair);
}
d.finish()
}
ExprNode::Apply(sym, args) => f.debug_tuple("Apply").field(sym).field(args).finish(),
ExprNode::Derivative(body, var) => {
f.debug_tuple("Derivative").field(body).field(var).finish()
}
ExprNode::Integral(body, var) => {
f.debug_tuple("Integral").field(body).field(var).finish()
}
ExprNode::DefiniteIntegral(body, var, lo, hi) => f
.debug_tuple("DefiniteIntegral")
.field(body)
.field(var)
.field(lo)
.field(hi)
.finish(),
ExprNode::Sum(body, var, lo, hi) => f
.debug_tuple("Sum")
.field(body)
.field(var)
.field(lo)
.field(hi)
.finish(),
ExprNode::Product_(body, var, lo, hi) => f
.debug_tuple("Product_")
.field(body)
.field(var)
.field(lo)
.field(hi)
.finish(),
ExprNode::EmptySet => write!(f, "EmptySet"),
ExprNode::UniversalSet => write!(f, "UniversalSet"),
ExprNode::Interval(a, b, flags) => f
.debug_tuple("Interval")
.field(a)
.field(b)
.field(flags)
.finish(),
ExprNode::FiniteSet(ids) => f.debug_tuple("FiniteSet").field(ids).finish(),
ExprNode::SetUnion(ids) => f.debug_tuple("SetUnion").field(ids).finish(),
ExprNode::SetIntersection(ids) => f.debug_tuple("SetIntersection").field(ids).finish(),
ExprNode::SetComplement(a, b) => {
f.debug_tuple("SetComplement").field(a).field(b).finish()
}
ExprNode::Limit(body, var, point) => f
.debug_tuple("Limit")
.field(body)
.field(var)
.field(point)
.finish(),
ExprNode::Series(body, var, point, order) => f
.debug_tuple("Series")
.field(body)
.field(var)
.field(point)
.field(order)
.finish(),
ExprNode::LaplaceTransform(body, t, s) => f
.debug_tuple("LaplaceTransform")
.field(body)
.field(t)
.field(s)
.finish(),
ExprNode::InverseLaplaceTransform(body, s, t) => f
.debug_tuple("InverseLaplaceTransform")
.field(body)
.field(s)
.field(t)
.finish(),
ExprNode::Residue(body, var, point) => f
.debug_tuple("Residue")
.field(body)
.field(var)
.field(point)
.finish(),
ExprNode::RootOf(poly, idx) => f.debug_tuple("RootOf").field(poly).field(idx).finish(),
ExprNode::DSolve(expr, func, var) => f
.debug_tuple("DSolve")
.field(expr)
.field(func)
.field(var)
.finish(),
ExprNode::RootSum(poly, body, sumvar) => f
.debug_tuple("RootSum")
.field(poly)
.field(body)
.field(sumvar)
.finish(),
ExprNode::ConditionSet(var, cond) => f
.debug_tuple("ConditionSet")
.field(var)
.field(cond)
.finish(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn debug_expr_id() {
let id = ExprId(42);
assert_eq!(format!("{id:?}"), "e42");
}
#[test]
fn debug_num_id() {
let id = NumId(7);
assert_eq!(format!("{id:?}"), "n7");
}
#[test]
fn debug_symbol_id() {
let id = SymbolId(3);
assert_eq!(format!("{id:?}"), "s3");
}
#[test]
fn debug_ctx_id() {
let id = CtxId(0);
assert_eq!(format!("{id:?}"), "ctx0");
}
#[test]
fn atom_has_no_children() {
let node = ExprNode::Pi;
assert!(node.is_atom());
assert!(node.children().is_empty());
}
#[test]
fn num_is_atom() {
let node = ExprNode::Num(NumId(0));
assert!(node.is_atom());
assert!(node.children().is_empty());
}
#[test]
fn symbol_is_atom() {
let node = ExprNode::Symbol(SymbolId(1));
assert!(node.is_atom());
}
#[test]
fn special_values_are_atoms() {
for node in [
ExprNode::Infinity,
ExprNode::NegInfinity,
ExprNode::ComplexInfinity,
ExprNode::NaN,
ExprNode::ImaginaryUnit,
ExprNode::E,
] {
assert!(node.is_atom(), "{node:?} should be an atom");
assert!(node.children().is_empty());
}
}
#[test]
fn add_children() {
let ids: SmallVec<[ExprId; 6]> = smallvec![ExprId(1), ExprId(2), ExprId(3)];
let node = ExprNode::Add(ids.clone());
assert!(!node.is_atom());
assert_eq!(node.children(), ids);
}
#[test]
fn mul_children() {
let ids: SmallVec<[ExprId; 6]> = smallvec![ExprId(4), ExprId(5)];
let node = ExprNode::Mul(ids.clone());
assert!(!node.is_atom());
assert_eq!(node.children(), ids);
}
#[test]
fn pow_children() {
let node = ExprNode::Pow(ExprId(10), ExprId(20));
assert!(!node.is_atom());
let kids = node.children();
assert_eq!(kids.len(), 2);
assert_eq!(kids[0], ExprId(10));
assert_eq!(kids[1], ExprId(20));
}
#[test]
fn unary_children() {
for node in [
ExprNode::Neg(ExprId(1)),
ExprNode::Sin(ExprId(1)),
ExprNode::Cos(ExprId(1)),
ExprNode::Tan(ExprId(1)),
ExprNode::Exp(ExprId(1)),
ExprNode::Ln(ExprId(1)),
ExprNode::Abs(ExprId(1)),
] {
assert!(!node.is_atom());
let kids = node.children();
assert_eq!(kids.len(), 1, "{node:?} should have exactly 1 child");
assert_eq!(kids[0], ExprId(1));
}
}
#[test]
fn apply_children() {
let args: SmallVec<[ExprId; 2]> = smallvec![ExprId(5), ExprId(6)];
let node = ExprNode::Apply(SymbolId(0), args);
assert!(!node.is_atom());
let kids = node.children();
assert_eq!(kids.len(), 2);
assert_eq!(kids[0], ExprId(5));
assert_eq!(kids[1], ExprId(6));
}
#[test]
fn derivative_children() {
let node = ExprNode::Derivative(ExprId(3), ExprId(7));
assert!(!node.is_atom());
let kids = node.children();
assert_eq!(kids.len(), 2);
assert_eq!(kids[0], ExprId(3));
assert_eq!(kids[1], ExprId(7));
}
#[test]
fn integral_children() {
let node = ExprNode::Integral(ExprId(8), ExprId(9));
assert!(!node.is_atom());
let kids = node.children();
assert_eq!(kids.len(), 2);
assert_eq!(kids[0], ExprId(8));
assert_eq!(kids[1], ExprId(9));
}
#[test]
fn named_constants_are_atoms() {
for node in [
ExprNode::EulerGamma,
ExprNode::Catalan,
ExprNode::GoldenRatio,
] {
assert!(node.is_atom(), "{node:?} should be an atom");
assert!(node.children().is_empty());
assert_eq!(node.child_count(), 0);
let mut count = 0;
node.for_each_child(|_| count += 1);
assert_eq!(count, 0);
}
assert_eq!(format!("{:?}", ExprNode::EulerGamma), "EulerGamma");
assert_eq!(format!("{:?}", ExprNode::Catalan), "Catalan");
assert_eq!(format!("{:?}", ExprNode::GoldenRatio), "GoldenRatio");
}
#[test]
fn complex_and_special_unary_children() {
for node in [
ExprNode::Re(ExprId(3)),
ExprNode::Im(ExprId(3)),
ExprNode::Conjugate(ExprId(3)),
ExprNode::Arg(ExprId(3)),
ExprNode::Si(ExprId(3)),
ExprNode::Ci(ExprId(3)),
ExprNode::Ei(ExprId(3)),
ExprNode::Li(ExprId(3)),
ExprNode::Zeta(ExprId(3)),
] {
assert!(!node.is_atom());
assert_eq!(node.child_count(), 1);
let kids = node.children();
assert_eq!(kids.len(), 1, "{node:?} should have exactly 1 child");
assert_eq!(kids[0], ExprId(3));
let mut seen = Vec::new();
node.for_each_child(|c| seen.push(c));
assert_eq!(seen, vec![ExprId(3)]);
}
}
#[test]
fn polygamma_and_kronecker_children() {
for node in [
ExprNode::Polygamma(ExprId(1), ExprId(2)),
ExprNode::KroneckerDelta(ExprId(1), ExprId(2)),
] {
assert!(!node.is_atom());
assert_eq!(node.child_count(), 2);
let kids = node.children();
assert_eq!(kids.len(), 2);
assert_eq!(kids[0], ExprId(1));
assert_eq!(kids[1], ExprId(2));
let mut seen = Vec::new();
node.for_each_child(|c| seen.push(c));
assert_eq!(seen, vec![ExprId(1), ExprId(2)]);
}
assert_eq!(
format!("{:?}", ExprNode::Polygamma(ExprId(1), ExprId(2))),
"Polygamma(e1, e2)"
);
assert_eq!(
format!("{:?}", ExprNode::KroneckerDelta(ExprId(1), ExprId(2))),
"KroneckerDelta(e1, e2)"
);
}
}