use crate::base::arena::Arena;
use crate::base::interval::Interval;
use crate::base::node::{ExprId, ExprNode, INTERVAL_BOTH_OPEN, SymbolId};
use crate::base::walk;
use crate::transforms::eval;
use crate::transforms::evalf;
use crate::transforms::inequalities::Relation;
use crate::transforms::solve;
#[allow(dead_code)]
pub(crate) fn continuous_domain(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
_var_sym: SymbolId,
domain: ExprId,
) -> ExprId {
tracing::debug!("continuous_domain: starting domain analysis");
let post_order = walk::post_order_ids(arena, expr);
let mut valid = domain;
for &id in &post_order {
let node = arena.node(id).clone();
let constraint = compute_node_constraint(arena, &node, var);
if let Some(c) = constraint {
valid = arena.set_intersection(&[valid, c]);
}
}
valid
}
#[allow(dead_code)]
fn compute_node_constraint(arena: &mut Arena, node: &ExprNode, var: ExprId) -> Option<ExprId> {
match node {
ExprNode::Pow(base, exp) => {
let base = *base;
let exp = *exp;
if !walk::contains(arena, base, var) {
return None;
}
if let Some(r) = arena.as_num(exp).cloned() {
use num_traits::Signed;
if r.is_negative() {
return Some(domain_exclude_zeros(arena, base, var));
}
let half = num_rational::Ratio::new(
num_bigint::BigInt::from(1),
num_bigint::BigInt::from(2),
);
if r == half {
return solve_ge_zero(arena, base, var);
}
}
if let ExprNode::Neg(_) = arena.node(exp)
&& walk::contains(arena, base, var)
{
return Some(domain_exclude_zeros(arena, base, var));
}
None
}
ExprNode::Ln(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return None;
}
solve_gt_zero(arena, inner, var)
}
ExprNode::Tan(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return None;
}
let cos_inner = arena.cos(inner);
Some(domain_exclude_zeros(arena, cos_inner, var))
}
ExprNode::Asin(inner) | ExprNode::Acos(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return None;
}
let neg_one = arena.neg_one;
let one = arena.one;
let shifted_low = arena.sub(inner, neg_one); let set_low = solve_ge_zero(arena, shifted_low, var);
let shifted_high = arena.sub(one, inner);
let set_high = solve_ge_zero(arena, shifted_high, var);
match (set_low, set_high) {
(Some(a), Some(b)) => Some(arena.set_intersection(&[a, b])),
(Some(a), None) => Some(a),
(None, Some(b)) => Some(b),
(None, None) => None,
}
}
ExprNode::Acosh(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return None;
}
let one = arena.one;
let shifted = arena.sub(inner, one); solve_ge_zero(arena, shifted, var)
}
ExprNode::Atanh(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return None;
}
let neg_one = arena.neg_one;
let one = arena.one;
let shifted_low = arena.sub(inner, neg_one);
let set_low = solve_gt_zero(arena, shifted_low, var);
let shifted_high = arena.sub(one, inner);
let set_high = solve_gt_zero(arena, shifted_high, var);
match (set_low, set_high) {
(Some(a), Some(b)) => Some(arena.set_intersection(&[a, b])),
(Some(a), None) => Some(a),
(None, Some(b)) => Some(b),
(None, None) => None,
}
}
_ => None,
}
}
#[allow(dead_code)]
fn solve_gt_zero(arena: &mut Arena, expr: ExprId, var: ExprId) -> Option<ExprId> {
arena.solve_inequality_expr(expr, var, Relation::Gt).ok()
}
#[allow(dead_code)]
fn solve_ge_zero(arena: &mut Arena, expr: ExprId, var: ExprId) -> Option<ExprId> {
arena.solve_inequality_expr(expr, var, Relation::Ge).ok()
}
#[allow(dead_code)]
fn domain_exclude_zeros(arena: &mut Arena, expr: ExprId, var: ExprId) -> ExprId {
let solutions = solve::solve(arena, expr, var);
if solutions.is_empty() {
return arena.interval(arena.neg_infinity, arena.infinity, INTERVAL_BOTH_OPEN);
}
let root_ids: Vec<ExprId> = solutions.into_iter().map(|s| s.value).collect();
let roots_set = arena.finite_set(&root_ids);
let reals = arena.interval(arena.neg_infinity, arena.infinity, INTERVAL_BOTH_OPEN);
arena.set_complement(reals, roots_set)
}
pub(crate) fn singularities(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
_var_sym: SymbolId,
range: Interval<f64>,
) -> Vec<f64> {
tracing::debug!(
range_lo = range.lower,
range_hi = range.upper,
"singularities: scanning for singularities"
);
let scan = scan_breakpoints(arena, expr, var, Some(range));
let mut sing_points: Vec<f64> = Vec::new();
for bp in scan.points {
if bp.kind != BreakKind::Singular {
continue;
}
if let Some(val) = bp.value
&& val.is_finite()
&& val >= range.lower
&& val <= range.upper
{
if !sing_points.iter().any(|&v| (v - val).abs() < 1e-12) {
sing_points.push(val);
}
}
}
sing_points.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
sing_points
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum BreakKind {
Singular,
Kink,
}
#[derive(Clone, Debug)]
pub(crate) struct Breakpoint {
pub point: ExprId,
pub value: Option<f64>,
pub kind: BreakKind,
}
#[derive(Clone, Debug, Default)]
pub(crate) struct BreakScan {
pub points: Vec<Breakpoint>,
pub complete: bool,
pub opaque: bool,
}
pub(crate) fn scan_breakpoints(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
range: Option<Interval<f64>>,
) -> BreakScan {
let post_order = walk::post_order_ids(arena, expr);
let mut scan = BreakScan {
points: Vec::new(),
complete: true,
opaque: false,
};
for &id in &post_order {
let node = arena.node(id).clone();
breakpoints_of_node(arena, &node, var, range, &mut scan);
}
scan
}
fn push_zeros(
arena: &mut Arena,
g: ExprId,
var: ExprId,
kind: BreakKind,
range: Option<Interval<f64>>,
scan: &mut BreakScan,
) {
if !walk::contains(arena, g, var) {
return;
}
if let Some(done) = push_trig_zeros(arena, g, var, kind, range, scan) {
if !done {
scan.complete = false;
}
return;
}
let solutions = solve::solve(arena, g, var);
if solutions.is_empty() {
if crate::poly::polybridge::expr_to_poly(arena, g, var).is_none() {
scan.complete = false;
}
return;
}
for s in solutions {
push_point(arena, s.value, kind, scan);
}
}
fn push_point(arena: &mut Arena, point: ExprId, kind: BreakKind, scan: &mut BreakScan) {
if walk::free_symbols(arena, point).is_empty() {
match expr_to_f64(arena, point) {
Some(v) if v.is_finite() => scan.points.push(Breakpoint {
point,
value: Some(v),
kind,
}),
_ => {}
}
} else {
let mut cache = crate::base::assumptions::AssumptionCache::new();
if cache.query(arena, point, crate::base::assumptions::Props::REAL) == Some(false) {
return;
}
let mut nonreal = false;
let mut stack = vec![point];
while let Some(id) = stack.pop() {
match arena.node(id).clone() {
ExprNode::ImaginaryUnit => {
nonreal = true;
break;
}
ExprNode::Pow(b, e) => {
if let Some(r) = arena.as_num(e)
&& !r.is_integer()
&& num_integer::Integer::is_even(r.denom())
&& cache.query(arena, b, crate::base::assumptions::Props::NEGATIVE)
== Some(true)
{
nonreal = true;
break;
}
stack.push(b);
stack.push(e);
}
node => node.for_each_child(|c| stack.push(c)),
}
}
if nonreal && cache.query(arena, point, crate::base::assumptions::Props::REAL) != Some(true)
{
return;
}
scan.points.push(Breakpoint {
point,
value: None,
kind,
});
}
}
fn push_trig_zeros(
arena: &mut Arena,
g: ExprId,
var: ExprId,
kind: BreakKind,
range: Option<Interval<f64>>,
scan: &mut BreakScan,
) -> Option<bool> {
let (inner, offset) = match arena.node(g).clone() {
ExprNode::Sin(i) | ExprNode::Tan(i) => (i, 0.0),
ExprNode::Cos(i) => (i, std::f64::consts::FRAC_PI_2),
_ => return None,
};
if !walk::contains(arena, inner, var) {
return None;
}
let LinearCoeffs {
slope: alpha,
intercept: beta,
} = match linear_coeffs_f64(arena, inner, var) {
Some(ab) => ab,
None => return Some(false),
};
let (lo, hi) = match range {
Some(r) => (r.lower, r.upper),
None => return Some(false),
};
if alpha == 0.0 || !lo.is_finite() || !hi.is_finite() {
return Some(false);
}
let t_lo = alpha * lo + beta;
let t_hi = alpha * hi + beta;
let (t_min, t_max) = if t_lo <= t_hi {
(t_lo, t_hi)
} else {
(t_hi, t_lo)
};
let k_min = ((t_min - offset) / std::f64::consts::PI).floor() as i64 - 1;
let k_max = ((t_max - offset) / std::f64::consts::PI).ceil() as i64 + 1;
if k_max - k_min > 10_000 {
return Some(false);
}
let LinearCoeffs {
slope: alpha_ex,
intercept: beta_ex,
} = match linear_coeffs_exact(arena, inner, var) {
Some(ab) => ab,
None => return Some(false),
};
for k in k_min..=k_max {
let t = offset + (k as f64) * std::f64::consts::PI;
let xv = (t - beta) / alpha;
if xv < lo - 1e-12 || xv > hi + 1e-12 {
continue;
}
let k_ex = if offset == 0.0 {
arena.int(k)
} else {
arena.rational(2 * k + 1, 2)
};
let pi = arena.pi();
let k_pi = arena.mul(&[k_ex, pi]);
let num = arena.sub(k_pi, beta_ex);
let point = arena.div(num, alpha_ex);
scan.points.push(Breakpoint {
point,
value: Some(xv),
kind,
});
}
Some(true)
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub(crate) struct LinearCoeffs<T> {
pub slope: T,
pub intercept: T,
}
fn linear_coeffs_f64(arena: &mut Arena, expr: ExprId, var: ExprId) -> Option<LinearCoeffs<f64>> {
let lin = linear_coeffs_exact(arena, expr, var)?;
let af = expr_to_f64(arena, lin.slope)?;
let bf = expr_to_f64(arena, lin.intercept)?;
if af.is_finite() && bf.is_finite() {
Some(LinearCoeffs {
slope: af,
intercept: bf,
})
} else {
None
}
}
pub(crate) fn linear_coeffs_exact(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
) -> Option<LinearCoeffs<ExprId>> {
let coeffs = poly_coeffs_symbolic(arena, expr, var)?;
if coeffs.len() != 2 {
return None;
}
Some(LinearCoeffs {
slope: coeffs[1],
intercept: coeffs[0],
})
}
pub(crate) fn poly_coeffs_symbolic(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
) -> Option<Vec<ExprId>> {
let expanded = crate::transforms::expand::expand(arena, expr);
let expanded = eval::eval(arena, expanded);
let terms: Vec<ExprId> = match arena.node(expanded).clone() {
ExprNode::Add(children) => children.to_vec(),
_ => vec![expanded],
};
let mut buckets: Vec<Vec<ExprId>> = Vec::new();
for term in terms {
let (k, coeff) = split_term_power(arena, term, var)?;
while buckets.len() <= k {
buckets.push(Vec::new());
}
buckets[k].push(coeff);
}
if buckets.is_empty() {
return None;
}
let mut result = Vec::with_capacity(buckets.len());
for bucket in buckets {
let c = if bucket.is_empty() {
arena.zero()
} else {
let s = arena.add(&bucket);
eval::eval(arena, s)
};
result.push(c);
}
while result.len() > 1 && arena.is_zero_structural(*result.last()?) {
result.pop();
}
if result.len() == 1 && arena.is_zero_structural(result[0]) {
return None;
}
Some(result)
}
fn split_term_power(arena: &mut Arena, term: ExprId, var: ExprId) -> Option<(usize, ExprId)> {
if !walk::contains(arena, term, var) {
return Some((0, term));
}
if term == var {
return Some((1, arena.one()));
}
match arena.node(term).clone() {
ExprNode::Pow(base, exp) if base == var => {
let r = arena.as_num(exp)?.clone();
if !r.is_integer() {
return None;
}
use num_traits::{Signed, ToPrimitive};
if r.is_negative() {
return None;
}
let k = r.to_integer().to_usize()?;
Some((k, arena.one()))
}
ExprNode::Mul(children) => {
let mut k_total = 0usize;
let mut coeff_parts: Vec<ExprId> = Vec::new();
for c in children.iter() {
if !walk::contains(arena, *c, var) {
coeff_parts.push(*c);
continue;
}
let (k, cc) = split_term_power(arena, *c, var)?;
if !arena.is_one_structural(cc) {
return None;
}
k_total += k;
}
let coeff = if coeff_parts.is_empty() {
arena.one()
} else {
arena.mul(&coeff_parts)
};
Some((k_total, coeff))
}
ExprNode::Neg(inner) => {
let (k, c) = split_term_power(arena, inner, var)?;
Some((k, arena.neg(c)))
}
_ => None,
}
}
fn breakpoints_of_node(
arena: &mut Arena,
node: &ExprNode,
var: ExprId,
range: Option<Interval<f64>>,
scan: &mut BreakScan,
) {
use BreakKind::{Kink, Singular};
match node {
ExprNode::Pow(base, exp) => {
let base = *base;
let exp = *exp;
if !walk::contains(arena, base, var) {
return;
}
if walk::contains(arena, exp, var) {
scan.complete = false;
return;
}
let matters = if let Some(r) = arena.as_num(exp).cloned() {
use num_traits::Signed;
r.is_negative() || !r.is_integer()
} else {
true
};
if matters {
push_zeros(arena, base, var, Singular, range, scan);
}
}
ExprNode::Ln(inner) => push_zeros(arena, *inner, var, Singular, range, scan),
ExprNode::Tan(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return;
}
let cos_inner = arena.cos(inner);
push_zeros(arena, cos_inner, var, Singular, range, scan);
}
ExprNode::Asin(inner) | ExprNode::Acos(inner) | ExprNode::Atanh(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return;
}
let one = arena.one();
let g_minus = arena.sub(inner, one);
let g_plus = arena.add(&[inner, one]);
push_zeros(arena, g_minus, var, Singular, range, scan);
push_zeros(arena, g_plus, var, Singular, range, scan);
}
ExprNode::Acosh(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return;
}
let one = arena.one();
let g_minus = arena.sub(inner, one);
push_zeros(arena, g_minus, var, Singular, range, scan);
}
ExprNode::Gamma(inner) | ExprNode::Digamma(inner) | ExprNode::LogGamma(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return;
}
push_integer_family(arena, inner, var, range, Singular, scan, |n| n <= 0);
}
ExprNode::Abs(inner)
| ExprNode::Sign(inner)
| ExprNode::Heaviside(inner)
| ExprNode::DiracDelta(inner) => push_zeros(arena, *inner, var, Kink, range, scan),
ExprNode::Floor(inner) | ExprNode::Ceiling(inner) => {
let inner = *inner;
if !walk::contains(arena, inner, var) {
return;
}
push_integer_family(arena, inner, var, range, Kink, scan, |_| true);
}
ExprNode::Gt(a, b) | ExprNode::Ge(a, b) | ExprNode::Eq_(a, b) | ExprNode::Ne(a, b) => {
let d = arena.sub(*a, *b);
push_zeros(arena, d, var, Kink, range, scan);
}
ExprNode::Apply(_, args) if args.iter().any(|&a| walk::contains(arena, a, var)) => {
scan.complete = false;
scan.opaque = true;
}
ExprNode::Integral(body, _)
| ExprNode::Derivative(body, _)
| ExprNode::Limit(body, _, _)
| ExprNode::Sum(body, _, _, _)
| ExprNode::Product_(body, _, _, _)
| ExprNode::Residue(body, _, _)
if walk::contains(arena, *body, var) =>
{
scan.complete = false;
scan.opaque = true;
}
ExprNode::RootOf(..) | ExprNode::RootSum(..) | ExprNode::LambertW(_)
if node
.children()
.iter()
.any(|&c| walk::contains(arena, c, var)) =>
{
scan.complete = false;
scan.opaque = true;
}
_ => {}
}
}
fn push_integer_family(
arena: &mut Arena,
inner: ExprId,
var: ExprId,
range: Option<Interval<f64>>,
kind: BreakKind,
scan: &mut BreakScan,
keep: impl Fn(i64) -> bool,
) {
let (Some(lin), Some(lin_ex), Some(r)) = (
linear_coeffs_f64(arena, inner, var),
linear_coeffs_exact(arena, inner, var),
range,
) else {
scan.complete = false;
return;
};
let (alpha, beta) = (lin.slope, lin.intercept);
let (alpha_ex, beta_ex) = (lin_ex.slope, lin_ex.intercept);
let (lo, hi) = (r.lower, r.upper);
if alpha == 0.0 || !lo.is_finite() || !hi.is_finite() {
scan.complete = false;
return;
}
let t_lo = alpha * lo + beta;
let t_hi = alpha * hi + beta;
let (t_min, t_max) = if t_lo <= t_hi {
(t_lo, t_hi)
} else {
(t_hi, t_lo)
};
let n_min = t_min.floor() as i64 - 1;
let n_max = t_max.ceil() as i64 + 1;
if n_max - n_min > 10_000 {
scan.complete = false;
return;
}
for n in n_min..=n_max {
if !keep(n) {
continue;
}
let xv = ((n as f64) - beta) / alpha;
if xv < lo - 1e-12 || xv > hi + 1e-12 {
continue;
}
let n_ex = arena.int(n);
let num = arena.sub(n_ex, beta_ex);
let point = arena.div(num, alpha_ex);
let point = eval::eval(arena, point);
scan.points.push(Breakpoint {
point,
value: Some(xv),
kind,
});
}
}
pub(crate) fn expr_to_f64(arena: &mut Arena, expr: ExprId) -> Option<f64> {
let evaled = eval::eval(arena, expr);
if let Some(r) = arena.as_num(evaled).cloned() {
use num_traits::ToPrimitive;
return r.to_f64();
}
let s = evalf::evalf(arena, evaled, 16).ok()?;
if s.contains('I') || s.contains('i') {
return None;
}
s.trim().parse::<f64>().ok()
}
pub(crate) fn estimate_frequency(
arena: &Arena,
expr: ExprId,
var: ExprId,
_var_sym: SymbolId,
) -> Option<f64> {
tracing::debug!("estimate_frequency: scanning for trig terms");
let post_order = walk::post_order_ids(arena, expr);
let mut max_omega: Option<f64> = None;
for &id in &post_order {
let node = arena.node(id);
let inner = match node {
ExprNode::Sin(i) => Some(*i),
ExprNode::Cos(i) => Some(*i),
ExprNode::Tan(i) => Some(*i),
_ => None,
};
if let Some(inner) = inner
&& let Some(omega) = extract_linear_coefficient(arena, inner, var)
{
let omega_abs = omega.abs();
match max_omega {
Some(cur) if cur >= omega_abs => {}
_ => max_omega = Some(omega_abs),
}
}
}
max_omega
}
fn extract_linear_coefficient(arena: &Arena, expr: ExprId, var: ExprId) -> Option<f64> {
if expr == var {
return Some(1.0);
}
let node = arena.node(expr);
match node {
ExprNode::Mul(args) => {
let mut has_var = false;
let mut coeff_ids: Vec<ExprId> = Vec::new();
for &arg in args.iter() {
if arg == var {
has_var = true;
} else if walk::contains(arena, arg, var) {
return None;
} else {
coeff_ids.push(arg);
}
}
if !has_var {
return None;
}
if coeff_ids.is_empty() {
return Some(1.0);
}
if coeff_ids.len() == 1 {
return num_value(arena, coeff_ids[0]);
}
None
}
ExprNode::Add(args) => {
let mut omega = None;
for &arg in args.iter() {
if walk::contains(arena, arg, var) {
omega = extract_linear_coefficient(arena, arg, var);
}
}
omega
}
ExprNode::Neg(inner) => extract_linear_coefficient(arena, *inner, var).map(|c| -c),
_ => None,
}
}
fn num_value(arena: &Arena, expr: ExprId) -> Option<f64> {
if let Some(r) = arena.as_num(expr) {
use num_traits::ToPrimitive;
r.to_f64()
} else {
match arena.node(expr) {
ExprNode::Neg(inner) => num_value(arena, *inner).map(|v| -v),
_ => None,
}
}
}
#[derive(Clone, Debug, Default)]
pub(crate) struct ContinuityScan {
pub singular: Vec<ExprId>,
pub kinks: Vec<ExprId>,
pub opaque: Option<&'static str>,
}
pub(crate) fn continuity_scan(arena: &mut Arena, expr: ExprId, var: ExprId) -> ContinuityScan {
let mut scan = ContinuityScan::default();
let mut assumptions = crate::base::assumptions::AssumptionCache::new();
for id in walk::post_order_ids(arena, expr) {
let node = arena.node(id).clone();
let depends = |arena: &Arena, child: ExprId| walk::contains(arena, child, var);
match node {
ExprNode::Pow(base, exp) => {
if !depends(arena, base) || depends(arena, exp) {
continue;
}
let negative = match arena.as_num(exp) {
Some(r) => {
use num_traits::Signed;
r.is_negative()
}
None => {
assumptions.query(arena, exp, crate::base::assumptions::Props::NEGATIVE)
== Some(true)
}
};
if negative {
push_unique(&mut scan.singular, base);
}
}
ExprNode::Ln(g) if depends(arena, g) => push_unique(&mut scan.singular, g),
ExprNode::Tan(g) if depends(arena, g) => {
let cos_g = arena.cos(g);
push_unique(&mut scan.singular, cos_g);
}
ExprNode::Atanh(g) if depends(arena, g) => {
let one = arena.one();
let minus = arena.sub(g, one);
let plus = arena.add(&[g, one]);
push_unique(&mut scan.singular, minus);
push_unique(&mut scan.singular, plus);
}
ExprNode::Abs(g) if depends(arena, g) => push_unique(&mut scan.kinks, g),
_ => {
if scan.opaque.is_none()
&& let Some(name) = opaque_node_name(&node)
&& node.children().iter().any(|&c| depends(arena, c))
{
scan.opaque = Some(name);
}
}
}
}
scan
}
fn push_unique(list: &mut Vec<ExprId>, id: ExprId) {
if !list.contains(&id) {
list.push(id);
}
}
fn opaque_node_name(node: &ExprNode) -> Option<&'static str> {
Some(match node {
ExprNode::Floor(_) => "floor",
ExprNode::Ceiling(_) => "ceiling",
ExprNode::Sign(_) => "sign",
ExprNode::Heaviside(_) => "Heaviside",
ExprNode::DiracDelta(_) => "DiracDelta",
ExprNode::Piecewise(_) => "Piecewise",
ExprNode::Min(_) => "min",
ExprNode::Max(_) => "max",
ExprNode::Re(_) => "re",
ExprNode::Im(_) => "im",
ExprNode::Conjugate(_) => "conjugate",
ExprNode::Arg(_) => "arg",
ExprNode::Atan2(..) => "atan2",
ExprNode::Gamma(_) => "Gamma",
ExprNode::LogGamma(_) => "loggamma",
ExprNode::Digamma(_) => "digamma",
ExprNode::Polygamma(..) => "polygamma",
ExprNode::Zeta(_) => "zeta",
ExprNode::Beta(..) => "Beta",
ExprNode::Factorial(_) => "factorial",
ExprNode::Binomial(..) => "binomial",
ExprNode::KroneckerDelta(..) => "KroneckerDelta",
ExprNode::LambertW(_) => "LambertW",
ExprNode::Apply(..) => "an unknown function",
ExprNode::Derivative(..) => "Derivative",
ExprNode::Integral(..) => "Integral",
ExprNode::DefiniteIntegral(..) => "DefiniteIntegral",
ExprNode::Sum(..) => "Sum",
ExprNode::Product_(..) => "Product",
ExprNode::Limit(..) => "Limit",
ExprNode::Series(..) => "Series",
ExprNode::LaplaceTransform(..) => "LaplaceTransform",
ExprNode::InverseLaplaceTransform(..) => "InverseLaplaceTransform",
ExprNode::Residue(..) => "Residue",
ExprNode::RootOf(..) => "RootOf",
ExprNode::RootSum(..) => "RootSum",
ExprNode::DSolve(..) => "DSolve",
ExprNode::Gt(..)
| ExprNode::Ge(..)
| ExprNode::Eq_(..)
| ExprNode::Ne(..)
| ExprNode::And(_)
| ExprNode::Or(_)
| ExprNode::Not(_) => "a boolean",
ExprNode::ConditionSet(..)
| ExprNode::Interval(..)
| ExprNode::FiniteSet(_)
| ExprNode::SetUnion(_)
| ExprNode::SetIntersection(_)
| ExprNode::SetComplement(..) => "a set",
_ => return None,
})
}
pub(crate) fn sign_nodes_of(arena: &Arena, expr: ExprId, var: ExprId) -> Vec<(ExprId, ExprId)> {
let mut out: Vec<(ExprId, ExprId)> = Vec::new();
for id in walk::post_order_ids(arena, expr) {
if let ExprNode::Sign(h) = arena.node(id)
&& walk::contains(arena, *h, var)
&& !out.iter().any(|(s, _)| *s == id)
{
out.push((id, *h));
}
}
out
}
pub(crate) fn has_trig_of(arena: &Arena, expr: ExprId, var: ExprId) -> bool {
walk::post_order_ids(arena, expr)
.into_iter()
.any(|id| match arena.node(id) {
ExprNode::Sin(g) | ExprNode::Cos(g) | ExprNode::Tan(g) => {
walk::contains(arena, *g, var)
}
_ => false,
})
}
pub(crate) fn periodicity(arena: &mut Arena, expr: ExprId, var: ExprId) -> Option<ExprId> {
if !walk::contains(arena, expr, var) {
return Some(arena.zero());
}
let mut memo: rustc_hash::FxHashMap<ExprId, Option<ExprId>> = rustc_hash::FxHashMap::default();
for id in walk::post_order_ids(arena, expr) {
if memo.contains_key(&id) {
continue;
}
let period = if walk::contains(arena, id, var) {
node_period(arena, id, var, &memo)
} else {
Some(arena.zero())
};
memo.insert(id, period);
}
memo.get(&expr).copied().flatten()
}
fn node_period(
arena: &mut Arena,
id: ExprId,
var: ExprId,
memo: &rustc_hash::FxHashMap<ExprId, Option<ExprId>>,
) -> Option<ExprId> {
let node = arena.node(id).clone();
match &node {
ExprNode::Symbol(_) => None,
ExprNode::Sin(g) | ExprNode::Cos(g) | ExprNode::Tan(g) => {
if let Some(p) = trig_power_period(arena, id, var) {
return Some(p);
}
memo.get(g).copied().flatten()
}
ExprNode::Abs(g) => {
if let ExprNode::Sin(h) | ExprNode::Cos(h) = arena.node(*g).clone()
&& let Some(lin) = linear_coeffs_exact(arena, h, var)
{
return Some(pi_over_abs(arena, lin.slope, 1));
}
memo.get(g).copied().flatten()
}
ExprNode::Pow(_, exp) => {
if !walk::contains(arena, *exp, var)
&& let Some(p) = trig_power_period(arena, id, var)
{
return Some(p);
}
lcm_of_children(arena, &node, memo)
}
ExprNode::Mul(children) => {
let children = children.clone();
mul_period(arena, &children, var, memo)
}
ExprNode::Integral(..)
| ExprNode::DefiniteIntegral(..)
| ExprNode::Sum(..)
| ExprNode::Product_(..)
| ExprNode::Limit(..)
| ExprNode::Series(..)
| ExprNode::LaplaceTransform(..)
| ExprNode::InverseLaplaceTransform(..)
| ExprNode::Residue(..)
| ExprNode::RootOf(..)
| ExprNode::RootSum(..)
| ExprNode::DSolve(..)
| ExprNode::Piecewise(_)
| ExprNode::Gt(..)
| ExprNode::Ge(..)
| ExprNode::Eq_(..)
| ExprNode::Ne(..)
| ExprNode::And(_)
| ExprNode::Or(_)
| ExprNode::Not(_)
| ExprNode::ConditionSet(..)
| ExprNode::Interval(..)
| ExprNode::FiniteSet(_)
| ExprNode::SetUnion(_)
| ExprNode::SetIntersection(_)
| ExprNode::SetComplement(..) => None,
_ => lcm_of_children(arena, &node, memo),
}
}
fn lcm_of_children(
arena: &mut Arena,
node: &ExprNode,
memo: &rustc_hash::FxHashMap<ExprId, Option<ExprId>>,
) -> Option<ExprId> {
let mut acc = arena.zero();
for c in node.children() {
let p = memo.get(&c).copied().flatten()?;
acc = lcm_periods(arena, acc, p)?;
}
Some(acc)
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum TrigKind {
SinCos,
Tan,
}
fn as_trig_power(arena: &Arena, factor: ExprId) -> Option<(ExprId, TrigKind, i64)> {
let (base, k) = match arena.node(factor) {
ExprNode::Pow(base, exp) => {
let r = arena.as_num(*exp)?;
if !r.is_integer() {
return None;
}
use num_traits::ToPrimitive;
(*base, r.to_integer().to_i64()?)
}
_ => (factor, 1),
};
match arena.node(base) {
ExprNode::Sin(g) | ExprNode::Cos(g) => Some((*g, TrigKind::SinCos, k)),
ExprNode::Tan(g) => Some((*g, TrigKind::Tan, k)),
_ => None,
}
}
fn trig_power_period(arena: &mut Arena, factor: ExprId, var: ExprId) -> Option<ExprId> {
let (g, kind, k) = as_trig_power(arena, factor)?;
let a = linear_coeffs_exact(arena, g, var)?.slope;
let halves = kind == TrigKind::Tan || k % 2 == 0;
Some(pi_over_abs(arena, a, if halves { 1 } else { 2 }))
}
fn mul_period(
arena: &mut Arena,
children: &[ExprId],
var: ExprId,
memo: &rustc_hash::FxHashMap<ExprId, Option<ExprId>>,
) -> Option<ExprId> {
let mut groups: Vec<(ExprId, ExprId, i64)> = Vec::new();
let mut acc = arena.zero();
for &c in children {
if !walk::contains(arena, c, var) {
continue;
}
if let Some((g, kind, k)) = as_trig_power(arena, c)
&& let Some(LinearCoeffs { slope: a, .. }) = linear_coeffs_exact(arena, g, var)
{
let weight = if kind == TrigKind::Tan { 0 } else { k };
match groups.iter_mut().find(|(gg, _, _)| *gg == g) {
Some(entry) => entry.2 += weight,
None => groups.push((g, a, weight)),
}
continue;
}
let p = memo.get(&c).copied().flatten()?;
acc = lcm_periods(arena, acc, p)?;
}
for (_, a, sum) in groups {
let p = pi_over_abs(arena, a, if sum % 2 == 0 { 1 } else { 2 });
acc = lcm_periods(arena, acc, p)?;
}
Some(acc)
}
fn pi_over_abs(arena: &mut Arena, a: ExprId, k: i64) -> ExprId {
let k = arena.int(k);
let pi = arena.pi();
let num = arena.mul(&[k, pi]);
let mut assumptions = crate::base::assumptions::AssumptionCache::new();
let den = if assumptions.query(arena, a, crate::base::assumptions::Props::POSITIVE)
== Some(true)
{
a
} else if assumptions.query(arena, a, crate::base::assumptions::Props::NEGATIVE) == Some(true) {
arena.neg(a)
} else {
arena.abs(a)
};
let q = arena.div(num, den);
eval::eval(arena, q)
}
fn lcm_periods(arena: &mut Arena, a: ExprId, b: ExprId) -> Option<ExprId> {
if arena.is_zero_structural(a) {
return Some(b);
}
if arena.is_zero_structural(b) || a == b {
return Some(a);
}
let ratio = arena.div(a, b);
let ratio = eval::eval(arena, ratio);
let r = arena.as_num(ratio)?.clone();
use num_traits::Signed;
if !r.is_positive() {
return None;
}
let n = arena.big_int(r.denom().clone());
let l = arena.mul(&[a, n]);
Some(eval::eval(arena, l))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::base::arena::Arena;
use crate::base::node::{ExprNode, INTERVAL_BOTH_CLOSED, INTERVAL_BOTH_OPEN};
use std::f64::consts::PI;
fn sym_id(arena: &Arena, expr: ExprId) -> SymbolId {
match arena.node(expr) {
ExprNode::Symbol(sid) => *sid,
_ => panic!("expected Symbol node"),
}
}
fn set_display(arena: &Arena, set: ExprId) -> String {
use crate::output::display;
display::format_expr(arena, set)
}
#[test]
fn domain_sqrt_x() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let expr = arena.sqrt(x);
let reals = arena.interval(arena.neg_infinity, arena.infinity, INTERVAL_BOTH_OPEN);
let result = continuous_domain(&mut arena, expr, x, x_sym, reals);
let s = set_display(&arena, result);
assert!(
s.contains("0") && (s.contains("oo") || s.contains("∞")),
"sqrt(x) domain should be [0, ∞), got: {s}"
);
assert!(
!s.contains("EmptySet"),
"sqrt(x) domain should not be empty: {s}"
);
}
#[test]
fn domain_1_over_x() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let expr = arena.div(arena.one, x);
let reals = arena.interval(arena.neg_infinity, arena.infinity, INTERVAL_BOTH_OPEN);
let result = continuous_domain(&mut arena, expr, x, x_sym, reals);
let s = set_display(&arena, result);
assert!(
!s.contains("EmptySet"),
"1/x domain should not be empty: {s}"
);
assert!(
s.contains("0"),
"1/x domain should reference 0 as excluded point: {s}"
);
}
#[test]
fn domain_ln_x() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let expr = arena.ln(x);
let reals = arena.interval(arena.neg_infinity, arena.infinity, INTERVAL_BOTH_OPEN);
let result = continuous_domain(&mut arena, expr, x, x_sym, reals);
let s = set_display(&arena, result);
assert!(
!s.contains("EmptySet"),
"ln(x) domain should not be empty: {s}"
);
assert!(
s.contains("0") && (s.contains("oo") || s.contains("∞")),
"ln(x) domain should be (0, ∞), got: {s}"
);
}
#[test]
fn domain_sqrt_x_minus_2() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let two = arena.int(2);
let inner = arena.sub(x, two);
let expr = arena.sqrt(inner);
let neg5 = arena.int(-5);
let five = arena.int(5);
let domain = arena.interval(neg5, five, INTERVAL_BOTH_CLOSED);
let result = continuous_domain(&mut arena, expr, x, x_sym, domain);
let s = set_display(&arena, result);
assert!(
!s.contains("EmptySet"),
"sqrt(x-2) on [-5,5] should not be empty: {s}"
);
assert!(
s.contains("2") && s.contains("5"),
"sqrt(x-2) on [-5,5] should give [2, 5], got: {s}"
);
}
#[test]
fn singularities_tan_x() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let expr = arena.tan(x);
let sings = singularities(&mut arena, expr, x, x_sym, Interval::closed(0.0, 5.0));
assert!(
!sings.is_empty(),
"tan(x) should have singularities in [0, 5]"
);
let has_pi_half = sings.iter().any(|&v| (v - PI / 2.0).abs() < 0.1);
assert!(
has_pi_half,
"tan(x) singularities should include ≈π/2, got: {:?}",
sings
);
}
#[test]
fn singularities_1_over_x() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let expr = arena.div(arena.one, x);
let sings = singularities(&mut arena, expr, x, x_sym, Interval::closed(-2.0, 2.0));
assert_eq!(
sings.len(),
1,
"1/x should have one singularity: {:?}",
sings
);
assert!(
sings[0].abs() < 1e-10,
"1/x singularity should be at 0, got: {}",
sings[0]
);
}
#[test]
fn frequency_sin_100x() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let hundred = arena.int(100);
let inner = arena.mul(&[hundred, x]);
let expr = arena.sin(inner);
let freq = estimate_frequency(&arena, expr, x, x_sym);
assert!(freq.is_some(), "sin(100*x) should have a frequency");
let omega = freq.unwrap();
assert!(
(omega - 100.0).abs() < 1e-10,
"sin(100*x) angular frequency should be 100 rad/s, got: {}",
omega
);
}
#[test]
fn frequency_no_trig() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let x_sym = sym_id(&arena, x);
let two = arena.int(2);
let x2 = arena.pow(x, two);
let one = arena.one;
let expr = arena.add(&[x2, one]);
let freq = estimate_frequency(&arena, expr, x, x_sym);
assert!(
freq.is_none(),
"x^2 + 1 should have no frequency, got: {:?}",
freq
);
}
}