use ocas_atom::normalize::normalize;
use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use super::rules::rat_of;
use super::{
contains_integral, contains_symbol, int_pow, integrate_raw, inv, is_constant, linear_form,
node_count, pick_subst_symbol, rat_atom, replace_symbol,
};
const MAX_NODES: usize = 200;
const MAX_SUBST_NODES: usize = 300;
const MAX_REDUCED_NODES: usize = 600;
const MAX_POWER_K: i64 = 6;
const MAX_SUBST_K: i64 = 4;
const MAX_POLY_DEG: usize = 4;
const MAX_HYP_EXP: usize = 8;
const MAX_CANCEL_BUDGET: usize = 800;
const MAX_KERNEL_POW: i64 = 6;
const MAX_COFACTOR_DEG: usize = 6;
const MAX_RECIP_EXP: i64 = 8;
const MAX_KERNEL_NODES: usize = 400;
const MAX_RAT_Q: i64 = 4;
#[derive(Clone, Copy, PartialEq, Eq)]
enum InvFun {
Asin,
Acos,
Atan,
Asinh,
Acosh,
Atanh,
}
impl InvFun {
fn name(self) -> &'static str {
match self {
InvFun::Asin => "asin",
InvFun::Acos => "acos",
InvFun::Atan => "atan",
InvFun::Asinh => "asinh",
InvFun::Acosh => "acosh",
InvFun::Atanh => "atanh",
}
}
}
fn inv_fun(name: &str) -> Option<InvFun> {
Some(match name {
"asin" => InvFun::Asin,
"acos" => InvFun::Acos,
"atan" => InvFun::Atan,
"asinh" => InvFun::Asinh,
"acosh" => InvFun::Acosh,
"atanh" => InvFun::Atanh,
_ => return None,
})
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum KernelExp {
InvSqrt,
Recip,
}
fn family_kernel_exp(f: InvFun) -> KernelExp {
match f {
InvFun::Asin | InvFun::Acos | InvFun::Asinh | InvFun::Acosh => KernelExp::InvSqrt,
InvFun::Atan | InvFun::Atanh => KernelExp::Recip,
}
}
fn family_sign(f: InvFun) -> i64 {
match f {
InvFun::Acos => -1,
_ => 1,
}
}
fn kernel_sum<'a>(
ctx: &'a AtomArena<'a>,
f: InvFun,
sigma: Atom<'a>,
var: Symbol,
) -> (Atom<'a>, i64) {
let x = ctx.var(var.as_str());
let s2x2 = ctx.mul(&[int_pow(ctx, sigma, 2), int_pow(ctx, x, 2)]);
match f {
InvFun::Asin | InvFun::Acos | InvFun::Atanh => (
normalize(ctx, ctx.add(&[ctx.num(1), ctx.mul(&[ctx.num(-1), s2x2])])),
1,
),
InvFun::Atan | InvFun::Asinh => (normalize(ctx, ctx.add(&[ctx.num(1), s2x2])), 1),
InvFun::Acosh => (normalize(ctx, ctx.add(&[ctx.num(-1), s2x2])), -1),
}
}
pub(crate) fn folds_to_zero<'a>(ctx: &'a AtomArena<'a>, e: Atom<'a>) -> bool {
let n = normalize(ctx, e);
if matches!(n.node(), AtomNode::Num(0)) {
return true;
}
matches!(
crate::ode::util::collect_terms(ctx, n).node(),
AtomNode::Num(0)
)
}
fn is_zero_atom<'a>(ctx: &'a AtomArena<'a>, e: Atom<'a>) -> bool {
folds_to_zero(ctx, e)
}
struct InvPow<'a> {
f: InvFun,
arg: Atom<'a>,
#[allow(dead_code)] a: Atom<'a>,
b: Atom<'a>,
k: i64,
base: Atom<'a>,
}
fn match_inv_base<'a>(
ctx: &'a AtomArena<'a>,
base: Atom<'a>,
var: Symbol,
) -> Option<(InvFun, Atom<'a>, Atom<'a>, Atom<'a>)> {
match base.node() {
AtomNode::Fun(name, args) if args.len() == 1 => {
let f = inv_fun(name.as_str())?;
Some((f, args[0], ctx.num(0), ctx.num(1)))
}
AtomNode::Mul(args) => {
let mut coeff: Vec<Atom<'a>> = Vec::new();
let mut found: Option<(InvFun, Atom<'a>)> = None;
for fa in args.iter() {
match fa.node() {
AtomNode::Fun(name, fargs) if fargs.len() == 1 => {
let f = inv_fun(name.as_str())?;
if found.is_some() {
return None;
}
found = Some((f, fargs[0]));
}
_ if is_constant(*fa, var) => coeff.push(*fa),
_ => return None,
}
}
let (f, arg) = found?;
let b = if coeff.is_empty() {
ctx.num(1)
} else {
normalize(ctx, ctx.mul(&coeff))
};
Some((f, arg, ctx.num(0), b))
}
AtomNode::Add(args) => {
let mut consts: Vec<Atom<'a>> = Vec::new();
let mut found: Option<(InvFun, Atom<'a>, Atom<'a>, Atom<'a>)> = None;
for t in args.iter() {
if is_constant(*t, var) {
consts.push(*t);
continue;
}
let term = match_inv_base(ctx, *t, var)?;
if found.is_some() {
return None;
}
found = Some(term);
}
let (f, arg, _zero, b) = found?;
let a = if consts.is_empty() {
ctx.num(0)
} else {
normalize(ctx, ctx.add(&consts))
};
Some((f, arg, a, b))
}
_ => None,
}
}
fn match_inv_power<'a>(
ctx: &'a AtomArena<'a>,
factor: Atom<'a>,
var: Symbol,
k_cap: i64,
) -> Option<InvPow<'a>> {
if let AtomNode::Pow(b, e) = factor.node()
&& let AtomNode::Num(k) = e.node()
{
if !(1..=k_cap).contains(k) {
return None;
}
let (f, arg, a, bb) = match_inv_base(ctx, *b, var)?;
return Some(InvPow {
f,
arg,
a,
b: bb,
k: *k,
base: *b,
});
}
let (f, arg, a, bb) = match_inv_base(ctx, factor, var)?;
Some(InvPow {
f,
arg,
a,
b: bb,
k: 1,
base: factor,
})
}
fn match_kernel(factor: Atom<'_>) -> Option<(Atom<'_>, KernelExp)> {
let AtomNode::Pow(b, e) = factor.node() else {
return None;
};
if let AtomNode::Fun(name, args) = b.node()
&& name.as_str() == "sqrt"
&& args.len() == 1
&& matches!(e.node(), AtomNode::Num(-1))
{
return Some((args[0], KernelExp::InvSqrt));
}
match rat_of(*e)? {
(-1, 2) => Some((*b, KernelExp::InvSqrt)),
(-1, 1) => Some((*b, KernelExp::Recip)),
_ => None,
}
}
fn radicand_normalization<'a>(
ctx: &'a AtomArena<'a>,
r: Atom<'a>,
e_sum: Atom<'a>,
e0: i64,
var: Symbol,
) -> Option<Atom<'a>> {
let r0 = normalize(ctx, replace_symbol(ctx, r, var, ctx.num(0)));
let s_atom = if e0 == 1 {
r0
} else {
normalize(ctx, ctx.mul(&[ctx.num(-1), r0]))
};
if is_zero_atom(ctx, s_atom) {
return None;
}
let diff = ctx.add(&[r, ctx.mul(&[ctx.num(-1), s_atom, e_sum])]);
if !folds_to_zero(ctx, diff) {
return None;
}
Some(s_atom)
}
fn sqrt_norm_positive(s_atom: Atom<'_>) -> bool {
match rat_of(s_atom) {
Some((p, q)) => (p > 0) == (q > 0),
None => true,
}
}
pub(crate) fn integrate_kernel_power<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if is_constant(expr, var) || matches!(expr.node(), AtomNode::Add(_)) {
return None;
}
let factors: Vec<Atom<'a>> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![expr],
};
let mut rest: Vec<Atom<'a>> = Vec::new();
let mut inv_f: Option<InvPow<'a>> = None;
let mut kern: Option<(Atom<'a>, KernelExp)> = None;
for f in factors {
if is_constant(f, var) {
rest.push(f);
continue;
}
if inv_f.is_none()
&& let Some(m) = match_inv_power(ctx, f, var, MAX_POWER_K)
{
inv_f = Some(m);
continue;
}
if kern.is_none()
&& let Some(kd) = match_kernel(f)
{
kern = Some(kd);
continue;
}
return None;
}
let m = inv_f?;
let (r, ke) = kern?;
if ke != family_kernel_exp(m.f) {
return None;
}
let (sigma, intercept) = linear_form(ctx, m.arg, var)?;
if !is_zero_atom(ctx, intercept) || is_zero_atom(ctx, sigma) {
return None;
}
let (e_sum, e0) = kernel_sum(ctx, m.f, sigma, var);
let s_atom = radicand_normalization(ctx, r, e_sum, e0, var)?;
if ke == KernelExp::InvSqrt && !sqrt_norm_positive(s_atom) {
return None;
}
let s_factor = match ke {
KernelExp::Recip => inv(ctx, s_atom),
KernelExp::InvSqrt => ctx.pow(s_atom, rat_atom(ctx, -1, 2)),
};
let mut out = rest;
if family_sign(m.f) < 0 {
out.push(ctx.num(-1));
}
out.push(s_factor);
out.push(int_pow(ctx, m.base, m.k + 1));
let denom = normalize(ctx, ctx.mul(&[m.b, sigma, ctx.num(m.k + 1)]));
out.push(inv(ctx, denom));
Some(normalize(ctx, ctx.mul(&out)))
}
fn poly_deg(expr: Atom<'_>, var: Symbol) -> Option<usize> {
match expr.node() {
AtomNode::Num(_) => Some(0),
AtomNode::Var(v) => Some(if *v == var { 1 } else { 0 }),
AtomNode::Add(args) => args
.iter()
.try_fold(0usize, |d, a| Some(d.max(poly_deg(*a, var)?))),
AtomNode::Mul(args) => args
.iter()
.try_fold(0usize, |d, a| Some(d + poly_deg(*a, var)?)),
AtomNode::Pow(b, e) => {
let AtomNode::Num(n) = e.node() else {
return None;
};
if *n < 0 {
return None;
}
Some(poly_deg(*b, var)?.checked_mul(*n as usize)?)
}
AtomNode::Fun(_, _) => None,
}
}
fn match_sqrt_pow(factor: Atom<'_>) -> Option<(Atom<'_>, i64)> {
match factor.node() {
AtomNode::Fun(name, args) if name.as_str() == "sqrt" && args.len() == 1 => {
Some((args[0], 1))
}
AtomNode::Pow(b, e) => {
if let AtomNode::Fun(name, args) = b.node()
&& name.as_str() == "sqrt"
&& args.len() == 1
&& let AtomNode::Num(n) = e.node()
{
return match n {
1 | 3 => Some((args[0], *n)),
_ => None,
};
}
let (p, q) = rat_of(*e)?;
if q == 2 && matches!(p, 1 | 3) {
Some((*b, p))
} else {
None
}
}
_ => None,
}
}
pub(crate) fn integrate_invhyp_subst<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if is_constant(expr, var) || matches!(expr.node(), AtomNode::Add(_)) {
return None;
}
let factors: Vec<Atom<'a>> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![expr],
};
let mut rest: Vec<Atom<'a>> = Vec::new();
let mut inv_f: Option<InvPow<'a>> = None;
let mut polys: Vec<Atom<'a>> = Vec::new();
let mut sqrt_f: Option<(Atom<'a>, i64)> = None;
let mut total_deg = 0usize;
for f in factors {
if is_constant(f, var) {
rest.push(f);
continue;
}
if inv_f.is_none()
&& let Some(m) = match_inv_power(ctx, f, var, MAX_SUBST_K)
{
if !matches!(m.f, InvFun::Asinh | InvFun::Acosh) {
return None;
}
inv_f = Some(m);
continue;
}
if sqrt_f.is_none()
&& let Some(sp) = match_sqrt_pow(f)
{
sqrt_f = Some(sp);
continue;
}
if let Some(d) = poly_deg(f, var)
&& d >= 1
&& total_deg + d <= MAX_POLY_DEG
{
total_deg += d;
polys.push(f);
continue;
}
return None;
}
let m = inv_f?;
let (sigma, rho) = linear_form(ctx, m.arg, var)?;
if is_zero_atom(ctx, sigma) {
return None;
}
if sqrt_f.is_some() && !is_zero_atom(ctx, rho) {
return None;
}
let t_sym = pick_subst_symbol(expr, var)?;
let t = ctx.var(t_sym.as_str());
let (g_name, gp_name) = match m.f {
InvFun::Asinh => ("sinh", "cosh"),
InvFun::Acosh => ("cosh", "sinh"),
_ => return None,
};
let g_t = ctx.fun(g_name, &[t]);
let gp_t = ctx.fun(gp_name, &[t]);
let x_of_t = {
let num_t = if is_zero_atom(ctx, rho) {
g_t
} else {
ctx.add(&[g_t, ctx.mul(&[ctx.num(-1), rho])])
};
ctx.mul(&[num_t, inv(ctx, sigma)])
};
let mut t_factors = rest;
let base_t = normalize(ctx, ctx.add(&[m.a, ctx.mul(&[m.b, t])]));
t_factors.push(int_pow(ctx, base_t, m.k));
for p in &polys {
t_factors.push(replace_symbol(ctx, *p, var, x_of_t));
}
t_factors.push(ctx.mul(&[gp_t, inv(ctx, sigma)]));
if let Some((r, mpow)) = sqrt_f {
let (e_sum, e0) = kernel_sum(ctx, m.f, sigma, var);
let s_atom = radicand_normalization(ctx, r, e_sum, e0, var)?;
if !sqrt_norm_positive(s_atom) {
return None;
}
t_factors.push(ctx.pow(s_atom, rat_atom(ctx, mpow, 2)));
t_factors.push(int_pow(ctx, gp_t, mpow));
}
let product = normalize(ctx, ctx.mul(&t_factors));
let expanded = crate::expand::expand_bounded(ctx, product).unwrap_or(product);
let folded = crate::ode::util::collect_terms(ctx, normalize(ctx, expanded));
if node_count(folded) > MAX_SUBST_NODES {
return None;
}
let terms: Vec<Atom<'a>> = match folded.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![folded],
};
let mut reduced: Vec<Atom<'a>> = Vec::new();
for term in terms {
let (coeff, j, na, nb) = scan_hyp_term(ctx, term, t_sym)?;
reduce_hyp_term(ctx, &coeff, j, na, nb, t, &mut reduced)?;
}
let t_form = normalize(ctx, ctx.add(&reduced));
if node_count(t_form) > MAX_REDUCED_NODES {
return None;
}
let back = ctx.fun(m.f.name(), &[m.arg]);
let chain = integrate_raw(ctx, t_form, t_sym, 0, true, 0, 0);
if !contains_integral(chain) {
return Some(normalize(ctx, replace_symbol(ctx, chain, t_sym, back)));
}
let rterms: Vec<Atom<'a>> = match t_form.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![t_form],
};
let mut out = Vec::with_capacity(rterms.len());
for term in rterms {
out.push(integrate_hyp_term(ctx, term, t_sym)?);
}
let res_t = normalize(ctx, ctx.add(&out));
Some(normalize(ctx, replace_symbol(ctx, res_t, t_sym, back)))
}
fn binom(n: i64, mut k: i64) -> i64 {
if k < 0 || k > n {
return 0;
}
if k > n - k {
k = n - k;
}
let mut r: i64 = 1;
for i in 0..k {
r = r.saturating_mul(n - i) / (i + 1);
}
r
}
fn is_t_var(e: Atom<'_>, t_sym: Symbol) -> bool {
matches!(e.node(), AtomNode::Var(v) if *v == t_sym)
}
fn scan_hyp_term<'a>(
ctx: &'a AtomArena<'a>,
term: Atom<'a>,
t_sym: Symbol,
) -> Option<(Vec<Atom<'a>>, i64, usize, usize)> {
let _ = ctx;
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![term],
};
let mut coeff: Vec<Atom<'a>> = Vec::new();
let mut j = 0i64;
let mut na = 0usize;
let mut nb = 0usize;
for f in factors {
if is_constant(f, t_sym) {
coeff.push(f);
continue;
}
match f.node() {
AtomNode::Var(v) if *v == t_sym => j += 1,
AtomNode::Pow(b, e) if matches!(b.node(), AtomNode::Var(v) if *v == t_sym) => {
let AtomNode::Num(n) = e.node() else {
return None;
};
if *n < 0 {
return None;
}
j += n;
}
AtomNode::Fun(name, args) if args.len() == 1 && is_t_var(args[0], t_sym) => {
match name.as_str() {
"cosh" => na += 1,
"sinh" => nb += 1,
_ => return None,
}
}
AtomNode::Pow(b, e) => {
let AtomNode::Fun(name, args) = b.node() else {
return None;
};
if args.len() != 1 || !is_t_var(args[0], t_sym) {
return None;
}
let AtomNode::Num(n) = e.node() else {
return None;
};
if *n <= 0 {
return None;
}
match name.as_str() {
"cosh" => na += *n as usize,
"sinh" => nb += *n as usize,
_ => return None,
}
}
_ => return None,
}
}
if j > MAX_SUBST_K || na + nb > MAX_HYP_EXP {
return None;
}
Some((coeff, j, na, nb))
}
fn reduce_hyp_term<'a>(
ctx: &'a AtomArena<'a>,
coeff: &[Atom<'a>],
j: i64,
na: usize,
nb: usize,
t: Atom<'a>,
out: &mut Vec<Atom<'a>>,
) -> Option<()> {
if na == 0 && nb == 0 {
let mut fs = coeff.to_vec();
if j > 0 {
fs.push(int_pow(ctx, t, j));
}
out.push(normalize(ctx, ctx.mul(&fs)));
return Some(());
}
let nsum = na + nb;
if nsum > MAX_HYP_EXP {
return None;
}
let mut m = nsum % 2;
while m <= nsum {
let target = ((m + nsum) / 2) as i64;
let mut sum: i64 = 0;
for i in 0..=(na as i64) {
let jj = target - i;
if jj < 0 || jj > nb as i64 {
continue;
}
let c = binom(na as i64, i) * binom(nb as i64, jj);
sum += if (nb as i64 - jj) % 2 == 0 { c } else { -c };
}
if sum != 0 {
let num = if m == 0 { sum } else { 2 * sum };
let r = rat_atom(ctx, num, 1i64 << nsum);
let mut fs = coeff.to_vec();
fs.push(r);
if j > 0 {
fs.push(int_pow(ctx, t, j));
}
if m > 0 {
let marg = if m == 1 {
t
} else {
ctx.mul(&[ctx.num(m as i64), t])
};
fs.push(ctx.fun(if nb.is_multiple_of(2) { "cosh" } else { "sinh" }, &[marg]));
}
out.push(normalize(ctx, ctx.mul(&fs)));
}
m += 2;
}
Some(())
}
fn hyp_multiple(arg: Atom<'_>, t_sym: Symbol) -> Option<i64> {
match arg.node() {
AtomNode::Var(v) if *v == t_sym => Some(1),
AtomNode::Mul(args) if args.len() == 2 => {
let (a, b) = (args[0], args[1]);
match (a.node(), b.node()) {
(AtomNode::Num(m), AtomNode::Var(v)) if *v == t_sym && *m > 0 => Some(*m),
(AtomNode::Var(v), AtomNode::Num(m)) if *v == t_sym && *m > 0 => Some(*m),
_ => None,
}
}
_ => None,
}
}
fn integrate_hyp_term<'a>(
ctx: &'a AtomArena<'a>,
term: Atom<'a>,
t_sym: Symbol,
) -> Option<Atom<'a>> {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![term],
};
let mut coeff: Vec<Atom<'a>> = Vec::new();
let mut j = 0i64;
let mut hyp: Option<(bool, i64)> = None;
for f in factors {
if is_constant(f, t_sym) {
coeff.push(f);
continue;
}
match f.node() {
AtomNode::Var(v) if *v == t_sym => j += 1,
AtomNode::Pow(b, e) if matches!(b.node(), AtomNode::Var(v) if *v == t_sym) => {
let AtomNode::Num(n) = e.node() else {
return None;
};
if *n < 0 {
return None;
}
j += n;
}
AtomNode::Fun(name, args) if args.len() == 1 => {
let is_sinh = match name.as_str() {
"sinh" => true,
"cosh" => false,
_ => return None,
};
if hyp.is_some() {
return None;
}
hyp = Some((is_sinh, hyp_multiple(args[0], t_sym)?));
}
_ => return None,
}
}
if j > MAX_SUBST_K {
return None;
}
let t = ctx.var(t_sym.as_str());
let (is_sinh, m) = match hyp {
None => {
let mut fs = coeff;
fs.push(int_pow(ctx, t, j + 1));
fs.push(inv(ctx, ctx.num(j + 1)));
return Some(normalize(ctx, ctx.mul(&fs)));
}
Some(h) => h,
};
let mut terms = Vec::new();
let mut fall: i64 = 1; for i in 0..=j {
if i > 0 {
fall = fall.checked_mul(j - i + 1)?;
}
let mpow = m.checked_pow((i + 1) as u32)?;
let out_sinh = if is_sinh { i % 2 == 1 } else { i % 2 == 0 };
let marg = if m == 1 { t } else { ctx.mul(&[ctx.num(m), t]) };
let mut fs = coeff.clone();
if i % 2 == 1 {
fs.push(ctx.num(-1));
}
if fall != 1 {
fs.push(ctx.num(fall));
}
let deg = j - i;
if deg > 0 {
fs.push(int_pow(ctx, t, deg));
}
fs.push(ctx.pow(ctx.num(mpow), ctx.num(-1)));
fs.push(ctx.fun(if out_sinh { "sinh" } else { "cosh" }, &[marg]));
terms.push(normalize(ctx, ctx.mul(&fs)));
}
Some(normalize(ctx, ctx.add(&terms)))
}
pub(crate) fn integrate_bare_linear<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if is_constant(expr, var) || matches!(expr.node(), AtomNode::Add(_)) {
return None;
}
let factors: Vec<Atom<'a>> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![expr],
};
let mut rest: Vec<Atom<'a>> = Vec::new();
let mut bare: Option<(InvFun, Atom<'a>)> = None;
for f in factors {
if is_constant(f, var) {
rest.push(f);
continue;
}
if bare.is_none()
&& let AtomNode::Fun(name, args) = f.node()
&& args.len() == 1
&& inv_fun(name.as_str()).is_some()
{
bare = Some((inv_fun(name.as_str())?, args[0]));
continue;
}
return None;
}
let (f, u) = bare?;
let (slope, _intercept) = linear_form(ctx, u, var)?;
if is_zero_atom(ctx, slope) {
return None;
}
let u2 = int_pow(ctx, u, 2);
let one_minus_u2 = || ctx.add(&[ctx.num(1), ctx.mul(&[ctx.num(-1), u2])]);
let fu = ctx.fun(f.name(), &[u]);
let core = match f {
InvFun::Asin => ctx.add(&[ctx.mul(&[u, fu]), ctx.fun("sqrt", &[one_minus_u2()])]),
InvFun::Acos => ctx.add(&[
ctx.mul(&[u, fu]),
ctx.mul(&[ctx.num(-1), ctx.fun("sqrt", &[one_minus_u2()])]),
]),
InvFun::Atan => ctx.add(&[
ctx.mul(&[u, fu]),
ctx.mul(&[
rat_atom(ctx, -1, 2),
ctx.fun("log", &[ctx.add(&[ctx.num(1), u2])]),
]),
]),
InvFun::Asinh => ctx.add(&[
ctx.mul(&[u, fu]),
ctx.mul(&[ctx.num(-1), ctx.fun("sqrt", &[ctx.add(&[ctx.num(1), u2])])]),
]),
InvFun::Acosh => ctx.add(&[
ctx.mul(&[u, fu]),
ctx.mul(&[ctx.num(-1), ctx.fun("sqrt", &[ctx.add(&[ctx.num(-1), u2])])]),
]),
InvFun::Atanh => ctx.add(&[
ctx.mul(&[u, fu]),
ctx.mul(&[rat_atom(ctx, 1, 2), ctx.fun("log", &[one_minus_u2()])]),
]),
};
let mut out = rest;
out.push(inv(ctx, slope));
out.push(core);
Some(normalize(ctx, ctx.mul(&out)))
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum CancelGuard {
Unconditional,
AtanBand,
AcoshHalf,
}
fn cancel_guard(outer: &str, inner: &str) -> Option<CancelGuard> {
Some(match (outer, inner) {
("atanh", "tanh") | ("asinh", "sinh") | ("acoth", "coth") | ("log", "exp") => {
CancelGuard::Unconditional
}
("atan", "tan") => CancelGuard::AtanBand,
("acosh", "cosh") => CancelGuard::AcoshHalf,
_ => return None,
})
}
fn guard_certified(g: CancelGuard, u: Atom<'_>, var: Symbol) -> bool {
match g {
CancelGuard::Unconditional => true,
CancelGuard::AtanBand => {
is_constant(u, var)
&& rat_of(u).is_some_and(|(p, q)| {
q > 0 && p.saturating_abs().saturating_mul(2) <= q.saturating_mul(3)
})
}
CancelGuard::AcoshHalf => {
is_constant(u, var) && rat_of(u).is_some_and(|(p, q)| q > 0 && p >= 0)
}
}
}
fn rewrite_children<'a>(
ctx: &'a AtomArena<'a>,
args: &[Atom<'a>],
var: Symbol,
budget: &mut usize,
) -> (bool, Vec<Atom<'a>>) {
let mut changed = false;
let mut out: Vec<Atom<'a>> = Vec::with_capacity(args.len());
for a in args {
match cancel_compositions(ctx, *a, var, budget) {
Some(r) => {
changed = true;
out.push(r);
}
None => out.push(*a),
}
}
(changed, out)
}
fn cancel_compositions<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
budget: &mut usize,
) -> Option<Atom<'a>> {
if *budget == 0 {
return None;
}
*budget -= 1;
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => None,
AtomNode::Add(args) => {
let (changed, out) = rewrite_children(ctx, args, var, budget);
if changed { Some(ctx.add(&out)) } else { None }
}
AtomNode::Mul(args) => {
let (changed, out) = rewrite_children(ctx, args, var, budget);
if changed { Some(ctx.mul(&out)) } else { None }
}
AtomNode::Pow(b, e) => {
let nb = cancel_compositions(ctx, *b, var, budget);
let ne = cancel_compositions(ctx, *e, var, budget);
match (nb, ne) {
(None, None) => None,
(nb, ne) => Some(ctx.pow(nb.unwrap_or(*b), ne.unwrap_or(*e))),
}
}
AtomNode::Fun(name, args) => {
let (changed, new_args) = rewrite_children(ctx, args, var, budget);
if new_args.len() == 1
&& let AtomNode::Fun(inner, iargs) = new_args[0].node()
&& iargs.len() == 1
&& let Some(g) = cancel_guard(name.as_str(), inner.as_str())
&& guard_certified(g, iargs[0], var)
{
return Some(iargs[0]);
}
if changed {
Some(ctx.fun(name.as_str(), &new_args))
} else {
None
}
}
}
}
fn factor_var_power(factor: Atom<'_>, var: Symbol) -> Option<(i64, i64)> {
match factor.node() {
AtomNode::Var(v) if *v == var => Some((1, 1)),
AtomNode::Pow(b, e) => {
let (p, q) = rat_of(*e)?;
if q <= 0 {
return None;
}
let (bp, bq) = factor_var_power(*b, var)?;
Some((bp.checked_mul(p)?, bq.checked_mul(q)?))
}
AtomNode::Fun(name, args) if name.as_str() == "sqrt" && args.len() == 1 => {
let (p, q) = factor_var_power(args[0], var)?;
Some((p, q.checked_mul(2)?))
}
_ => None,
}
}
fn add_var_power(acc: &mut Option<(i64, i64)>, add: (i64, i64)) -> Option<()> {
match acc {
None => {
*acc = Some(add);
Some(())
}
Some((p, q)) => {
let nq = q.checked_mul(add.1)?;
if nq <= 0 || nq > MAX_RAT_Q * MAX_RAT_Q {
return None;
}
let np = p.checked_mul(add.1)?.checked_add(add.0.checked_mul(*q)?)?;
*acc = Some((np, nq));
Some(())
}
}
}
fn split_monomial<'a>(term: Atom<'a>, var: Symbol) -> Option<(Vec<Atom<'a>>, i64, i64)> {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![term],
};
let mut coeff: Vec<Atom<'a>> = Vec::new();
let mut power: Option<(i64, i64)> = None;
for f in factors {
if let Some(pq) = factor_var_power(f, var) {
add_var_power(&mut power, pq)?;
continue;
}
if is_constant(f, var) {
coeff.push(f);
continue;
}
return None;
}
let (p, q) = power?;
if q <= 0 {
return None;
}
Some((coeff, p, q))
}
fn integrate_expanded_terms<'a>(
ctx: &'a AtomArena<'a>,
folded: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let terms: Vec<Atom<'a>> = match folded.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![folded],
};
let x = ctx.var(var.as_str());
let mut out: Vec<Atom<'a>> = Vec::with_capacity(terms.len());
for term in terms {
if is_constant(term, var) {
out.push(ctx.mul(&[term, x]));
continue;
}
if let Some((coeff, p, q)) = split_monomial(term, var) {
let shifted = p.checked_add(q)?;
let mut fs = coeff;
if shifted == 0 {
fs.push(ctx.fun("log", &[x]));
} else {
let r = rat_atom(ctx, shifted, q);
fs.push(ctx.pow(x, r));
fs.push(inv(ctx, r));
}
out.push(normalize(ctx, ctx.mul(&fs)));
continue;
}
let r = integrate_raw(ctx, term, var, 0, true, 0, 0);
if contains_integral(r) {
return None;
}
out.push(r);
}
Some(normalize(ctx, ctx.add(&out)))
}
fn affine_kernel_pair(outer: &str, inner: &str) -> bool {
matches!(
(outer, inner),
("atanh", "tanh")
| ("asinh", "sinh")
| ("acoth", "coth")
| ("acoth", "tanh")
| ("atan", "tan")
| ("log", "exp")
)
}
struct AffineKernel<'a> {
base: Atom<'a>,
sigma: Atom<'a>,
}
fn match_affine_kernel<'a>(
ctx: &'a AtomArena<'a>,
factor: Atom<'a>,
var: Symbol,
) -> Option<AffineKernel<'a>> {
let AtomNode::Fun(name, args) = factor.node() else {
return None;
};
if args.len() != 1 {
return None;
}
let AtomNode::Fun(inner, iargs) = args[0].node() else {
return None;
};
if iargs.len() != 1 || !affine_kernel_pair(name.as_str(), inner.as_str()) {
return None;
}
let (sigma, _rho) = linear_form(ctx, iargs[0], var)?;
if is_zero_atom(ctx, sigma) {
return None;
}
Some(AffineKernel {
base: factor,
sigma,
})
}
fn reciprocal_power(factor: Atom<'_>, var: Symbol) -> Option<i64> {
let AtomNode::Pow(b, e) = factor.node() else {
return None;
};
if !matches!(b.node(), AtomNode::Var(v) if *v == var) {
return None;
}
let AtomNode::Num(j) = e.node() else {
return None;
};
if *j <= -1 && *j >= -MAX_RECIP_EXP {
Some(-*j)
} else {
None
}
}
fn pick_proxy_symbol(expr: Atom<'_>, var: Symbol) -> Option<Symbol> {
for name in ["w", "k", "r", "z"] {
let s = Symbol::new(name);
if s != var && !contains_symbol(expr, s) {
return Some(s);
}
}
None
}
pub(crate) fn integrate_affine_kernel<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if is_constant(expr, var) || matches!(expr.node(), AtomNode::Add(_)) {
return None;
}
let factors: Vec<Atom<'a>> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![expr],
};
let mut rest: Vec<Atom<'a>> = Vec::new();
let mut kern: Option<(AffineKernel<'a>, i64)> = None;
let mut cofactor: Vec<Atom<'a>> = Vec::new();
let mut cofactor_deg = 0usize;
let mut recip = 0i64;
for f in factors {
if is_constant(f, var) {
rest.push(f);
continue;
}
if kern.is_none() {
if let Some(k) = match_affine_kernel(ctx, f, var) {
kern = Some((k, 1));
continue;
}
if let AtomNode::Pow(b, e) = f.node()
&& let AtomNode::Num(m) = e.node()
&& *m != 0
&& (-MAX_KERNEL_POW..=MAX_KERNEL_POW).contains(m)
&& let Some(k) = match_affine_kernel(ctx, *b, var)
{
kern = Some((k, *m));
continue;
}
}
if recip == 0
&& let Some(j) = reciprocal_power(f, var)
{
recip = j;
cofactor.push(f);
continue;
}
if let Some(d) = poly_deg(f, var)
&& d >= 1
&& cofactor_deg + d <= MAX_COFACTOR_DEG
{
cofactor_deg += d;
cofactor.push(f);
continue;
}
return None;
}
let (k, m) = kern?;
let w_sym = pick_proxy_symbol(expr, var)?;
let x = ctx.var(var.as_str());
let w = ctx.var(w_sym.as_str());
let lin = normalize(ctx, ctx.add(&[ctx.mul(&[k.sigma, x]), w]));
let mut rewritten: Vec<Atom<'a>> = rest;
rewritten.push(int_pow(ctx, lin, m));
rewritten.extend(cofactor);
let product = normalize(ctx, ctx.mul(&rewritten));
if node_count(product) > MAX_KERNEL_NODES {
return None;
}
let expanded = crate::expand::expand_bounded(ctx, product).unwrap_or(product);
let folded = crate::ode::util::collect_terms(ctx, normalize(ctx, expanded));
if node_count(folded) > MAX_KERNEL_NODES {
return None;
}
let chain = integrate_expanded_terms(ctx, folded, var)?;
if contains_integral(chain) {
return None;
}
let back = normalize(ctx, ctx.add(&[k.base, ctx.mul(&[ctx.num(-1), k.sigma, x])]));
let res = normalize(ctx, replace_symbol(ctx, chain, w_sym, back));
if contains_symbol(res, w_sym) {
return None;
}
Some(res)
}
fn family_kernel_of_u<'a>(ctx: &'a AtomArena<'a>, f: InvFun, u: Atom<'a>) -> Atom<'a> {
let u2 = int_pow(ctx, u, 2);
match f {
InvFun::Asin | InvFun::Acos | InvFun::Atanh => {
normalize(ctx, ctx.add(&[ctx.num(1), ctx.mul(&[ctx.num(-1), u2])]))
}
InvFun::Atan | InvFun::Asinh => normalize(ctx, ctx.add(&[ctx.num(1), u2])),
InvFun::Acosh => normalize(ctx, ctx.add(&[ctx.num(-1), u2])),
}
}
fn match_inv_power_rational<'a>(
ctx: &'a AtomArena<'a>,
factor: Atom<'a>,
var: Symbol,
) -> Option<(InvFun, Atom<'a>, Atom<'a>, Atom<'a>, i64, i64)> {
if let AtomNode::Fun(name, args) = factor.node()
&& name.as_str() == "sqrt"
&& args.len() == 1
&& let Some((f, arg, a, b)) = match_inv_base(ctx, args[0], var)
{
return Some((f, arg, a, b, 1, 2));
}
if let AtomNode::Pow(b, e) = factor.node()
&& let Some((p, q)) = rat_of(*e)
&& q > 0
&& q <= MAX_RAT_Q
&& p >= 1
&& p <= MAX_POWER_K * q
&& let Some((f, arg, a, bb)) = match_inv_base(ctx, *b, var)
{
return Some((f, arg, a, bb, p, q));
}
let (f, arg, a, b) = match_inv_base(ctx, factor, var)?;
Some((f, arg, a, b, 1, 1))
}
fn flatten_powers<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, budget: &mut usize) -> Atom<'a> {
if *budget == 0 {
return expr;
}
*budget -= 1;
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => expr,
AtomNode::Add(args) => {
let out: Vec<Atom<'a>> = args
.iter()
.map(|a| flatten_powers(ctx, *a, budget))
.collect();
ctx.add(&out)
}
AtomNode::Mul(args) => {
let out: Vec<Atom<'a>> = args
.iter()
.map(|a| flatten_powers(ctx, *a, budget))
.collect();
ctx.mul(&out)
}
AtomNode::Pow(b, e) => {
let nb = flatten_powers(ctx, *b, budget);
let ne = flatten_powers(ctx, *e, budget);
if let AtomNode::Num(n) = ne.node()
&& *n >= 2
&& let AtomNode::Mul(fs) = nb.node()
{
let out: Vec<Atom<'a>> = fs
.iter()
.map(|f| ctx.pow(flatten_powers(ctx, *f, budget), ne))
.collect();
return ctx.mul(&out);
}
ctx.pow(nb, ne)
}
AtomNode::Fun(name, args) => {
let out: Vec<Atom<'a>> = args
.iter()
.map(|a| flatten_powers(ctx, *a, budget))
.collect();
ctx.fun(name.as_str(), &out)
}
}
}
fn companion_normal_form<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>) -> Atom<'a> {
let expanded = crate::expand::expand_bounded(ctx, expr).unwrap_or(expr);
let mut budget = MAX_CANCEL_BUDGET;
let flat = flatten_powers(ctx, expanded, &mut budget);
crate::ode::util::collect_terms(ctx, normalize(ctx, flat))
}
fn radicand_scale<'a>(
ctx: &'a AtomArena<'a>,
r: Atom<'a>,
e_u: Atom<'a>,
x_root: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let re = companion_normal_form(ctx, r);
let ee = companion_normal_form(ctx, e_u);
for point in [x_root, ctx.num(0)] {
let rp = companion_normal_form(ctx, replace_symbol(ctx, re, var, point));
let ep = companion_normal_form(ctx, replace_symbol(ctx, ee, var, point));
if !is_constant(rp, var) || is_zero_atom(ctx, ep) {
continue;
}
let s = normalize(ctx, ctx.mul(&[rp, inv(ctx, ep)]));
if is_zero_atom(ctx, s) {
continue;
}
let diff = ctx.add(&[re, ctx.mul(&[ctx.num(-1), s, ee])]);
let folded = companion_normal_form(ctx, diff);
if folds_to_zero(ctx, folded) {
return Some(s);
}
}
None
}
pub(crate) fn integrate_companion_form<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if is_constant(expr, var) || matches!(expr.node(), AtomNode::Add(_)) {
return None;
}
let factors: Vec<Atom<'a>> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![expr],
};
let mut rest: Vec<Atom<'a>> = Vec::new();
let mut inv_f: Option<(InvFun, Atom<'a>, Atom<'a>, Atom<'a>, i64, i64)> = None;
let mut kern: Option<(Atom<'a>, KernelExp)> = None;
for f in factors {
if is_constant(f, var) {
rest.push(f);
continue;
}
if inv_f.is_none()
&& let Some(m) = match_inv_power_rational(ctx, f, var)
{
inv_f = Some(m);
continue;
}
if kern.is_none()
&& let Some(kd) = match_kernel(f)
{
kern = Some(kd);
continue;
}
return None;
}
let (f, arg, a, b, p, q) = inv_f?;
let (r, ke) = kern?;
if ke != family_kernel_exp(f) {
return None;
}
let (sigma, rho) = linear_form(ctx, arg, var)?;
if is_zero_atom(ctx, sigma) {
return None;
}
let e_u = family_kernel_of_u(ctx, f, arg);
let x_root = normalize(ctx, ctx.mul(&[ctx.num(-1), rho, inv(ctx, sigma)]));
let s_atom = radicand_scale(ctx, r, e_u, x_root, var)?;
if ke == KernelExp::InvSqrt && !sqrt_norm_positive(s_atom) {
return None;
}
let s_factor = match ke {
KernelExp::Recip => inv(ctx, s_atom),
KernelExp::InvSqrt => ctx.pow(s_atom, rat_atom(ctx, -1, 2)),
};
let base = normalize(ctx, ctx.add(&[a, ctx.mul(&[b, ctx.fun(f.name(), &[arg])])]));
let mut out = rest;
if family_sign(f) < 0 {
out.push(ctx.num(-1));
}
out.push(s_factor);
out.push(ctx.pow(base, rat_atom(ctx, p + q, q)));
let denom = normalize(ctx, ctx.mul(&[b, sigma, rat_atom(ctx, p + q, q)]));
out.push(inv(ctx, denom));
Some(normalize(ctx, ctx.mul(&out)))
}
pub(crate) fn integrate_inverse_trig<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if is_constant(expr, var) || node_count(expr) > MAX_NODES {
return None;
}
let mut budget = MAX_CANCEL_BUDGET;
if let Some(rewritten) = cancel_compositions(ctx, expr, var, &mut budget) {
let rewritten = normalize(ctx, rewritten);
if node_count(rewritten) <= MAX_NODES {
let expanded = crate::expand::expand_bounded(ctx, rewritten).unwrap_or(rewritten);
let folded = crate::ode::util::collect_terms(ctx, normalize(ctx, expanded));
if node_count(folded) <= MAX_KERNEL_NODES
&& let Some(r) = integrate_expanded_terms(ctx, folded, var)
{
return Some(normalize(ctx, r));
}
}
}
integrate_bare_linear(ctx, expr, var)
.or_else(|| integrate_kernel_power(ctx, expr, var))
.or_else(|| integrate_companion_form(ctx, expr, var))
.or_else(|| integrate_affine_kernel(ctx, expr, var))
.or_else(|| integrate_invhyp_subst(ctx, expr, var))
}
#[cfg(test)]
mod tests {
use super::*;
use ocas_core::arena::Arena;
fn parse_norm<'a>(ctx: &'a AtomArena<'a>, s: &str) -> Atom<'a> {
normalize(ctx, ocas_parse::parse(ctx, s).unwrap())
}
fn eval_f64(expr: Atom<'_>, env: &[(Symbol, f64)]) -> Option<f64> {
match expr.node() {
AtomNode::Num(n) => Some(*n as f64),
AtomNode::Var(v) => env.iter().find(|(s, _)| s == v).map(|(_, val)| *val),
AtomNode::Add(args) => args
.iter()
.try_fold(0.0, |acc, a| Some(acc + eval_f64(*a, env)?)),
AtomNode::Mul(args) => args
.iter()
.try_fold(1.0, |acc, a| Some(acc * eval_f64(*a, env)?)),
AtomNode::Pow(b, e) => Some(eval_f64(*b, env)?.powf(eval_f64(*e, env)?)),
AtomNode::Fun(name, args) => {
let v = eval_f64(*args.first()?, env)?;
Some(match name.as_str() {
"sin" => v.sin(),
"cos" => v.cos(),
"tan" => v.tan(),
"cot" => 1.0 / v.tan(),
"sec" => 1.0 / v.cos(),
"csc" => 1.0 / v.sin(),
"exp" => v.exp(),
"log" => v.ln(),
"sqrt" => v.sqrt(),
"sinh" => v.sinh(),
"cosh" => v.cosh(),
"tanh" => v.tanh(),
"coth" => 1.0 / v.tanh(),
"sech" => 1.0 / v.cosh(),
"csch" => 1.0 / v.sinh(),
"asin" => v.asin(),
"acos" => v.acos(),
"atan" => v.atan(),
"acot" => std::f64::consts::FRAC_PI_2 - v.atan(),
"asinh" => v.asinh(),
"acosh" => v.acosh(),
"atanh" => v.atanh(),
"acoth" => (1.0 / v).atanh(),
_ => return None,
})
}
}
}
fn assert_antiderivative_num<'a>(
ctx: &'a AtomArena<'a>,
integrand: Atom<'a>,
var: Symbol,
consts: &[(Symbol, f64)],
samples: &[f64],
) {
let result = integrate_inverse_trig(ctx, integrand, var).expect("mechanism declined");
assert!(
!result.to_string().contains("Integral"),
"residue: {result}"
);
let d = crate::diff(ctx, result, var);
for &xv in samples {
let mut env = consts.to_vec();
env.push((var, xv));
let lhs = eval_f64(d, &env).expect("eval diff");
let rhs = eval_f64(integrand, &env).expect("eval integrand");
let tol = 1e-6 * rhs.abs().max(1.0);
assert!(
(lhs - rhs).abs() < tol,
"at x={xv}: diff={lhs} integrand={rhs} (result: {result})"
);
}
}
fn assert_declined<'a>(ctx: &'a AtomArena<'a>, input: &str) {
let expr = parse_norm(ctx, input);
assert!(
integrate_inverse_trig(ctx, expr, Symbol::new("x")).is_none(),
"expected None for {input}"
);
}
fn replace_atom<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
target: Atom<'a>,
repl: Atom<'a>,
) -> Atom<'a> {
if expr == target {
return repl;
}
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => expr,
AtomNode::Add(args) => {
let out: Vec<Atom<'a>> = args
.iter()
.map(|a| replace_atom(ctx, *a, target, repl))
.collect();
ctx.add(&out)
}
AtomNode::Mul(args) => {
let out: Vec<Atom<'a>> = args
.iter()
.map(|a| replace_atom(ctx, *a, target, repl))
.collect();
ctx.mul(&out)
}
AtomNode::Pow(b, e) => {
let nb = replace_atom(ctx, *b, target, repl);
let ne = replace_atom(ctx, *e, target, repl);
ctx.pow(nb, ne)
}
AtomNode::Fun(name, args) => {
let out: Vec<Atom<'a>> = args
.iter()
.map(|a| replace_atom(ctx, *a, target, repl))
.collect();
ctx.fun(name.as_str(), &out)
}
}
}
fn assert_antiderivative_branch_shifted<'a>(
ctx: &'a AtomArena<'a>,
integrand: Atom<'a>,
kernel: Atom<'a>,
proxy: Atom<'a>,
var: Symbol,
consts: &[(Symbol, f64)],
samples: &[f64],
) {
let result = integrate_inverse_trig(ctx, integrand, var).expect("mechanism declined");
assert!(
!result.to_string().contains("Integral"),
"residue: {result}"
);
for leaked in ["w", "k", "r", "z"] {
assert!(
!contains_symbol(result, Symbol::new(leaked)),
"proxy symbol {leaked} leaked: {result}"
);
}
let lhs_expr = replace_atom(ctx, result, kernel, proxy);
let rhs_expr = replace_atom(ctx, integrand, kernel, proxy);
let d = crate::diff(ctx, lhs_expr, var);
for &xv in samples {
let mut env = consts.to_vec();
env.push((var, xv));
let lhs = eval_f64(d, &env).expect("eval diff");
let rhs = eval_f64(rhs_expr, &env).expect("eval integrand");
let tol = 1e-6 * rhs.abs().max(1.0);
assert!(
(lhs - rhs).abs() < tol,
"at x={xv}: diff={lhs} integrand={rhs} (result: {result})"
);
}
}
#[test]
fn m1_asinh_kernel_sqrt_fun_form() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*asinh(c*x))^2/sqrt(1 + c^2*x^2)");
let env = [
(Symbol::new("a"), 1.2),
(Symbol::new("b"), -0.7),
(Symbol::new("c"), 1.5),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.2, 0.5, 0.9]);
}
#[test]
fn m1_asin_kernel_pow_form() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*asin(c*x))^3*(1 - c^2*x^2)^(-1/2)");
let env = [
(Symbol::new("a"), 0.8),
(Symbol::new("b"), 1.1),
(Symbol::new("c"), 2.0),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.1, 0.3, 0.45]);
}
#[test]
fn m1_acos_kernel_negative_sign() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*acos(c*x))/sqrt(1 - c^2*x^2)");
let env = [
(Symbol::new("a"), 0.9),
(Symbol::new("b"), -1.2),
(Symbol::new("c"), 1.6),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.1, 0.3, 0.55]);
}
#[test]
fn m1_atanh_kernel_reciprocal() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*atanh(c*x))^2/(1 - c^2*x^2)");
let env = [
(Symbol::new("a"), 1.3),
(Symbol::new("b"), 0.7),
(Symbol::new("c"), 1.9),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.1, 0.25, 0.5]);
}
#[test]
fn m1_atan_kernel_symbolic_norm() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*atan(c*x))^4/(d + c^2*d*x^2)");
let env = [
(Symbol::new("a"), 1.1),
(Symbol::new("b"), -0.8),
(Symbol::new("c"), 1.7),
(Symbol::new("d"), 2.3),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.2, 0.7, 1.3]);
}
#[test]
fn m1_acosh_kernel() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*acosh(c*x))^2/sqrt(c^2*x^2 - 1)");
let env = [
(Symbol::new("a"), 0.6),
(Symbol::new("b"), 1.4),
(Symbol::new("c"), 0.8),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[1.5, 2.0, 3.0]);
}
#[test]
fn m1_numeric_norm_factor() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*asinh(c*x))^2/sqrt(4 + 4*c^2*x^2)");
let env = [
(Symbol::new("a"), 1.2),
(Symbol::new("b"), -0.7),
(Symbol::new("c"), 1.5),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.2, 0.5, 0.9]);
}
#[test]
fn m2_corpus_acosh_poly() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(d + e*x)*(a + b*acosh(c*x))^2");
let env = [
(Symbol::new("a"), 0.9),
(Symbol::new("b"), -1.3),
(Symbol::new("c"), 0.7),
(Symbol::new("d"), 1.1),
(Symbol::new("e"), 0.6),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[1.6, 2.2, 3.0]);
}
#[test]
fn m2_corpus_asinh_sqrt_pow() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(d + c^2*d*x^2)^(3/2)*(a + b*asinh(c*x))");
let env = [
(Symbol::new("a"), 0.8),
(Symbol::new("b"), 1.1),
(Symbol::new("c"), 1.3),
(Symbol::new("d"), 1.7),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.2, 0.6, 1.1]);
}
#[test]
fn m2_corpus_acosh_linear_arg_k4() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(c*e + d*e*x)*(a + b*acosh(c + d*x))^4");
let env = [
(Symbol::new("a"), 0.7),
(Symbol::new("b"), -1.1),
(Symbol::new("c"), 0.4),
(Symbol::new("d"), 0.9),
(Symbol::new("e"), 1.4),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.8, 1.2, 1.8]);
}
#[test]
fn m2_asinh_linear_arg_bare_power() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = parse_norm(&ctx, "(a + b*asinh(c + d*x))^2");
let env = [
(Symbol::new("a"), 1.1),
(Symbol::new("b"), 0.7),
(Symbol::new("c"), 0.4),
(Symbol::new("d"), 1.2),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.3, 0.8, 1.5]);
}
#[test]
fn m3_trig_family_linear_arg() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 0.3), (Symbol::new("b"), 1.2)];
for input in ["asin(a + b*x)", "acos(a + b*x)", "atan(a + b*x)"] {
let expr = parse_norm(&ctx, input);
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.1, 0.3, 0.5]);
}
}
#[test]
fn m3_hyper_family_linear_arg() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 0.3), (Symbol::new("b"), 1.2)];
for input in ["asinh(a + b*x)", "atanh(a + b*x)"] {
let expr = parse_norm(&ctx, input);
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.1, 0.3, 0.5]);
}
}
#[test]
fn m3_acosh_linear_arg_domain() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 1.4), (Symbol::new("b"), 0.9)];
let expr = parse_norm(&ctx, "acosh(a + b*x)");
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.1, 0.4, 0.8]);
}
#[test]
fn m3_constant_multiple() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 0.3), (Symbol::new("b"), 1.2)];
let expr = parse_norm(&ctx, "q*asin(a + b*x)");
let mut env2 = env.to_vec();
env2.push((Symbol::new("q"), 2.5));
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env2, &[0.1, 0.3, 0.5]);
}
#[test]
fn declines_out_of_scope() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
assert_declined(&ctx, "x/acosh(a*x)^4");
assert_declined(&ctx, "(a + b*atanh(c*x))^2/(d + c*d*x)");
assert_declined(&ctx, "x^2/((c + a^2*c*x^2)^(5/2)*atan(a*x)^2)");
assert_declined(&ctx, "a + b*asin(1 + d*x^2)");
assert_declined(&ctx, "asin(1 + d*x^2)");
assert_declined(&ctx, "(d + e*x)*(a + b*acosh(c*x))^6");
assert_declined(&ctx, "(-4 - 4*c^2*x^2)^(3/2)*(a + b*asinh(c*x))");
assert_declined(&ctx, "(a + b*asinh(c*x))^2/sqrt(1 - c^2*x^2)");
assert_declined(&ctx, "(1 + c*x)^3*(a + b*atanh(c*x))^3");
}
#[test]
fn p1_atanh_tanh_numeric_and_symbolic() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let expr = parse_norm(&ctx, "atanh(tanh(b*x))^2");
let env = [(Symbol::new("b"), 1.2)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.1, 0.4, 0.9]);
let expr = parse_norm(&ctx, "atanh(tanh(a + b*x))^2");
let env = [(Symbol::new("a"), 0.4), (Symbol::new("b"), 1.1)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.1, 0.3, 0.5]);
}
#[test]
fn p1_corpus_atanh_tanh_shapes() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let env = [(Symbol::new("a"), 0.4), (Symbol::new("b"), 0.9)];
for input in ["atanh(tanh(a + b*x))^3", "atanh(tanh(a + b*x))^2"] {
let expr = parse_norm(&ctx, input);
assert_antiderivative_num(&ctx, expr, x, &env, &[0.1, 0.3, 0.5]);
}
let expr = parse_norm(&ctx, "atanh(tanh(a + b*x))^4/x^4");
assert_antiderivative_num(&ctx, expr, x, &env, &[0.3, 0.6, 1.1]);
let expr = parse_norm(&ctx, "atanh(tanh(a + b*x))^2*sqrt(x)");
assert_antiderivative_num(&ctx, expr, x, &env, &[0.2, 0.5, 0.9]);
}
#[test]
fn p1_acoth_coth_numeric_and_symbolic() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let expr = parse_norm(&ctx, "acoth(coth(b*x))^2");
let env = [(Symbol::new("b"), 1.1)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.3, 0.7, 1.1]);
let expr = parse_norm(&ctx, "acoth(coth(a + b*x))^3");
let env = [(Symbol::new("a"), 0.5), (Symbol::new("b"), 1.1)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.3, 0.6, 0.9]);
}
#[test]
fn p1_asinh_sinh_numeric_and_symbolic() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let expr = parse_norm(&ctx, "asinh(sinh(b*x))^3");
let env = [(Symbol::new("b"), 1.3)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.2, 0.6, 1.0]);
let expr = parse_norm(&ctx, "asinh(sinh(a + b*x))^2/x^2");
let env = [(Symbol::new("a"), -0.3), (Symbol::new("b"), 0.8)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.4, 0.9, 1.4]);
}
#[test]
fn p1_log_exp_numeric_and_symbolic() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let expr = parse_norm(&ctx, "log(exp(b*x))^2");
let env = [(Symbol::new("b"), 1.4)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.2, 0.5, 0.9]);
let expr = parse_norm(&ctx, "log(exp(a + b*x))^3/x");
let env = [(Symbol::new("a"), 0.6), (Symbol::new("b"), -1.1)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.3, 0.8, 1.5]);
}
#[test]
fn p1_band_guards_certify_only_inside_the_band() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
for input in ["atan(tan(1/2))", "acosh(cosh(1/2))"] {
let e = parse_norm(&ctx, input);
let mut budget = MAX_CANCEL_BUDGET;
let r = cancel_compositions(&ctx, e, x, &mut budget).expect("certified band");
assert_eq!(r, parse_norm(&ctx, "1/2"), "{input}");
}
for input in ["atan(tan(2))", "atan(tan(a + b*x))", "acosh(cosh(a + b*x))"] {
let e = parse_norm(&ctx, input);
let mut budget = MAX_CANCEL_BUDGET;
assert!(
cancel_compositions(&ctx, e, x, &mut budget).is_none(),
"expected no cancellation for {input}"
);
}
}
#[test]
fn p1_nested_pairs_reach_the_fixpoint_in_one_pass() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let e = parse_norm(&ctx, "atanh(tanh(asinh(sinh(x))))");
let mut budget = MAX_CANCEL_BUDGET;
let r = cancel_compositions(&ctx, e, x, &mut budget).expect("nested pair");
assert_eq!(r, ctx.var("x"));
let mut budget = MAX_CANCEL_BUDGET;
assert!(cancel_compositions(&ctx, r, x, &mut budget).is_none());
}
#[test]
fn p2_acoth_tanh_corpus_shapes() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let kernel = parse_norm(&ctx, "acoth(tanh(a + b*x))");
let proxy = parse_norm(&ctx, "a + b*x");
let env = [(Symbol::new("a"), 0.4), (Symbol::new("b"), 1.1)];
for input in ["acoth(tanh(a + b*x))^3/x^3", "x*acoth(tanh(a + b*x))^3"] {
let expr = parse_norm(&ctx, input);
assert_antiderivative_branch_shifted(
&ctx,
expr,
kernel,
proxy,
x,
&env,
&[0.4, 0.8, 1.3],
);
}
let expr = parse_norm(&ctx, "acoth(tanh(a + b*x))^3/x^6");
assert_antiderivative_branch_shifted(&ctx, expr, kernel, proxy, x, &env, &[0.4, 0.8, 1.3]);
let expr = parse_norm(&ctx, "acoth(tanh(2*x))^2/x^2");
let kernel = parse_norm(&ctx, "acoth(tanh(2*x))");
let proxy = parse_norm(&ctx, "2*x");
assert_antiderivative_branch_shifted(&ctx, expr, kernel, proxy, x, &[], &[0.5, 1.0, 1.6]);
}
#[test]
fn p2_atan_tan_is_branch_independent() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let kernel = parse_norm(&ctx, "atan(tan(a + b*x))");
let proxy = parse_norm(&ctx, "a + b*x");
let env = [(Symbol::new("a"), 0.7), (Symbol::new("b"), 1.3)];
for input in ["x^2*atan(tan(a + b*x))^2", "atan(tan(a + b*x))"] {
let expr = parse_norm(&ctx, input);
assert_antiderivative_branch_shifted(
&ctx,
expr,
kernel,
proxy,
x,
&env,
&[1.0, 2.0, 3.0],
);
}
}
#[test]
fn p3_companion_radical_with_intercept() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let expr = parse_norm(&ctx, "asin(a + b*x)/sqrt(1 - (a + b*x)^2)");
let env = [(Symbol::new("a"), 0.2), (Symbol::new("b"), 0.7)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.1, 0.4, 0.8]);
let expr = parse_norm(&ctx, "asin(c*x)/sqrt(1 - c^2*x^2)");
let env = [(Symbol::new("c"), 0.8)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.1, 0.5, 1.0]);
}
#[test]
fn p3_rational_inverse_power() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let expr = parse_norm(&ctx, "sqrt(asin(a*x))/(c - a^2*c*x^2)^(1/2)");
let env = [(Symbol::new("a"), 0.6), (Symbol::new("c"), 1.4)];
assert_antiderivative_num(&ctx, expr, x, &env, &[0.2, 0.5, 0.9]);
let expr = parse_norm(&ctx, "sqrt(acosh(x))/(x^2 - 1)^(1/2)");
assert_antiderivative_num(&ctx, expr, x, &[], &[1.4, 1.9, 2.4]);
}
#[test]
fn p3_declines_non_companion_radicals() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
assert_declined(&ctx, "atan(a + b*x)/sqrt(1 + a^2 + 2*a*b*x + b^2*x^2)");
assert_declined(&ctx, "atanh(a + b*x)/sqrt(1 - (a + b*x)^2)");
assert_declined(&ctx, "asinh(a + b*x)/sqrt(1 - (a + b*x)^2)");
}
#[test]
fn declines_out_of_family_compositions() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
assert_declined(&ctx, "atan(x^2)");
assert_declined(&ctx, "atanh(sin(x))");
assert_declined(&ctx, "sqrt(asin(x))");
assert_declined(&ctx, "asin(x)*exp(x)");
assert_declined(&ctx, "acosh(cosh(a + b*x))");
assert_declined(&ctx, "x^2*atanh(a + b*f^((c + d*x)))");
assert_declined(&ctx, "acoth(a + b*f^((c + d*x)))");
}
#[test]
fn declines_hang_attributed_shapes_quickly() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
assert_declined(&ctx, "(c + d*x)^(1/2)/(a + b*x)^2");
assert_declined(&ctx, "(e + f*x)^2*cos(c + d*x)/(a + b*sin(c + d*x))^3");
}
#[test]
fn stress_repeat_is_stable_and_budget_bounded() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = Symbol::new("x");
let shapes = [
"atanh(tanh(a + b*x))^2",
"atanh(tanh(a + b*x))^4/x^4",
"atanh(tanh(a + b*x))^2*sqrt(x)",
"x*acoth(tanh(a + b*x))^3",
"acoth(tanh(a + b*x))^3/x^3",
"atan(x^2)",
"atanh(sin(x))",
"(c + d*x)^(1/2)/(a + b*x)^2",
];
let solved = [true, true, true, true, true, false, false, false];
let mut first: Vec<Option<(String, usize)>> = Vec::with_capacity(shapes.len());
for round in 0..300 {
for (i, input) in shapes.iter().enumerate() {
let expr = parse_norm(&ctx, input);
let observed =
integrate_inverse_trig(&ctx, expr, x).map(|r| (r.to_string(), node_count(r)));
if round == 0 {
first.push(observed);
} else {
assert_eq!(observed, first[i], "{input} diverged at round {round}");
}
}
}
for (i, input) in shapes.iter().enumerate() {
if solved[i] {
let r = first[i].as_ref().expect("expected a closed form");
assert!(!r.0.contains("Integral"), "{input} left a residue: {}", r.0);
} else {
assert!(first[i].is_none(), "{input}: {:?}", first[i]);
}
}
}
}