use ocas_atom::normalize::normalize;
use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use super::{is_constant, linear_form};
pub(crate) fn special_integrate<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
x: Atom<'a>,
) -> Option<Atom<'a>> {
let mut factors: Vec<Atom> = Vec::new();
flat_factors(expr, &mut factors);
erf_family(ctx, &factors, x)
.or_else(|| ei_family(ctx, &factors, x))
.or_else(|| trig_integral_family(ctx, &factors, x))
.or_else(|| fresnel_family(ctx, &factors, x))
.or_else(|| reduction_families(ctx, &factors, x))
}
fn flat_factors<'a>(expr: Atom<'a>, out: &mut Vec<Atom<'a>>) {
match expr.node() {
AtomNode::Mul(args) => {
for a in args.iter() {
flat_factors(*a, out);
}
}
_ => out.push(expr),
}
}
fn as_exp(f: Atom) -> Option<Atom> {
if let AtomNode::Fun(name, args) = f.node() {
if name.as_str() == "exp" && args.len() == 1 {
return Some(args[0]);
}
}
None
}
fn is_x_inv<'a>(f: Atom<'a>, x: Atom<'a>) -> bool {
matches!(f.node(), AtomNode::Pow(b, e) if *b == x && matches!(e.node(), AtomNode::Num(-1)))
}
fn as_quadratic<'a>(u: Atom<'a>, x: Atom<'a>) -> Option<Atom<'a>> {
if matches!(u.node(), AtomNode::Pow(b, e) if *b == x && matches!(e.node(), AtomNode::Num(2))) {
return None; }
if let AtomNode::Mul(factors) = u.node() {
if factors.len() == 2 {
if matches!(factors[1].node(), AtomNode::Pow(b, e)
if *b == x && matches!(e.node(), AtomNode::Num(2)))
{
return Some(factors[0]);
}
}
}
None
}
fn is_x_squared<'a>(u: Atom<'a>, x: Atom<'a>) -> bool {
if matches!(u.node(), AtomNode::Pow(b, e) if *b == x && matches!(e.node(), AtomNode::Num(2))) {
return true;
}
if let AtomNode::Pow(b, e) = u.node() {
if matches!(e.node(), AtomNode::Num(2)) {
if let AtomNode::Mul(factors) = b.node() {
let all_neg_one_or_x = factors
.iter()
.all(|f| matches!(f.node(), AtomNode::Num(-1)) || *f == x);
let has_x = factors.contains(&x);
return all_neg_one_or_x && has_x;
}
}
}
false
}
fn is_x<'a>(u: Atom<'a>, x: Atom<'a>) -> bool {
u == x
}
fn erf_family<'a>(ctx: &'a AtomArena<'a>, factors: &[Atom<'a>], x: Atom<'a>) -> Option<Atom<'a>> {
if factors.len() != 1 {
return None;
}
let u = as_exp(factors[0])?;
if is_x_squared(u, x) {
let sqrt_pi = ctx.fun("sqrt", &[ctx.var("pi")]);
let erfi = ctx.fun("erfi", &[x]);
return Some(ctx.mul(&[sqrt_pi, ctx.pow(ctx.num(2), ctx.num(-1)), erfi]));
}
if let Some(c) = as_quadratic(u, x) {
let neg_c = ctx.mul(&[ctx.num(-1), c]);
let root = ctx.fun("sqrt", &[neg_c]);
let sqrt_pi = ctx.fun("sqrt", &[ctx.var("pi")]);
let erf = ctx.fun("erf", &[ctx.mul(&[root, x])]);
return Some(ctx.mul(&[
sqrt_pi,
ctx.pow(ctx.mul(&[ctx.num(2), root]), ctx.num(-1)),
erf,
]));
}
None
}
fn ei_family<'a>(ctx: &'a AtomArena<'a>, factors: &[Atom<'a>], x: Atom<'a>) -> Option<Atom<'a>> {
if factors.len() != 2 {
return None;
}
let (exp_f, inv_f) = if as_exp(factors[0]).is_some() {
(factors[0], factors[1])
} else if as_exp(factors[1]).is_some() {
(factors[1], factors[0])
} else {
return None;
};
if !is_x_inv(inv_f, x) {
return None;
}
let u = as_exp(exp_f)?;
if is_x(u, x) {
return Some(ctx.fun("Ei", &[x]));
}
if let AtomNode::Mul(mf) = u.node() {
if mf.len() == 2 && mf[1] == x && matches!(mf[0].node(), AtomNode::Num(_)) {
return Some(ctx.fun("Ei", &[u]));
}
}
None
}
fn trig_integral_family<'a>(
ctx: &'a AtomArena<'a>,
factors: &[Atom<'a>],
x: Atom<'a>,
) -> Option<Atom<'a>> {
if factors.len() != 2 {
return None;
}
let (fun_f, inv_f) = if is_x_inv(factors[1], x) {
(factors[0], factors[1])
} else if is_x_inv(factors[0], x) {
(factors[1], factors[0])
} else {
return None;
};
let _ = inv_f;
let AtomNode::Fun(name, args) = fun_f.node() else {
return None;
};
if args.len() != 1 || !is_x(args[0], x) {
return None;
}
let target = match name.as_str() {
"sin" => "Si",
"cos" => "Ci",
"sinh" => "Shi",
"cosh" => "Chi",
_ => return None,
};
Some(ctx.fun(target, &[x]))
}
fn fresnel_family<'a>(
ctx: &'a AtomArena<'a>,
factors: &[Atom<'a>],
x: Atom<'a>,
) -> Option<Atom<'a>> {
if factors.len() != 1 {
return None;
}
let AtomNode::Fun(name, args) = factors[0].node() else {
return None;
};
if args.len() != 1 || !is_x_squared(args[0], x) {
return None;
}
let target = match name.as_str() {
"sin" => "fresnels",
"cos" => "fresnelc",
_ => return None,
};
let pi = ctx.var("pi");
let two = ctx.num(2);
let prefactor = ctx.fun("sqrt", &[ctx.mul(&[pi, ctx.pow(two, ctx.num(-1))])]);
let inner = ctx.mul(&[
ctx.fun("sqrt", &[ctx.mul(&[two, ctx.pow(pi, ctx.num(-1))])]),
x,
]);
Some(ctx.mul(&[prefactor, ctx.fun(target, &[inner])]))
}
const MAX_SPECIAL_DEG: i64 = 8;
const MAX_SPECIAL_STEPS: usize = 16;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
enum Head {
Erf,
Erfc,
Erfi,
Si,
Ci,
Shi,
Chi,
Ei,
}
impl Head {
fn from_name(name: &str) -> Option<Head> {
Some(match name {
"erf" => Head::Erf,
"erfc" => Head::Erfc,
"erfi" => Head::Erfi,
"Si" => Head::Si,
"Ci" => Head::Ci,
"Shi" => Head::Shi,
"Chi" => Head::Chi,
"Ei" => Head::Ei,
_ => return None,
})
}
fn name(self) -> &'static str {
match self {
Head::Erf => "erf",
Head::Erfc => "erfc",
Head::Erfi => "erfi",
Head::Si => "Si",
Head::Ci => "Ci",
Head::Shi => "Shi",
Head::Chi => "Chi",
Head::Ei => "Ei",
}
}
fn square_capable(self) -> bool {
!matches!(self, Head::Ei | Head::Erfc)
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Kind {
Sin,
Cos,
Sinh,
Cosh,
}
impl Kind {
fn name(self) -> &'static str {
match self {
Kind::Sin => "sin",
Kind::Cos => "cos",
Kind::Sinh => "sinh",
Kind::Cosh => "cosh",
}
}
fn atom<'a>(self, ctx: &'a AtomArena<'a>, t: Atom<'a>) -> Atom<'a> {
ctx.fun(self.name(), &[t])
}
}
#[derive(Clone, Copy, Debug)]
enum Kernel<'a> {
Plain(Head, Atom<'a>),
Square(Head, Atom<'a>),
Order(i64, Atom<'a>),
}
#[derive(Clone, Copy)]
struct Lin<'a> {
u: Atom<'a>,
a: Atom<'a>,
b: Atom<'a>,
x: Atom<'a>,
a_zero: bool,
}
fn is_zero(a: Atom<'_>) -> bool {
matches!(a.node(), AtomNode::Num(0))
}
fn muln<'a>(ctx: &'a AtomArena<'a>, args: &[Atom<'a>]) -> Atom<'a> {
let mut out = Vec::with_capacity(args.len());
for a in args {
match a.node() {
AtomNode::Num(0) => return ctx.num(0),
AtomNode::Num(1) => {}
_ => out.push(*a),
}
}
match out.len() {
0 => ctx.num(1),
1 => out[0],
_ => ctx.mul(&out),
}
}
fn addn<'a>(ctx: &'a AtomArena<'a>, args: &[Atom<'a>]) -> Atom<'a> {
let mut out = Vec::with_capacity(args.len());
for a in args {
if !is_zero(*a) {
out.push(*a);
}
}
match out.len() {
0 => ctx.num(0),
1 => out[0],
_ => ctx.add(&out),
}
}
fn neg<'a>(ctx: &'a AtomArena<'a>, a: Atom<'a>) -> Atom<'a> {
match a.node() {
AtomNode::Num(v) => ctx.num(-v),
_ => ctx.mul(&[ctx.num(-1), a]),
}
}
fn powi_clean<'a>(ctx: &'a AtomArena<'a>, base: Atom<'a>, e: i64) -> Atom<'a> {
if e == 0 {
return ctx.num(1);
}
if let AtomNode::Num(v) = base.node() {
match *v {
1 => return ctx.num(1),
0 => {
return if e > 0 {
ctx.num(0)
} else {
ctx.pow(base, ctx.num(e))
};
}
-1 => return ctx.num(if e % 2 == 0 { 1 } else { -1 }),
_ => {
if e > 0
&& let Ok(e32) = u32::try_from(e)
&& let Some(p) = v.checked_pow(e32)
{
return ctx.num(p);
}
}
}
}
if e == 1 {
return base;
}
ctx.pow(base, ctx.num(e))
}
fn q_frac<'a>(ctx: &'a AtomArena<'a>, p: i64, q: i64) -> Atom<'a> {
debug_assert!(q != 0, "q_frac with zero denominator");
if p % q == 0 {
return ctx.num(p / q);
}
if p == 1 {
return powi_clean(ctx, ctx.num(q), -1);
}
if p == -1 {
return muln(ctx, &[ctx.num(-1), powi_clean(ctx, ctx.num(q), -1)]);
}
muln(ctx, &[ctx.num(p), powi_clean(ctx, ctx.num(q), -1)])
}
fn exp_atom<'a>(ctx: &'a AtomArena<'a>, arg: Atom<'a>) -> Atom<'a> {
ctx.fun("exp", &[arg])
}
fn sqrt_pi<'a>(ctx: &'a AtomArena<'a>) -> Atom<'a> {
ctx.fun("sqrt", &[ctx.var("pi")])
}
fn inv_sqrt_pi<'a>(ctx: &'a AtomArena<'a>) -> Atom<'a> {
powi_clean(ctx, sqrt_pi(ctx), -1)
}
fn binom_k(n: i64, k: i64) -> i64 {
if k < 0 || k > n {
return 0;
}
let k = k.min(n - k);
let mut acc = 1i64;
for i in 0..k {
acc = acc * (n - i) / (i + 1);
}
acc
}
fn lin_of<'a>(ctx: &'a AtomArena<'a>, u: Atom<'a>, x: Symbol) -> Option<Lin<'a>> {
let (b, a) = linear_form(ctx, u, x)?;
if !is_constant(a, x) || !is_constant(b, x) || is_zero(b) {
return None;
}
let b = normalize(ctx, b);
let a = normalize(ctx, a);
Some(Lin {
u,
a_zero: is_zero(a),
a,
b,
x: ctx.var(x.as_str()),
})
}
fn match_plain<'a>(ctx: &'a AtomArena<'a>, f: Atom<'a>, x: Symbol) -> Option<Kernel<'a>> {
let AtomNode::Fun(name, args) = f.node() else {
return None;
};
if name.as_str() == "Ei" {
if args.len() == 2 {
if let AtomNode::Num(n) = args[0].node()
&& (-MAX_SPECIAL_DEG..=MAX_SPECIAL_DEG).contains(n)
&& lin_of(ctx, args[1], x).is_some()
{
return Some(Kernel::Order(*n, args[1]));
}
}
if args.len() == 1 && lin_of(ctx, args[0], x).is_some() {
return Some(Kernel::Plain(Head::Ei, args[0]));
}
return None;
}
let h = Head::from_name(name.as_str())?;
if args.len() != 1 {
return None;
}
lin_of(ctx, args[0], x)?;
Some(Kernel::Plain(h, args[0]))
}
fn match_square<'a>(ctx: &'a AtomArena<'a>, f: Atom<'a>, x: Symbol) -> Option<Kernel<'a>> {
let AtomNode::Pow(base, exp) = f.node() else {
return None;
};
if !matches!(exp.node(), AtomNode::Num(2)) {
return None;
}
let AtomNode::Fun(name, args) = base.node() else {
return None;
};
let h = Head::from_name(name.as_str()).filter(|h| h.square_capable())?;
if args.len() != 1 {
return None;
}
lin_of(ctx, args[0], x)?;
Some(Kernel::Square(h, args[0]))
}
fn split_kernel<'a>(
ctx: &'a AtomArena<'a>,
factors: &[Atom<'a>],
x: Symbol,
) -> Option<(Kernel<'a>, Vec<Atom<'a>>)> {
let mut found: Option<(usize, Kernel<'a>)> = None;
for (i, f) in factors.iter().enumerate() {
if let Some(k) = match_plain(ctx, *f, x).or_else(|| match_square(ctx, *f, x)) {
if found.is_some() {
return None;
}
found = Some((i, k));
}
}
let (idx, kernel) = found?;
let rest = factors
.iter()
.enumerate()
.filter(|(j, _)| *j != idx)
.map(|(_, a)| *a)
.collect();
Some((kernel, rest))
}
fn reduction_families<'a>(
ctx: &'a AtomArena<'a>,
factors: &[Atom<'a>],
x: Atom<'a>,
) -> Option<Atom<'a>> {
let xs = match x.node() {
AtomNode::Var(v) => *v,
_ => return None,
};
let (kernel, rest) = split_kernel(ctx, factors, xs)?;
let rest = muln(ctx, &rest);
match kernel {
Kernel::Order(n, u) => order_family(ctx, n, u, rest, xs),
Kernel::Plain(h, u) => plain_family(ctx, h, u, rest, xs),
Kernel::Square(h, u) => square_family(ctx, h, u, rest, xs),
}
}
fn poly_add<'a>(ctx: &'a AtomArena<'a>, a: &[Atom<'a>], b: &[Atom<'a>]) -> Option<Vec<Atom<'a>>> {
let len = a.len().max(b.len());
if len as i64 > MAX_SPECIAL_DEG + 1 {
return None;
}
let mut out = Vec::with_capacity(len);
for i in 0..len {
let x = a.get(i).copied().unwrap_or_else(|| ctx.num(0));
let y = b.get(i).copied().unwrap_or_else(|| ctx.num(0));
out.push(addn(ctx, &[x, y]));
}
Some(out)
}
fn poly_mul<'a>(ctx: &'a AtomArena<'a>, a: &[Atom<'a>], b: &[Atom<'a>]) -> Option<Vec<Atom<'a>>> {
if a.len() + b.len() - 1 > (MAX_SPECIAL_DEG + 1) as usize {
return None;
}
let mut out = vec![ctx.num(0); a.len() + b.len() - 1];
for (i, x) in a.iter().enumerate() {
if is_zero(*x) {
continue;
}
for (j, y) in b.iter().enumerate() {
if is_zero(*y) {
continue;
}
out[i + j] = addn(ctx, &[out[i + j], muln(ctx, &[*x, *y])]);
}
}
Some(out)
}
fn poly_coeffs<'a>(ctx: &'a AtomArena<'a>, e: Atom<'a>, x: Symbol) -> Option<Vec<Atom<'a>>> {
if is_constant(e, x) {
return Some(vec![e]);
}
match e.node() {
AtomNode::Var(v) if *v == x => Some(vec![ctx.num(0), ctx.num(1)]),
AtomNode::Add(args) => {
let mut acc = vec![ctx.num(0)];
for a in args.iter() {
let c = poly_coeffs(ctx, *a, x)?;
acc = poly_add(ctx, &acc, &c)?;
}
Some(acc)
}
AtomNode::Mul(args) => {
let mut acc = vec![ctx.num(1)];
for a in args.iter() {
let c = poly_coeffs(ctx, *a, x)?;
acc = poly_mul(ctx, &acc, &c)?;
}
Some(acc)
}
AtomNode::Pow(b, ex) => {
let AtomNode::Num(n) = ex.node() else {
return None;
};
if *n < 0 || *n > MAX_SPECIAL_DEG {
return None;
}
let base = poly_coeffs(ctx, *b, x)?;
let mut acc = vec![ctx.num(1)];
for _ in 0..*n {
acc = poly_mul(ctx, &acc, &base)?;
}
Some(acc)
}
_ => None,
}
}
fn x_power(e: Atom<'_>, x: Symbol) -> Option<i64> {
match e.node() {
AtomNode::Var(v) if *v == x => Some(1),
AtomNode::Pow(b, ex) => {
let AtomNode::Num(k) = ex.node() else {
return None;
};
if *k >= 1 && x_power(*b, x) == Some(1) {
Some(*k)
} else {
None
}
}
_ => None,
}
}
fn neg_power_monomial<'a>(
ctx: &'a AtomArena<'a>,
e: Atom<'a>,
x: Symbol,
) -> Option<(Atom<'a>, i64)> {
let neg_pow = |f: Atom<'a>| -> Option<i64> {
let AtomNode::Pow(b, ex) = f.node() else {
return None;
};
let AtomNode::Num(m) = ex.node() else {
return None;
};
if *m >= 0 {
return None;
}
let k = x_power(*b, x)?;
let eff = k.checked_mul(-*m)?;
if eff >= 1 { Some(eff) } else { None }
};
match e.node() {
AtomNode::Pow(_, _) => Some((ctx.num(1), neg_pow(e)?)),
AtomNode::Mul(args) => {
let mut coeff: Vec<Atom<'a>> = Vec::new();
let mut power: Option<i64> = None;
for a in args.iter() {
if is_constant(*a, x) {
coeff.push(*a);
continue;
}
if power.is_some() {
return None;
}
power = Some(neg_pow(*a)?);
}
Some((muln(ctx, &coeff), power?))
}
_ => None,
}
}
fn int_pow_exp<'a>(
ctx: &'a AtomArena<'a>,
t: Atom<'a>,
c: Atom<'a>,
m: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || m < 0 {
return None;
}
let cinv = powi_clean(ctx, c, -1);
if m == 0 {
return Some(muln(ctx, &[exp_atom(ctx, muln(ctx, &[c, t])), cinv]));
}
let rec = int_pow_exp(ctx, t, c, m - 1, budget - 1)?;
Some(addn(
ctx,
&[
muln(
ctx,
&[
powi_clean(ctx, t, m),
exp_atom(ctx, muln(ctx, &[c, t])),
cinv,
],
),
neg(ctx, muln(ctx, &[ctx.num(m), cinv, rec])),
],
))
}
fn int_pow_exp_quad<'a>(
ctx: &'a AtomArena<'a>,
t: Atom<'a>,
m: i64,
alpha: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || m < 0 || alpha == 0 {
return None;
}
let t2 = powi_clean(ctx, t, 2);
let ex = exp_atom(ctx, muln(ctx, &[ctx.num(alpha), t2]));
if m == 0 {
let root = match alpha.abs() {
1 => None,
2 => Some(ctx.fun("sqrt", &[ctx.num(2)])),
_ => return None,
};
let arg = match root {
Some(r) => muln(ctx, &[r, t]),
None => t,
};
let f = if alpha > 0 {
ctx.fun("erfi", &[arg])
} else {
ctx.fun("erf", &[arg])
};
let coef = match root {
Some(r) => muln(
ctx,
&[q_frac(ctx, 1, 2), sqrt_pi(ctx), powi_clean(ctx, r, -1)],
),
None => muln(ctx, &[q_frac(ctx, 1, 2), sqrt_pi(ctx)]),
};
return Some(muln(ctx, &[coef, f]));
}
if m % 2 == 1 {
let half = (m - 1) / 2;
let ie = int_pow_exp(ctx, t2, ctx.num(alpha), half, budget - 1)?;
return Some(muln(ctx, &[q_frac(ctx, 1, 2), ie]));
}
let rec = int_pow_exp_quad(ctx, t, m - 2, alpha, budget - 1)?;
let inner = addn(
ctx,
&[
muln(ctx, &[powi_clean(ctx, t, m - 1), ex]),
neg(ctx, muln(ctx, &[ctx.num(m - 1), rec])),
],
);
Some(muln(ctx, &[q_frac(ctx, 1, 2 * alpha), inner]))
}
fn int_pow_trig<'a>(
ctx: &'a AtomArena<'a>,
t: Atom<'a>,
m: i64,
kind: Kind,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || m < 0 {
return None;
}
if m == 0 {
return Some(match kind {
Kind::Sin => neg(ctx, Kind::Cos.atom(ctx, t)),
Kind::Cos => Kind::Sin.atom(ctx, t),
Kind::Sinh => Kind::Cosh.atom(ctx, t),
Kind::Cosh => Kind::Sinh.atom(ctx, t),
});
}
let (next, sign) = match kind {
Kind::Sin => (Kind::Cos, 1),
Kind::Cos => (Kind::Sin, -1),
Kind::Sinh => (Kind::Cosh, -1),
Kind::Cosh => (Kind::Sinh, -1),
};
let rec = int_pow_trig(ctx, t, m - 1, next, budget - 1)?;
let prim = match kind {
Kind::Sin => neg(ctx, Kind::Cos.atom(ctx, t)),
Kind::Cos => Kind::Sin.atom(ctx, t),
Kind::Sinh => Kind::Cosh.atom(ctx, t),
Kind::Cosh => Kind::Sinh.atom(ctx, t),
};
let lead = muln(ctx, &[powi_clean(ctx, t, m), prim]);
let tail = muln(ctx, &[ctx.num(m), rec]);
Some(addn(
ctx,
&[lead, if sign > 0 { tail } else { neg(ctx, tail) }],
))
}
fn int_exp_over_pow<'a>(
ctx: &'a AtomArena<'a>,
u: Atom<'a>,
c: Atom<'a>,
j: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || j < 1 {
return None;
}
let ex = exp_atom(ctx, muln(ctx, &[c, u]));
if j == 1 {
return Some(ctx.fun("Ei", &[muln(ctx, &[c, u])]));
}
let rec = int_exp_over_pow(ctx, u, c, j - 1, budget - 1)?;
Some(addn(
ctx,
&[
neg(
ctx,
muln(ctx, &[ex, powi_clean(ctx, u, 1 - j), q_frac(ctx, 1, j - 1)]),
),
muln(ctx, &[c, q_frac(ctx, 1, j - 1), rec]),
],
))
}
fn plain_family<'a>(
ctx: &'a AtomArena<'a>,
h: Head,
u: Atom<'a>,
rest: Atom<'a>,
x: Symbol,
) -> Option<Atom<'a>> {
let lin = lin_of(ctx, u, x)?;
if let Some(coeffs) = poly_coeffs(ctx, rest, x) {
return poly_plain(ctx, h, lin, &coeffs);
}
if h == Head::Ei && lin.a_zero {
let (c, m) = neg_power_monomial(ctx, rest, x)?;
let d = d_ei(ctx, lin.u, m)?;
return Some(muln(ctx, &[c, powi_clean(ctx, lin.b, m - 1), d]));
}
None
}
fn poly_plain<'a>(
ctx: &'a AtomArena<'a>,
h: Head,
lin: Lin<'a>,
coeffs: &[Atom<'a>],
) -> Option<Atom<'a>> {
let mut terms = Vec::new();
for (i, c) in coeffs.iter().enumerate() {
if is_zero(*c) {
continue;
}
let i = i as i64;
let mut inner = Vec::new();
for j in 0..=i {
let s = s_plain(ctx, h, lin.u, j, MAX_SPECIAL_STEPS)?;
let w = muln(
ctx,
&[
ctx.num(binom_k(i, j)),
powi_clean(ctx, neg(ctx, lin.a), i - j),
s,
],
);
inner.push(w);
}
let scale = muln(ctx, &[*c, powi_clean(ctx, lin.b, -(i + 1))]);
terms.push(muln(ctx, &[scale, addn(ctx, &inner)]));
}
Some(addn(ctx, &terms))
}
fn s_plain<'a>(
ctx: &'a AtomArena<'a>,
h: Head,
u: Atom<'a>,
j: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || j < 0 {
return None;
}
let jp = j + 1;
let lead = muln(
ctx,
&[
powi_clean(ctx, u, jp),
ctx.fun(h.name(), &[u]),
q_frac(ctx, 1, jp),
],
);
let r = r_plain(ctx, h, u, jp, budget - 1)?;
Some(addn(
ctx,
&[lead, neg(ctx, muln(ctx, &[q_frac(ctx, 1, jp), r]))],
))
}
fn r_plain<'a>(
ctx: &'a AtomArena<'a>,
h: Head,
u: Atom<'a>,
k: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || k < 1 {
return None;
}
let two_over_sqrt_pi = muln(ctx, &[ctx.num(2), inv_sqrt_pi(ctx)]);
Some(match h {
Head::Erf => muln(
ctx,
&[
two_over_sqrt_pi,
int_pow_exp_quad(ctx, u, k, -1, budget - 1)?,
],
),
Head::Erfc => neg(
ctx,
muln(
ctx,
&[
two_over_sqrt_pi,
int_pow_exp_quad(ctx, u, k, -1, budget - 1)?,
],
),
),
Head::Erfi => muln(
ctx,
&[
two_over_sqrt_pi,
int_pow_exp_quad(ctx, u, k, 1, budget - 1)?,
],
),
Head::Si => int_pow_trig(ctx, u, k - 1, Kind::Sin, budget - 1)?,
Head::Ci => int_pow_trig(ctx, u, k - 1, Kind::Cos, budget - 1)?,
Head::Shi => int_pow_trig(ctx, u, k - 1, Kind::Sinh, budget - 1)?,
Head::Chi => int_pow_trig(ctx, u, k - 1, Kind::Cosh, budget - 1)?,
Head::Ei => int_pow_exp(ctx, u, ctx.num(1), k - 1, budget - 1)?,
})
}
fn d_ei<'a>(ctx: &'a AtomArena<'a>, u: Atom<'a>, m: i64) -> Option<Atom<'a>> {
if m < 2 {
return None;
}
let lead = neg(
ctx,
muln(
ctx,
&[
ctx.fun("Ei", &[u]),
powi_clean(ctx, u, 1 - m),
q_frac(ctx, 1, m - 1),
],
),
);
let tail = muln(
ctx,
&[
q_frac(ctx, 1, m - 1),
int_exp_over_pow(ctx, u, ctx.num(1), m, MAX_SPECIAL_STEPS)?,
],
);
Some(addn(ctx, &[lead, tail]))
}
fn order_family<'a>(
ctx: &'a AtomArena<'a>,
n: i64,
u: Atom<'a>,
rest: Atom<'a>,
x: Symbol,
) -> Option<Atom<'a>> {
let lin = lin_of(ctx, u, x)?;
if let Some(coeffs) = poly_coeffs(ctx, rest, x) {
if coeffs.len() == 1 || lin.a_zero {
let mut terms = Vec::new();
for (k, c) in coeffs.iter().enumerate() {
if is_zero(*c) {
continue;
}
let j = j_ei(ctx, &lin, k as i64, n, MAX_SPECIAL_STEPS)?;
terms.push(muln(ctx, &[*c, j]));
}
return Some(addn(ctx, &terms));
}
return None;
}
if lin.a_zero {
let (c, m) = neg_power_monomial(ctx, rest, x)?;
let d = d_order(ctx, lin.u, m, n, MAX_SPECIAL_STEPS)?;
return Some(muln(ctx, &[c, powi_clean(ctx, lin.b, m - 1), d]));
}
None
}
fn j_ei<'a>(
ctx: &'a AtomArena<'a>,
lin: &Lin<'a>,
k: i64,
n: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || k < 0 {
return None;
}
if k == 0 {
let head = ctx.fun("Ei", &[ctx.num(n + 1), lin.u]);
return Some(neg(ctx, muln(ctx, &[head, powi_clean(ctx, lin.b, -1)])));
}
if !lin.a_zero {
return None;
}
let cc = neg(ctx, lin.b);
let binv = powi_clean(ctx, lin.b, -1);
let base = int_pow_exp(ctx, lin.x, cc, k - 1, MAX_SPECIAL_STEPS)?;
if n == 0 {
return Some(muln(ctx, &[binv, base]));
}
let rec = j_ei(ctx, lin, k - 1, n + 1, budget - 1)?;
Some(addn(
ctx,
&[
muln(ctx, &[binv, base]),
neg(ctx, muln(ctx, &[ctx.num(n), binv, rec])),
],
))
}
fn d_order<'a>(
ctx: &'a AtomArena<'a>,
u: Atom<'a>,
m: i64,
n: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || m < 1 {
return None;
}
if n == 0 {
return int_exp_over_pow(ctx, u, ctx.num(-1), m + 1, MAX_SPECIAL_STEPS);
}
if n > 0 {
if m <= n {
return None;
}
let lead = neg(
ctx,
muln(
ctx,
&[
ctx.fun("Ei", &[ctx.num(n), u]),
powi_clean(ctx, u, 1 - m),
q_frac(ctx, 1, m - 1),
],
),
);
let rec = d_order(ctx, u, m - 1, n - 1, budget - 1)?;
return Some(addn(
ctx,
&[lead, neg(ctx, muln(ctx, &[q_frac(ctx, 1, m - 1), rec]))],
));
}
let lead = int_exp_over_pow(ctx, u, ctx.num(-1), m + 1, MAX_SPECIAL_STEPS)?;
let rec = d_order(ctx, u, m + 1, n + 1, budget - 1)?;
Some(addn(ctx, &[lead, neg(ctx, muln(ctx, &[ctx.num(n), rec]))]))
}
fn square_family<'a>(
ctx: &'a AtomArena<'a>,
h: Head,
u: Atom<'a>,
rest: Atom<'a>,
x: Symbol,
) -> Option<Atom<'a>> {
let lin = lin_of(ctx, u, x)?;
let coeffs = poly_coeffs(ctx, rest, x)?;
let mut terms = Vec::new();
for (i, c) in coeffs.iter().enumerate() {
if is_zero(*c) {
continue;
}
let i = i as i64;
let mut inner = Vec::new();
for j in 0..=i {
let t = t_square(ctx, h, lin.u, j, MAX_SPECIAL_STEPS)?;
let w = muln(
ctx,
&[
ctx.num(binom_k(i, j)),
powi_clean(ctx, neg(ctx, lin.a), i - j),
t,
],
);
inner.push(w);
}
let scale = muln(ctx, &[*c, powi_clean(ctx, lin.b, -(i + 1))]);
terms.push(muln(ctx, &[scale, addn(ctx, &inner)]));
}
Some(addn(ctx, &terms))
}
fn t_square<'a>(
ctx: &'a AtomArena<'a>,
h: Head,
u: Atom<'a>,
j: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || j < 0 {
return None;
}
let jp = j + 1;
let sq = ctx.pow(ctx.fun(h.name(), &[u]), ctx.num(2));
let lead = muln(ctx, &[powi_clean(ctx, u, jp), sq, q_frac(ctx, 1, jp)]);
let us = u_square(ctx, h, u, jp, budget - 1)?;
Some(addn(
ctx,
&[
lead,
neg(ctx, muln(ctx, &[ctx.num(2), q_frac(ctx, 1, jp), us])),
],
))
}
fn u_square<'a>(
ctx: &'a AtomArena<'a>,
h: Head,
u: Atom<'a>,
m: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || m < 1 {
return None;
}
let two_over_sqrt_pi = muln(ctx, &[ctx.num(2), inv_sqrt_pi(ctx)]);
Some(match h {
Head::Erf => muln(
ctx,
&[two_over_sqrt_pi, m_gauss(ctx, u, -1, m, budget - 1)?],
),
Head::Erfi => muln(ctx, &[two_over_sqrt_pi, m_gauss(ctx, u, 1, m, budget - 1)?]),
Head::Ci => ab_pair(ctx, u, false, true, m - 1, budget - 1)?.0,
Head::Si => ab_pair(ctx, u, false, false, m - 1, budget - 1)?.1,
Head::Chi => ab_pair(ctx, u, true, true, m - 1, budget - 1)?.0,
Head::Shi => ab_pair(ctx, u, true, false, m - 1, budget - 1)?.1,
_ => return None,
})
}
fn m_gauss<'a>(
ctx: &'a AtomArena<'a>,
u: Atom<'a>,
sigma: i64,
j: i64,
budget: usize,
) -> Option<Atom<'a>> {
if budget == 0 || j < 0 || (sigma != 1 && sigma != -1) {
return None;
}
let f = if sigma > 0 {
ctx.fun("erfi", &[u])
} else {
ctx.fun("erf", &[u])
};
let ex = exp_atom(ctx, muln(ctx, &[ctx.num(sigma), powi_clean(ctx, u, 2)]));
if j == 0 {
return Some(muln(
ctx,
&[q_frac(ctx, 1, 4), sqrt_pi(ctx), ctx.pow(f, ctx.num(2))],
));
}
if j == 1 {
let t1 = muln(ctx, &[f, ex, q_frac(ctx, 1, 2 * sigma)]);
let g = int_pow_exp_quad(ctx, u, 0, 2 * sigma, budget - 1)?;
let t2 = muln(ctx, &[ctx.num(sigma), inv_sqrt_pi(ctx), g]);
return Some(addn(ctx, &[t1, neg(ctx, t2)]));
}
let lead = muln(ctx, &[powi_clean(ctx, u, j - 1), f, ex]);
let rec = m_gauss(ctx, u, sigma, j - 2, budget - 1)?;
let g = int_pow_exp_quad(ctx, u, j - 1, 2 * sigma, budget - 1)?;
let two_over_sqrt_pi = muln(ctx, &[ctx.num(2), inv_sqrt_pi(ctx)]);
let inner = addn(
ctx,
&[
lead,
neg(ctx, muln(ctx, &[ctx.num(j - 1), rec])),
neg(ctx, muln(ctx, &[two_over_sqrt_pi, g])),
],
);
Some(muln(ctx, &[q_frac(ctx, 1, 2 * sigma), inner]))
}
fn w_pair<'a>(ctx: &'a AtomArena<'a>, u: Atom<'a>, hyp: bool, p: i64) -> (Atom<'a>, Atom<'a>) {
let sigma = if hyp { 1 } else { -1 };
let two_u = muln(ctx, &[ctx.num(2), u]);
let s2 = if hyp {
Kind::Sinh.atom(ctx, two_u)
} else {
Kind::Sin.atom(ctx, two_u)
};
let c2 = if hyp {
Kind::Cosh.atom(ctx, two_u)
} else {
Kind::Cos.atom(ctx, two_u)
};
let mut ws = muln(ctx, &[ctx.num(sigma), q_frac(ctx, 1, 2), c2]);
let mut wc = muln(ctx, &[q_frac(ctx, 1, 2), s2]);
for q in 1..=p.max(0) {
let nws = addn(
ctx,
&[
muln(
ctx,
&[powi_clean(ctx, u, q), ctx.num(sigma), q_frac(ctx, 1, 2), c2],
),
neg(
ctx,
muln(ctx, &[ctx.num(q), ctx.num(sigma), q_frac(ctx, 1, 2), wc]),
),
],
);
let nwc = addn(
ctx,
&[
muln(ctx, &[powi_clean(ctx, u, q), q_frac(ctx, 1, 2), s2]),
neg(ctx, muln(ctx, &[ctx.num(q), q_frac(ctx, 1, 2), ws])),
],
);
ws = nws;
wc = nwc;
}
(ws, wc)
}
fn ab_pair<'a>(
ctx: &'a AtomArena<'a>,
u: Atom<'a>,
hyp: bool,
cos_type: bool,
j: i64,
budget: usize,
) -> Option<(Atom<'a>, Atom<'a>)> {
if budget == 0 || j < 0 {
return None;
}
let sigma = if hyp { 1 } else { -1 };
let (s_kind, c_kind) = if hyp {
(Kind::Sinh, Kind::Cosh)
} else {
(Kind::Sin, Kind::Cos)
};
let two_u = muln(ctx, &[ctx.num(2), u]);
let f_name = match (hyp, cos_type) {
(false, true) => "Ci",
(false, false) => "Si",
(true, true) => "Chi",
(true, false) => "Shi",
};
let f = ctx.fun(f_name, &[u]);
if j == 0 {
let a0 = match (hyp, cos_type) {
(false, true) => addn(
ctx,
&[
muln(ctx, &[f, s_kind.atom(ctx, u)]),
neg(
ctx,
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Si", &[two_u])]),
),
],
),
(false, false) => addn(
ctx,
&[
muln(ctx, &[f, s_kind.atom(ctx, u)]),
neg(ctx, muln(ctx, &[q_frac(ctx, 1, 2), log_u(ctx, u)])),
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Ci", &[two_u])]),
],
),
(true, true) => addn(
ctx,
&[
muln(ctx, &[f, s_kind.atom(ctx, u)]),
neg(
ctx,
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Shi", &[two_u])]),
),
],
),
(true, false) => addn(
ctx,
&[
muln(ctx, &[f, s_kind.atom(ctx, u)]),
neg(
ctx,
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Chi", &[two_u])]),
),
muln(ctx, &[q_frac(ctx, 1, 2), log_u(ctx, u)]),
],
),
};
let b0 = match (hyp, cos_type) {
(false, true) => addn(
ctx,
&[
neg(ctx, muln(ctx, &[f, c_kind.atom(ctx, u)])),
muln(ctx, &[q_frac(ctx, 1, 2), log_u(ctx, u)]),
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Ci", &[two_u])]),
],
),
(false, false) => addn(
ctx,
&[
neg(ctx, muln(ctx, &[f, c_kind.atom(ctx, u)])),
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Si", &[two_u])]),
],
),
(true, true) => addn(
ctx,
&[
muln(ctx, &[f, c_kind.atom(ctx, u)]),
neg(ctx, muln(ctx, &[q_frac(ctx, 1, 2), log_u(ctx, u)])),
neg(
ctx,
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Chi", &[two_u])]),
),
],
),
(true, false) => addn(
ctx,
&[
muln(ctx, &[f, c_kind.atom(ctx, u)]),
neg(
ctx,
muln(ctx, &[q_frac(ctx, 1, 2), ctx.fun("Shi", &[two_u])]),
),
],
),
};
return Some((a0, b0));
}
let (ap, bp) = ab_pair(ctx, u, hyp, cos_type, j - 1, budget - 1)?;
let (ws, wc) = w_pair(ctx, u, hyp, j - 1);
let uj = powi_clean(ctx, u, j);
let uj_over_j = muln(ctx, &[uj, q_frac(ctx, 1, j)]);
let sc = muln(ctx, &[q_frac(ctx, 1, 2), ws]);
let c_sq = muln(ctx, &[q_frac(ctx, 1, 2), addn(ctx, &[uj_over_j, wc])]);
let s_sq = if hyp {
muln(
ctx,
&[q_frac(ctx, 1, 2), addn(ctx, &[wc, neg(ctx, uj_over_j)])],
)
} else {
muln(
ctx,
&[q_frac(ctx, 1, 2), addn(ctx, &[uj_over_j, neg(ctx, wc)])],
)
};
let (ta, tb) = if cos_type { (sc, c_sq) } else { (s_sq, sc) };
let a = addn(
ctx,
&[
muln(ctx, &[uj, f, s_kind.atom(ctx, u)]),
neg(ctx, muln(ctx, &[ctx.num(j), bp])),
neg(ctx, ta),
],
);
let b_inner = addn(
ctx,
&[
muln(ctx, &[uj, f, c_kind.atom(ctx, u)]),
neg(ctx, muln(ctx, &[ctx.num(j), ap])),
neg(ctx, tb),
],
);
let b = muln(ctx, &[ctx.num(sigma), b_inner]);
Some((a, b))
}
fn log_u<'a>(ctx: &'a AtomArena<'a>, u: Atom<'a>) -> Atom<'a> {
ctx.fun("log", &[u])
}
#[cfg(test)]
mod tests {
use ocas_core::arena::Arena;
use super::*;
fn int<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>) -> Option<Atom<'a>> {
special_integrate(ctx, expr, ctx.var("x"))
}
#[test]
fn exp_neg_x_squared_gives_erf() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("exp", &[ctx.mul(&[ctx.num(-1), ctx.pow(x, ctx.num(2))])]);
let r = int(&ctx, expr).expect("erf form");
assert!(r.to_string().contains("erf"), "got {r}");
}
#[test]
fn exp_x_over_x_gives_ei() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.fun("exp", &[x]), ctx.pow(x, ctx.num(-1))]);
let r = int(&ctx, expr).expect("Ei form");
assert_eq!(r.to_string(), "Ei(x)");
}
#[test]
fn sin_x_over_x_gives_si() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.fun("sin", &[x]), ctx.pow(x, ctx.num(-1))]);
let r = int(&ctx, expr).expect("Si form");
assert_eq!(r.to_string(), "Si(x)");
}
#[test]
fn cos_x_over_x_gives_ci() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.fun("cos", &[x]), ctx.pow(x, ctx.num(-1))]);
let r = int(&ctx, expr).expect("Ci form");
assert_eq!(r.to_string(), "Ci(x)");
}
#[test]
fn sin_x_squared_gives_fresnel() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("sin", &[ctx.pow(x, ctx.num(2))]);
let r = int(&ctx, expr).expect("Fresnel form");
assert!(r.to_string().contains("fresnels"), "got {r}");
}
#[test]
fn unmatched_returns_none() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
assert!(int(&ctx, ctx.fun("exp", &[x])).is_none());
let _ = Symbol::new("x");
}
const EULER_GAMMA: f64 = 0.577_215_664_901_532_9;
fn quad(f: &dyn Fn(f64) -> f64, a: f64, b: f64, n: usize) -> f64 {
let n = if n % 2 == 1 { n + 1 } else { n };
let h = (b - a) / n as f64;
let mut s = f(a) + f(b);
for i in 1..n {
s += if i % 2 == 1 { 4.0 } else { 2.0 } * f(a + i as f64 * h);
}
s * h / 3.0
}
fn sin_over(t: f64) -> f64 {
if t == 0.0 { 1.0 } else { t.sin() / t }
}
fn sinh_over(t: f64) -> f64 {
if t == 0.0 { 1.0 } else { t.sinh() / t }
}
fn expm1_over(t: f64) -> f64 {
if t == 0.0 { 1.0 } else { (t.exp() - 1.0) / t }
}
fn cosm1_over(t: f64) -> f64 {
if t == 0.0 { 0.0 } else { (t.cos() - 1.0) / t }
}
fn coshm1_over(t: f64) -> f64 {
if t == 0.0 { 0.0 } else { (t.cosh() - 1.0) / t }
}
fn erf_o(v: f64) -> f64 {
2.0 / std::f64::consts::PI.sqrt() * quad(&|t| (-t * t).exp(), 0.0, v, 800)
}
fn erfi_o(v: f64) -> f64 {
2.0 / std::f64::consts::PI.sqrt() * quad(&|t| (t * t).exp(), 0.0, v, 800)
}
fn si_o(v: f64) -> f64 {
quad(&sin_over, 0.0, v, 800)
}
fn ci_o(v: f64) -> f64 {
EULER_GAMMA + v.abs().ln() + quad(&cosm1_over, 0.0, v.abs(), 800)
}
fn shi_o(v: f64) -> f64 {
quad(&sinh_over, 0.0, v, 800)
}
fn chi_o(v: f64) -> f64 {
EULER_GAMMA + v.abs().ln() + quad(&coshm1_over, 0.0, v.abs(), 800)
}
fn ei_o(v: f64) -> f64 {
EULER_GAMMA + v.abs().ln() + quad(&expm1_over, 0.0, v, 800)
}
fn factorial_o(n: u32) -> f64 {
(1..=n).map(f64::from).product()
}
fn ei_order_o(n: i64, z: f64) -> Option<f64> {
if n == 0 {
return if z == 0.0 { None } else { Some((-z).exp() / z) };
}
if n < 0 {
let m = (-n) as u32;
let mut sum = 0.0;
for k in 0..=m {
sum += z.powi(k as i32 - m as i32 - 1) / factorial_o(k);
}
return Some((-z).exp() * factorial_o(m) * sum);
}
if z <= 0.0 {
return None;
}
let mut e = -ei_o(-z);
for k in 1..n {
e = ((-z).exp() - z * e) / k as f64;
}
Some(e)
}
fn eval(e: Atom<'_>, env: &[(Symbol, f64)]) -> Option<f64> {
match e.node() {
AtomNode::Num(n) => Some(*n as f64),
AtomNode::Var(v) => {
if v.as_str() == "pi" {
return Some(std::f64::consts::PI);
}
env.iter().find(|(s, _)| s == v).map(|(_, val)| *val)
}
AtomNode::Add(args) => {
let mut acc = 0.0;
for a in args.iter() {
acc += eval(*a, env)?;
}
Some(acc)
}
AtomNode::Mul(args) => {
let mut acc = 1.0;
for a in args.iter() {
acc *= eval(*a, env)?;
}
Some(acc)
}
AtomNode::Pow(b, ex) => {
let base = eval(*b, env)?;
let e = eval(*ex, env)?;
if base == 0.0 && e < 0.0 {
return None;
}
let v = base.powf(e);
if v.is_finite() { Some(v) } else { None }
}
AtomNode::Fun(name, args) => {
if name.as_str() == "Ei" && args.len() == 2 {
let n = eval(args[0], env)?;
let z = eval(args[1], env)?;
return ei_order_o(n as i64, z);
}
let v = eval(*args.first()?, env)?;
let out = match name.as_str() {
"sin" => v.sin(),
"cos" => v.cos(),
"sinh" => v.sinh(),
"cosh" => v.cosh(),
"exp" => v.exp(),
"log" => v.abs().ln(),
"sqrt" => {
if v < 0.0 {
return None;
}
v.sqrt()
}
"erf" => erf_o(v),
"erfi" => erfi_o(v),
"Si" => si_o(v),
"Ci" => {
if v == 0.0 {
return None;
}
ci_o(v)
}
"Shi" => {
if v == 0.0 {
return None;
}
shi_o(v)
}
"Chi" => {
if v == 0.0 {
return None;
}
chi_o(v)
}
"Ei" if args.len() == 1 => {
if v == 0.0 {
return None;
}
ei_o(v)
}
_ => return None,
};
if out.is_finite() { Some(out) } else { None }
}
}
}
fn diff_local<'a>(ctx: &'a AtomArena<'a>, e: Atom<'a>, var: Symbol) -> Atom<'a> {
match e.node() {
AtomNode::Num(_) => ctx.num(0),
AtomNode::Var(v) => {
if *v == var {
ctx.num(1)
} else {
ctx.num(0)
}
}
AtomNode::Add(args) => {
let d: Vec<Atom<'a>> = args.iter().map(|a| diff_local(ctx, *a, var)).collect();
ctx.add(&d)
}
AtomNode::Mul(args) => {
let mut terms = Vec::with_capacity(args.len());
for i in 0..args.len() {
let mut factors = Vec::with_capacity(args.len());
for (j, a) in args.iter().enumerate() {
if i == j {
factors.push(diff_local(ctx, *a, var));
} else {
factors.push(*a);
}
}
terms.push(ctx.mul(&factors));
}
ctx.add(&terms)
}
AtomNode::Pow(b, ex) => {
let (base, exp) = (*b, *ex);
let db = diff_local(ctx, base, var);
let de = diff_local(ctx, exp, var);
let t1 = ctx.mul(&[exp, ctx.pow(base, ctx.add(&[exp, ctx.num(-1)])), db]);
let t2 = ctx.mul(&[ctx.pow(base, exp), ctx.fun("log", &[base]), de]);
ctx.add(&[t1, t2])
}
AtomNode::Fun(name, args) => {
let a0 = if name.as_str() == "Ei" && args.len() == 2 {
args[1]
} else {
*args.first().unwrap()
};
let d0 = diff_local(ctx, a0, var);
let inner = match name.as_str() {
"sin" => ctx.fun("cos", &[a0]),
"cos" => ctx.mul(&[ctx.num(-1), ctx.fun("sin", &[a0])]),
"sinh" => ctx.fun("cosh", &[a0]),
"cosh" => ctx.fun("sinh", &[a0]),
"exp" => ctx.fun("exp", &[a0]),
"log" => ctx.pow(a0, ctx.num(-1)),
"sqrt" => ctx.pow(ctx.mul(&[ctx.num(2), ctx.fun("sqrt", &[a0])]), ctx.num(-1)),
"erf" => ctx.mul(&[
ctx.num(2),
inv_sqrt_pi(ctx),
ctx.fun("exp", &[ctx.mul(&[ctx.num(-1), ctx.pow(a0, ctx.num(2))])]),
]),
"erfi" => ctx.mul(&[
ctx.num(2),
inv_sqrt_pi(ctx),
ctx.fun("exp", &[ctx.pow(a0, ctx.num(2))]),
]),
"Si" => ctx.mul(&[ctx.fun("sin", &[a0]), ctx.pow(a0, ctx.num(-1))]),
"Ci" => ctx.mul(&[ctx.fun("cos", &[a0]), ctx.pow(a0, ctx.num(-1))]),
"Shi" => ctx.mul(&[ctx.fun("sinh", &[a0]), ctx.pow(a0, ctx.num(-1))]),
"Chi" => ctx.mul(&[ctx.fun("cosh", &[a0]), ctx.pow(a0, ctx.num(-1))]),
"Ei" if args.len() == 1 => {
ctx.mul(&[ctx.fun("exp", &[a0]), ctx.pow(a0, ctx.num(-1))])
}
"Ei" if args.len() == 2 => {
let AtomNode::Num(n) = args[0].node() else {
panic!("non-numeric Ei order in emitted form");
};
if *n == 0 {
ctx.mul(&[
ctx.fun("exp", &[ctx.mul(&[ctx.num(-1), a0])]),
ctx.pow(a0, ctx.num(-1)),
])
} else {
ctx.mul(&[ctx.num(-1), ctx.fun("Ei", &[ctx.num(n - 1), a0])])
}
}
_ => panic!("diff_local: unsupported head {}", name.as_str()),
};
ctx.mul(&[inner, d0])
}
}
}
fn check(s: &str, consts: &[(&str, f64)], samples: &[f64]) {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let e = ocas_parse::parse(&ctx, s).unwrap_or_else(|_| panic!("parse: {s}"));
let x = ctx.var("x");
let r = special_integrate(&ctx, e, x).unwrap_or_else(|| panic!("declined: {s}"));
assert!(!r.to_string().contains("Integral("), "residue for {s}: {r}");
let d = diff_local(&ctx, r, Symbol::new("x"));
let base: Vec<(Symbol, f64)> = consts.iter().map(|(n, v)| (Symbol::new(n), *v)).collect();
let mut checked = 0usize;
let mut worst = 0.0f64;
for &xv in samples {
let mut env = base.clone();
env.push((Symbol::new("x"), xv));
let lhs = match eval(d, &env) {
Some(v) => v,
None => continue,
};
let rhs = eval(e, &env).expect("integrand evaluated");
let tol = 1e-5 * rhs.abs().max(1.0);
assert!(
(lhs - rhs).abs() < tol,
"{s} at x={xv}: d/dx F = {lhs}, integrand = {rhs}\n F = {r}"
);
checked += 1;
worst = worst.max((lhs - rhs).abs() / rhs.abs().max(1.0));
}
assert!(checked >= 2, "only {checked} usable samples for {s} -> {r}");
println!(" checked={checked} worst_rel={worst:e} {s} -> {r}");
}
fn declines(s: &str) {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let e = ocas_parse::parse(&ctx, s).unwrap_or_else(|_| panic!("parse: {s}"));
assert!(
special_integrate(&ctx, e, ctx.var("x")).is_none(),
"expected decline for {s}"
);
}
#[test]
fn plain_family_polynomial_cases() {
check("Ci(b*x)", &[("b", 1.5)], &[0.6, 1.1, 1.9]);
check("x*Chi(b*x)", &[("b", 1.3)], &[0.5, 1.0, 1.7]);
check(
"x^3*Shi(a + b*x)",
&[("a", 0.4), ("b", 1.2)],
&[0.4, 0.9, 1.6],
);
check("x^4*erf(b*x)", &[("b", 1.1)], &[0.5, 1.0, 1.6]);
check("x^2*Si(b*x)", &[("b", 1.4)], &[0.5, 1.1, 1.8]);
check("x*erfi(b*x)", &[("b", 0.9)], &[0.5, 1.0, 1.5]);
check("x^2*Ei(b*x)", &[("b", 1.2)], &[0.6, 1.1, 1.7]);
check(
"x^2*Shi(a + b*x)",
&[("a", 0.3), ("b", 1.1)],
&[0.5, 1.0, 1.5],
);
}
#[test]
fn order_family_cases() {
check("x^4*Ei(-2, b*x)", &[("b", 1.2)], &[0.6, 1.1, 1.7]);
check("x^3*Ei(-1, b*x)", &[("b", 1.3)], &[0.6, 1.1, 1.7]);
check(
"Ei(1, a + b*x)",
&[("a", 0.5), ("b", 1.2)],
&[0.4, 0.9, 1.6],
);
check("Ei(2, b*x)", &[("b", 1.2)], &[0.6, 1.1, 1.7]);
check("x^2*Ei(3, b*x)", &[("b", 1.1)], &[0.6, 1.1, 1.7]);
check(
"Ei(-2, a + b*x)",
&[("a", 0.4), ("b", 1.2)],
&[0.4, 0.9, 1.5],
);
}
#[test]
fn power_descent_cases() {
check("Ei(b*x)/x^4", &[("b", 1.2)], &[0.7, 1.2, 1.9]);
check("Ei(-3, b*x)/x^3", &[("b", 1.2)], &[0.7, 1.2, 1.9]);
check("Ei(1, b*x)/x^3", &[("b", 1.2)], &[0.7, 1.2, 1.9]);
check("Ei(b*x)/x^2", &[("b", 1.3)], &[0.7, 1.2, 1.9]);
check("Ei(2, b*x)/x^3", &[("b", 1.1)], &[0.7, 1.2, 1.9]);
check("Ei(-1, b*x)/x^4", &[("b", 1.1)], &[0.7, 1.2, 1.9]);
check("3*Ei(b*x)/x^3", &[("b", 1.1)], &[0.7, 1.2, 1.9]);
}
#[test]
fn square_family_cases() {
check("erfi(b*x)^2", &[("b", 1.2)], &[0.5, 1.0, 1.5]);
check("x^2*Ci(b*x)^2", &[("b", 1.2)], &[0.7, 1.2, 1.8]);
check(
"Chi(a + b*x)^2",
&[("a", 0.5), ("b", 1.2)],
&[0.4, 0.9, 1.5],
);
check(
"(c + d*x)*erf(a + b*x)^2",
&[("a", 0.4), ("b", 1.2), ("c", 0.7), ("d", 1.3)],
&[0.4, 0.9, 1.5],
);
check(
"x*Si(a + b*x)^2",
&[("a", 0.5), ("b", 1.2)],
&[0.4, 0.9, 1.5],
);
check("erf(b*x)^2", &[("b", 1.1)], &[0.5, 1.0, 1.5]);
check("x*Ci(b*x)^2", &[("b", 1.2)], &[0.7, 1.2, 1.8]);
check(
"Shi(a + b*x)^2",
&[("a", 0.5), ("b", 1.1)],
&[0.4, 0.9, 1.5],
);
check("x^3*Shi(b*x)^2", &[("b", 1.1)], &[0.6, 1.1, 1.6]);
check(
"x^2*Ci(a + b*x)^2",
&[("a", 0.4), ("b", 1.2)],
&[0.5, 1.0, 1.5],
);
check(
"x^3*Chi(a + b*x)^2",
&[("a", 0.3), ("b", 1.1)],
&[0.5, 1.0, 1.5],
);
check(
"x^3*erf(a + b*x)^2",
&[("a", 0.4), ("b", 1.2)],
&[0.4, 0.9, 1.5],
);
check(
"x^2*Si(a + b*x)^2",
&[("a", 0.4), ("b", 1.2)],
&[0.4, 0.9, 1.5],
);
check("x^2*erfi(b*x)^2", &[("b", 1.1)], &[0.5, 1.0, 1.5]);
}
#[test]
fn unscoped_shapes_decline() {
declines("erf(a + b*x)^2/(c + d*x)");
declines("Ei(n, a + b*x)/(c + d*x)^2");
declines("fresnels(b*x)/x^10");
declines("fresnels(b*x)/x^8");
declines("Ci(b*x)*Si(b*x)");
declines("Ci(b*x^2)");
declines("Ci(b*x)/x");
declines("Ei(b*x)/x");
declines("exp(x)");
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let e = ocas_parse::parse(&ctx, "sin(x)/x").expect("parse");
assert_eq!(
special_integrate(&ctx, e, ctx.var("x"))
.expect("Si entry")
.to_string(),
"Si(x)"
);
}
#[test]
fn budget_and_builder_edges() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let u = ctx.var("u");
assert!(int_pow_exp(&ctx, u, ctx.num(1), 3, 0).is_none());
assert!(int_pow_exp_quad(&ctx, u, 3, -1, 0).is_none());
assert!(int_pow_trig(&ctx, u, 3, Kind::Sin, 0).is_none());
assert!(int_exp_over_pow(&ctx, u, ctx.num(1), 2, 0).is_none());
assert!(d_ei(&ctx, u, 1).is_none());
assert!(d_order(&ctx, u, 2, 3, MAX_SPECIAL_STEPS).is_none());
assert!(m_gauss(&ctx, u, 0, 2, MAX_SPECIAL_STEPS).is_none());
assert_eq!(powi_clean(&ctx, ctx.num(0), 0).to_string(), "1");
assert_eq!(powi_clean(&ctx, ctx.num(0), 2).to_string(), "0");
assert_eq!(powi_clean(&ctx, ctx.num(-1), 3).to_string(), "-1");
assert_eq!(powi_clean(&ctx, ctx.num(-1), 4).to_string(), "1");
assert_eq!(muln(&ctx, &[ctx.num(1), ctx.num(1)]).to_string(), "1");
assert_eq!(muln(&ctx, &[ctx.num(0), ctx.var("x")]).to_string(), "0");
assert_eq!(addn(&ctx, &[ctx.num(0), ctx.num(0)]).to_string(), "0");
let xv = ctx.var("x");
assert!(poly_coeffs(&ctx, ctx.pow(xv, ctx.num(-1)), Symbol::new("x")).is_none());
assert!(poly_coeffs(&ctx, ctx.pow(xv, ctx.num(9)), Symbol::new("x")).is_none());
assert_eq!(binom_k(5, 2), 10);
assert_eq!(binom_k(0, 0), 1);
assert_eq!(binom_k(3, 4), 0);
}
}