use ocas_atom::normalize::normalize;
use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use ocas_domain::{Domain, Integer, IntegerDomain, Rational, RationalDomain};
use ocas_poly::{DenseUnivariatePolynomial, Lex, SparseMultivariatePolynomial};
use crate::tower::convert::{GeneratorField, atom_to_rational, rational_to_atom};
use super::rules::rat_of;
use super::{
binomial, contains_integral, gcd_i64, int_pow, inv, is_constant, lcm_i64, linear_form,
node_count, pick_subst_symbol, rat_atom, rational, replace_symbol, sqrt_quadratic,
};
type DPoly = DenseUnivariatePolynomial<RationalDomain>;
type Sparse = SparseMultivariatePolynomial<RationalDomain, Lex>;
const MAX_NODES: usize = 300;
const MAX_SITES: usize = 24;
const MAX_EXP_DEN: i64 = 6;
const MAX_REWRITE_ROUNDS: usize = 12;
const MAX_EXPAND_DEPTH: usize = 2;
const MAX_POW_GROUPS: usize = 32;
const MAX_POW_EXP: i64 = 64;
const MAX_RADICAL_DEG: i64 = 4;
const MAX_RADICAL_SITES: usize = 8;
const MAX_RAT_DEG: usize = 8;
const MAX_RAT_NUM_DEG: usize = 12;
const MAX_FACTORS: usize = 8;
const MAX_PIECES: usize = 16;
fn rat_exponent(atom: Atom<'_>) -> Option<(i64, i64)> {
let (p, q) = rat_of(atom)?;
if q == 0 {
return None;
}
let (p, q) = if q < 0 {
(p.checked_neg()?, q.checked_neg()?)
} else {
(p, q)
};
let g = gcd_i64(p, q).max(1);
Some((p / g, q / g))
}
pub(crate) fn integrate_kernel_subst<'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;
}
if matches!(expr.node(), AtomNode::Add(_)) {
return None;
}
let original = expr;
let rewritten = pythagorean_rewrite(ctx, expr, var);
let (expr, info) = match rewritten {
Some(r) => match match_family(ctx, r, var) {
Some(info) => (r, info),
None => (expr, match_family(ctx, expr, var)?),
},
None => (expr, match_family(ctx, expr, var)?),
};
for primary in info.candidates() {
if let Some(r) = try_primary(ctx, original, expr, var, &info, primary) {
return Some(r);
}
}
None
}
fn pythagorean_rewrite<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let mut current = normalize(ctx, expr);
let mut changed = false;
for _ in 0..MAX_REWRITE_ROUNDS {
let (next, hit) = rewrite_pass(ctx, current, var);
if !hit {
break;
}
changed = true;
current = next;
}
if changed { Some(current) } else { None }
}
fn rewrite_pass<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, var: Symbol) -> (Atom<'a>, bool) {
match expr.node() {
AtomNode::Add(args) => {
let mut kids = Vec::with_capacity(args.len());
let mut changed = false;
for a in args.iter() {
let (na, hit) = rewrite_pass(ctx, *a, var);
changed |= hit;
kids.push(na);
}
if let Some(r) = rewrite_add(ctx, &kids, var) {
return (r, true);
}
if !changed {
return (expr, false);
}
(ctx.add(&kids), true)
}
AtomNode::Mul(args) => {
let mut kids = Vec::with_capacity(args.len());
let mut changed = false;
for a in args.iter() {
let (na, hit) = rewrite_pass(ctx, *a, var);
changed |= hit;
kids.push(na);
}
if !changed {
return (expr, false);
}
(ctx.mul(&kids), true)
}
AtomNode::Pow(b, e) => {
let (nb, c1) = rewrite_pass(ctx, *b, var);
let (ne, c2) = rewrite_pass(ctx, *e, var);
if !c1 && !c2 {
return (expr, false);
}
(ctx.pow(nb, ne), true)
}
AtomNode::Num(_) | AtomNode::Var(_) | AtomNode::Fun(_, _) => (expr, false),
}
}
fn rewrite_add<'a>(ctx: &'a AtomArena<'a>, kids: &[Atom<'a>], var: Symbol) -> Option<Atom<'a>> {
if kids.len() != 2 {
return None;
}
let (x, y) = (kids[0], kids[1]);
if let Some(r) = rewrite_const_pair(ctx, x, y, var) {
return Some(r);
}
if let Some(r) = rewrite_const_pair(ctx, y, x, var) {
return Some(r);
}
if let Some(r) = rewrite_square_pair(ctx, x, y, var) {
return Some(r);
}
rewrite_square_pair(ctx, y, x, var)
}
fn rewrite_const_pair<'a>(
ctx: &'a AtomArena<'a>,
constant_term: Atom<'a>,
square_term: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let c = constant_term_of(ctx, constant_term, var)?;
let (c2, name, arg) = square_term_of(ctx, square_term, var)?;
let neg_c = normalize(ctx, ctx.mul(&[ctx.num(-1), c]));
let target = match name {
"tanh" | "sin" => ("sech", "cos", neg_c),
"cot" | "tan" => ("csc", "sec", c),
_ => return None,
};
if c2 != target.2 {
return None;
}
let result_name = if name == "tanh" {
target.0
} else if name == "sin" {
target.1
} else if name == "cot" {
target.0
} else {
target.1
};
let sq = int_pow(ctx, ctx.fun(result_name, &[arg]), 2);
Some(normalize(ctx, ctx.mul(&[c, sq])))
}
fn rewrite_square_pair<'a>(
ctx: &'a AtomArena<'a>,
first: Atom<'a>,
second: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let (c1, name1, arg1) = square_term_of(ctx, first, var)?;
if name1 == "sec"
&& let Some(c2) = constant_term_of(ctx, second, var)
&& c2 == normalize(ctx, ctx.mul(&[ctx.num(-1), c1]))
{
let sq = int_pow(ctx, ctx.fun("tan", &[arg1]), 2);
return Some(normalize(ctx, ctx.mul(&[c1, sq])));
}
if name1 == "cosh"
&& let Some((c2, name2, arg2)) = square_term_of(ctx, second, var)
&& name2 == "sinh"
&& arg2 == arg1
&& c2 == normalize(ctx, ctx.mul(&[ctx.num(-1), c1]))
{
return Some(c1);
}
None
}
fn constant_term_of<'a>(ctx: &'a AtomArena<'a>, term: Atom<'a>, var: Symbol) -> Option<Atom<'a>> {
if is_constant(term, var) {
Some(normalize(ctx, term))
} else {
None
}
}
fn square_term_of<'a>(
ctx: &'a AtomArena<'a>,
term: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, &'static str, Atom<'a>)> {
let (coeff, rest) = split_coeff(ctx, term, var);
let AtomNode::Pow(base, exp) = rest.node() else {
return None;
};
if !matches!(exp.node(), AtomNode::Num(2)) {
return None;
}
let AtomNode::Fun(name, args) = base.node() else {
return None;
};
if args.len() != 1 {
return None;
}
let name = match name.as_str() {
"tan" => "tan",
"cot" => "cot",
"tanh" => "tanh",
"sin" => "sin",
"cos" => "cos",
"sinh" => "sinh",
"cosh" => "cosh",
"sec" => "sec",
"csc" => "csc",
_ => return None,
};
Some((coeff, name, args[0]))
}
fn split_coeff<'a>(ctx: &'a AtomArena<'a>, term: Atom<'a>, var: Symbol) -> (Atom<'a>, Atom<'a>) {
let AtomNode::Mul(args) = term.node() else {
return (ctx.num(1), term);
};
let mut coeff: Vec<Atom<'a>> = Vec::new();
let mut rest: Vec<Atom<'a>> = Vec::new();
for a in args.iter() {
if is_constant(*a, var) {
coeff.push(*a);
} else {
rest.push(*a);
}
}
let c = match coeff.len() {
0 => ctx.num(1),
1 => normalize(ctx, coeff[0]),
_ => normalize(ctx, ctx.mul(&coeff)),
};
let r = match rest.len() {
0 => ctx.num(1),
1 => rest[0],
_ => normalize(ctx, ctx.mul(&rest)),
};
(c, r)
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Family {
Trig,
Hyper,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Kernel {
Tan,
Cot,
Tanh,
Coth,
}
impl Kernel {
fn name(self) -> &'static str {
match self {
Kernel::Tan => "tan",
Kernel::Cot => "cot",
Kernel::Tanh => "tanh",
Kernel::Coth => "coth",
}
}
fn family(self) -> Family {
match self {
Kernel::Tan | Kernel::Cot => Family::Trig,
Kernel::Tanh | Kernel::Coth => Family::Hyper,
}
}
fn quad<'a>(self, ctx: &'a AtomArena<'a>, t: Atom<'a>) -> Atom<'a> {
let t2 = int_pow(ctx, t, 2);
match self {
Kernel::Tan | Kernel::Cot => ctx.add(&[ctx.num(1), t2]),
Kernel::Tanh | Kernel::Coth => {
ctx.add(&[ctx.num(1), normalize(ctx, ctx.mul(&[ctx.num(-1), t2]))])
}
}
}
}
struct KernelInfo<'a> {
u: Atom<'a>,
slope: Atom<'a>,
family: Family,
seen: Vec<&'static str>,
}
impl<'a> KernelInfo<'a> {
fn candidates(&self) -> Vec<Kernel> {
let all: [Kernel; 2] = match self.family {
Family::Trig => [Kernel::Tan, Kernel::Cot],
Family::Hyper => [Kernel::Tanh, Kernel::Coth],
};
let mut out: Vec<Kernel> = Vec::with_capacity(2);
for k in all {
if self.seen.iter().any(|n| *n == k.name()) {
out.push(k);
}
}
for k in all {
if !out.contains(&k) {
out.push(k);
}
}
out
}
}
fn match_family<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, var: Symbol) -> Option<KernelInfo<'a>> {
let mut u: Option<Atom<'a>> = None;
let mut family: Option<Family> = None;
let mut seen: Vec<&'static str> = Vec::new();
let mut sites = 0usize;
scan_family(ctx, expr, var, &mut u, &mut family, &mut seen, &mut sites)?;
let u = u?;
let family = family?;
let (slope, _b) = linear_form(ctx, u, var)?;
if matches!(normalize(ctx, slope).node(), AtomNode::Num(0)) {
return None;
}
Some(KernelInfo {
u: normalize(ctx, u),
slope: normalize(ctx, slope),
family,
seen,
})
}
fn scan_family<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
u: &mut Option<Atom<'a>>,
family: &mut Option<Family>,
seen: &mut Vec<&'static str>,
sites: &mut usize,
) -> Option<()> {
if is_constant(expr, var) {
return Some(());
}
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => None,
AtomNode::Add(args) | AtomNode::Mul(args) => {
for a in args.iter() {
scan_family(ctx, *a, var, u, family, seen, sites)?;
}
Some(())
}
AtomNode::Pow(b, e) => {
scan_family(ctx, *b, var, u, family, seen, sites)?;
scan_family(ctx, *e, var, u, family, seen, sites)
}
AtomNode::Fun(name, args) => {
if name.as_str() == "sqrt" && args.len() == 1 {
return scan_family(ctx, args[0], var, u, family, seen, sites);
}
let (k, fam) = match name.as_str() {
"tan" => (Kernel::Tan, Family::Trig),
"cot" => (Kernel::Cot, Family::Trig),
"tanh" => (Kernel::Tanh, Family::Hyper),
"coth" => (Kernel::Coth, Family::Hyper),
_ => return None,
};
if args.len() != 1 {
return None;
}
match *family {
Some(f) if f != fam => return None,
None => *family = Some(fam),
_ => {}
}
let arg = normalize(ctx, args[0]);
match *u {
Some(u0) if u0 != arg => return None,
None => *u = Some(arg),
_ => {}
}
*sites += 1;
if *sites > MAX_SITES {
return None;
}
if !seen.contains(&k.name()) {
seen.push(k.name());
}
Some(())
}
}
}
fn try_primary<'a>(
ctx: &'a AtomArena<'a>,
original: Atom<'a>,
expr: Atom<'a>,
var: Symbol,
info: &KernelInfo<'a>,
primary: Kernel,
) -> Option<Atom<'a>> {
let t_sym = pick_subst_symbol(expr, var)?;
let t = ctx.var(t_sym.as_str());
let substituted = subst_expr(ctx, expr, var, info, primary, t)?;
if !is_admissible_t(substituted, var, t_sym) {
return None;
}
let mut factors: Vec<Atom<'a>> = vec![substituted];
let slope = info.slope;
if !matches!(slope.node(), AtomNode::Num(1)) {
factors.push(inv(ctx, slope));
}
factors.push(inv(ctx, primary.quad(ctx, t)));
if primary == Kernel::Cot {
factors.push(ctx.num(-1));
}
let t_form = normalize(ctx, ctx.mul(&factors));
let t_form = combine_like_powers(ctx, t_form, t_sym);
if node_count(t_form) > MAX_NODES || !is_admissible_t(t_form, var, t_sym) {
return None;
}
let (consts, core) = split_constants(ctx, t_form, t_sym);
let anti = if is_constant(core, t_sym) {
if matches!(normalize(ctx, core).node(), AtomNode::Num(0)) {
return None;
}
normalize(ctx, ctx.mul(&[core, t]))
} else {
let anti = integrate_algebraic(ctx, core, t_sym, 0)?;
if contains_integral(anti) {
return None;
}
anti
};
let back = replace_symbol(ctx, anti, t_sym, ctx.fun(primary.name(), &[info.u]));
let back = normalize(ctx, back);
let candidate = if consts.is_empty() {
back
} else {
let mut out: Vec<Atom<'a>> = consts;
out.push(back);
normalize(ctx, ctx.mul(&out))
};
if verify_candidate(ctx, original, candidate, var) {
Some(candidate)
} else {
None
}
}
const VERIFY_PARAMS: [f64; 10] = [0.7, 1.3, 0.37, 0.91, 1.73, 0.53, 1.11, 0.83, 1.47, 0.61];
const VERIFY_SAMPLES: [f64; 8] = [-1.7, -0.9, -0.37, 0.31, 0.77, 1.23, 1.91, 2.53];
const VERIFY_T_SAMPLES: [f64; 5] = [-0.77, -0.41, 0.31, 0.77, 1.23];
const VERIFY_TOL: f64 = 1e-4;
const VERIFY_MIN_SAMPLES: usize = 2;
fn verify_derivative<'a>(
_ctx: &'a AtomArena<'a>,
form: Atom<'a>,
candidate: Atom<'a>,
var: Symbol,
samples: &[f64],
) -> bool {
let mut symbols: Vec<Symbol> = Vec::new();
collect_symbols(form, var, &mut symbols);
collect_symbols(candidate, var, &mut symbols);
symbols.sort_by_key(|s| s.as_str().to_string());
symbols.dedup();
if symbols.len() > VERIFY_PARAMS.len() {
return false;
}
let mut usable = 0usize;
for &xv in samples.iter() {
let mut env: Vec<(Symbol, f64)> = symbols
.iter()
.enumerate()
.map(|(i, s)| (*s, VERIFY_PARAMS[i]))
.collect();
env.push((var, xv));
let Some(rhs) = eval_real(form, &env) else {
continue;
};
if !rhs.is_finite() {
continue;
}
let h = 1e-4 * xv.abs().max(1.0);
let mut stencil = [0.0f64; 4];
let mut ok = true;
for (slot, factor) in stencil.iter_mut().zip([-2.0f64, -1.0, 1.0, 2.0]) {
if let Some((_, last)) = env.last_mut() {
*last = xv + factor * h;
}
let Some(v) = eval_real(candidate, &env) else {
ok = false;
break;
};
if !v.is_finite() {
ok = false;
break;
}
*slot = v;
}
if !ok {
continue;
}
let lhs = (stencil[0] - 8.0 * stencil[1] + 8.0 * stencil[2] - stencil[3]) / (12.0 * h);
usable += 1;
if (lhs - rhs).abs() > VERIFY_TOL * rhs.abs().max(1.0) {
if std::env::var_os("OCAS_KERNEL_VERIFY_DEBUG").is_some() {
eprintln!(
"[kernel_subst verify] mismatch on {form} at {var:?}={xv}:\n \
diff={lhs} form={rhs}\n candidate: {candidate}"
);
}
return false;
}
}
if usable < VERIFY_MIN_SAMPLES && std::env::var_os("OCAS_KERNEL_VERIFY_DEBUG").is_some() {
eprintln!("[kernel_subst verify] only {usable} usable samples for {form}\n {candidate}");
}
usable >= VERIFY_MIN_SAMPLES
}
fn verify_local<'a>(
ctx: &'a AtomArena<'a>,
form: Atom<'a>,
candidate: Atom<'a>,
var: Symbol,
) -> bool {
verify_derivative(ctx, form, candidate, var, &VERIFY_T_SAMPLES)
}
fn verify_candidate<'a>(
ctx: &'a AtomArena<'a>,
integrand: Atom<'a>,
candidate: Atom<'a>,
var: Symbol,
) -> bool {
verify_derivative(ctx, integrand, candidate, var, &VERIFY_SAMPLES)
}
fn collect_symbols<'a>(expr: Atom<'a>, var: Symbol, out: &mut Vec<Symbol>) {
match expr.node() {
AtomNode::Num(_) => {}
AtomNode::Var(v) => {
if *v != var {
out.push(*v);
}
}
AtomNode::Add(args) | AtomNode::Mul(args) | AtomNode::Fun(_, args) => {
for a in args.iter() {
collect_symbols(*a, var, out);
}
}
AtomNode::Pow(b, e) => {
collect_symbols(*b, var, out);
collect_symbols(*e, var, out);
}
}
}
fn eval_real(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_real(*a, env)?)),
AtomNode::Mul(args) => args
.iter()
.try_fold(1.0, |acc, a| Some(acc * eval_real(*a, env)?)),
AtomNode::Pow(b, e) => Some(eval_real(*b, env)?.powf(eval_real(*e, env)?)),
AtomNode::Fun(name, args) => {
let v = eval_real(*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(),
"sinh" => v.sinh(),
"cosh" => v.cosh(),
"tanh" => v.tanh(),
"coth" => 1.0 / v.tanh(),
"sech" => 1.0 / v.cosh(),
"csch" => 1.0 / v.sinh(),
"exp" => v.exp(),
"log" => v.abs().ln(),
"sqrt" => v.sqrt(),
"atan" => v.atan(),
"asin" => v.asin(),
"acos" => v.acos(),
"atanh" => v.atanh(),
"asinh" => v.asinh(),
"acosh" => v.acosh(),
_ => return None,
})
}
}
}
fn subst_expr<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
info: &KernelInfo<'a>,
primary: Kernel,
t: Atom<'a>,
) -> Option<Atom<'a>> {
if is_constant(expr, var) {
return Some(expr);
}
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => None,
AtomNode::Add(args) => {
let mut kids = Vec::with_capacity(args.len());
for a in args.iter() {
kids.push(subst_expr(ctx, *a, var, info, primary, t)?);
}
Some(ctx.add(&kids))
}
AtomNode::Mul(args) => {
let mut kids = Vec::with_capacity(args.len());
for a in args.iter() {
kids.push(subst_expr(ctx, *a, var, info, primary, t)?);
}
Some(ctx.mul(&kids))
}
AtomNode::Pow(b, e) => {
if let AtomNode::Fun(name, args) = b.node()
&& args.len() == 1
&& let Some(img) = kernel_image(ctx, name.as_str(), primary, t)
{
if normalize(ctx, args[0]) != info.u {
return None;
}
let ne = subst_expr(ctx, *e, var, info, primary, t)?;
return Some(normalize(ctx, ctx.pow(img, ne)));
}
let nb = subst_expr(ctx, *b, var, info, primary, t)?;
let ne = subst_expr(ctx, *e, var, info, primary, t)?;
Some(normalize(ctx, ctx.pow(nb, ne)))
}
AtomNode::Fun(name, args) => {
if args.len() != 1 {
return None;
}
if name.as_str() == "sqrt" {
let inner = subst_expr(ctx, args[0], var, info, primary, t)?;
return Some(normalize(ctx, ctx.pow(inner, rat_atom(ctx, 1, 2))));
}
let img = kernel_image(ctx, name.as_str(), primary, t)?;
if normalize(ctx, args[0]) != info.u {
return None;
}
Some(img)
}
}
}
fn kernel_image<'a>(
ctx: &'a AtomArena<'a>,
name: &str,
primary: Kernel,
t: Atom<'a>,
) -> Option<Atom<'a>> {
let in_family = matches!(name, "tan" | "cot" | "tanh" | "coth");
let same_group = match primary.family() {
Family::Trig => matches!(name, "tan" | "cot"),
Family::Hyper => matches!(name, "tanh" | "coth"),
};
if !in_family || !same_group {
return None;
}
if name == primary.name() {
return Some(t);
}
Some(inv(ctx, t))
}
fn is_admissible_t<'a>(expr: Atom<'a>, var: Symbol, t_sym: Symbol) -> bool {
match expr.node() {
AtomNode::Num(_) => true,
AtomNode::Var(v) => *v != var,
AtomNode::Add(args) | AtomNode::Mul(args) => {
args.iter().all(|a| is_admissible_t(*a, var, t_sym))
}
AtomNode::Pow(b, e) => {
if !is_admissible_t(*e, var, t_sym) {
return false;
}
if is_constant(*b, t_sym) {
return true;
}
match rat_exponent(*e) {
Some((_p, q)) => q <= MAX_EXP_DEN && is_admissible_t(*b, var, t_sym),
None => false,
}
}
AtomNode::Fun(_, _) => is_constant(expr, t_sym),
}
}
fn split_constant_powers<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, t_sym: Symbol) -> Atom<'a> {
match expr.node() {
AtomNode::Add(args) => {
let kids: Vec<Atom<'a>> = args
.iter()
.map(|a| split_constant_powers(ctx, *a, t_sym))
.collect();
ctx.add(&kids)
}
AtomNode::Mul(args) => {
let kids: Vec<Atom<'a>> = args
.iter()
.map(|a| split_constant_powers(ctx, *a, t_sym))
.collect();
ctx.mul(&kids)
}
AtomNode::Pow(b, e) => {
let nb = split_constant_powers(ctx, *b, t_sym);
let ne = split_constant_powers(ctx, *e, t_sym);
if let Some((_p, q)) = rat_exponent(ne)
&& q > 1
&& let AtomNode::Mul(factors) = nb.node()
{
let mut consts: Vec<Atom<'a>> = Vec::new();
let mut rest: Vec<Atom<'a>> = Vec::new();
for f in factors.iter() {
if is_constant(*f, t_sym) {
consts.push(*f);
} else {
rest.push(*f);
}
}
if !consts.is_empty() && !rest.is_empty() {
let c = normalize(ctx, ctx.pow(normalize(ctx, ctx.mul(&consts)), ne));
let r = normalize(ctx, ctx.pow(normalize(ctx, ctx.mul(&rest)), ne));
return normalize(ctx, ctx.mul(&[c, r]));
}
}
normalize(ctx, ctx.pow(nb, ne))
}
AtomNode::Num(_) | AtomNode::Var(_) | AtomNode::Fun(_, _) => expr,
}
}
fn combine_like_powers<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, t_sym: Symbol) -> Atom<'a> {
let expr = split_constant_powers(ctx, expr, t_sym);
match expr.node() {
AtomNode::Add(args) => {
let kids: Vec<Atom<'a>> = args
.iter()
.map(|a| combine_like_powers(ctx, *a, t_sym))
.collect();
ctx.add(&kids)
}
AtomNode::Mul(args) => {
let kids: Vec<Atom<'a>> = args
.iter()
.map(|a| combine_like_powers(ctx, *a, t_sym))
.collect();
combine_mul(ctx, &kids, t_sym)
}
AtomNode::Pow(b, e) => {
let nb = combine_like_powers(ctx, *b, t_sym);
let ne = combine_like_powers(ctx, *e, t_sym);
match rat_exponent(ne) {
Some((p, 1)) => int_pow(ctx, nb, p),
Some((p, q)) => {
if let AtomNode::Pow(inner, inner_exp) = nb.node()
&& let Some((m, 1)) = rat_exponent(*inner_exp)
&& m % 2 == 0
&& let Some(n) = m.checked_mul(p)
&& n % q == 0
&& (n / q) % 2 == 0
{
return int_pow(ctx, *inner, n / q);
}
ctx.pow(nb, ne)
}
None => ctx.pow(nb, ne),
}
}
AtomNode::Num(_) | AtomNode::Var(_) | AtomNode::Fun(_, _) => expr,
}
}
fn combine_mul<'a>(ctx: &'a AtomArena<'a>, factors: &[Atom<'a>], t_sym: Symbol) -> Atom<'a> {
let mut groups: Vec<(Atom<'a>, i64, i64)> = Vec::new();
let mut rest: Vec<Atom<'a>> = Vec::new();
for f in factors {
let Some((base, p, q)) = groupable_factor(*f, t_sym) else {
rest.push(*f);
continue;
};
if let Some(slot) = groups.iter_mut().find(|(b, _, _)| *b == base) {
let num = slot
.1
.checked_mul(q)
.and_then(|v| v.checked_add(p * slot.2));
let den = slot.2.checked_mul(q);
let (Some(num), Some(den)) = (num, den) else {
rest.push(*f);
continue;
};
let g = gcd_i64(num, den).max(1);
slot.1 = num / g;
slot.2 = den / g;
} else if groups.len() < MAX_POW_GROUPS {
groups.push((base, p, q));
} else {
rest.push(*f);
}
}
let mut out: Vec<Atom<'a>> = Vec::with_capacity(rest.len() + groups.len());
for (base, p, q) in groups {
if p == 0 {
continue;
}
if p.abs() > MAX_POW_EXP || q.abs() > MAX_POW_EXP {
out.push(base);
continue;
}
if q == 1 {
out.push(int_pow(ctx, base, p));
} else {
out.push(ctx.pow(base, rat_atom(ctx, p, q)));
}
}
out.extend(rest);
if out.is_empty() {
return ctx.num(1);
}
normalize(ctx, ctx.mul(&out))
}
fn groupable_factor<'a>(f: Atom<'a>, t_sym: Symbol) -> Option<(Atom<'a>, i64, i64)> {
match f.node() {
AtomNode::Pow(b, e) if !is_constant(*b, t_sym) => {
let (p, q) = rat_exponent(*e)?;
Some((*b, p, q))
}
_ if !is_constant(f, t_sym) => Some((f, 1, 1)),
_ => None,
}
}
fn split_constants<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
t_sym: Symbol,
) -> (Vec<Atom<'a>>, Atom<'a>) {
let AtomNode::Mul(args) = expr.node() else {
if is_constant(expr, t_sym) {
return (vec![expr], ctx.num(1));
}
return (Vec::new(), expr);
};
let mut consts: Vec<Atom<'a>> = Vec::new();
let mut rest: Vec<Atom<'a>> = Vec::new();
for a in args.iter() {
if is_constant(*a, t_sym) {
consts.push(*a);
} else {
rest.push(*a);
}
}
let core = match rest.len() {
0 => ctx.num(1),
1 => rest[0],
_ => normalize(ctx, ctx.mul(&rest)),
};
(consts, core)
}
fn integrate_algebraic<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
depth: usize,
) -> Option<Atom<'a>> {
if node_count(expr) > MAX_NODES {
return None;
}
let (consts, core) = split_constants(ctx, expr, var);
if is_constant(core, var) {
return None;
}
if let Some(r) = try_engines(ctx, core, var) {
return Some(scale_by(ctx, &consts, r));
}
if depth < MAX_EXPAND_DEPTH
&& let Some(expanded) = crate::expand::expand_bounded(ctx, core)
{
let folded = crate::ode::util::collect_terms(ctx, expanded);
if let Some(r) = integrate_terms(ctx, folded, var, depth + 1)
&& verify_local(ctx, core, r, var)
{
return Some(scale_by(ctx, &consts, r));
}
}
if let Some(r) = try_radical_rationalize(ctx, core, var)
&& verify_local(ctx, core, r, var)
{
return Some(scale_by(ctx, &consts, r));
}
if let Some(r) = integrate_rational_complete(ctx, core, var)
&& verify_local(ctx, core, r, var)
{
return Some(scale_by(ctx, &consts, r));
}
None
}
fn scale_by<'a>(ctx: &'a AtomArena<'a>, consts: &[Atom<'a>], r: Atom<'a>) -> Atom<'a> {
if consts.is_empty() {
return r;
}
let mut factors: Vec<Atom<'a>> = consts.to_vec();
factors.push(r);
normalize(ctx, ctx.mul(&factors))
}
fn integrate_terms<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
depth: usize,
) -> Option<Atom<'a>> {
match expr.node() {
AtomNode::Add(args) => {
let mut out: Vec<Atom<'a>> = Vec::with_capacity(args.len());
for a in args.iter() {
out.push(integrate_algebraic(ctx, *a, var, depth)?);
}
Some(normalize(ctx, ctx.add(&out)))
}
_ => integrate_algebraic(ctx, expr, var, depth),
}
}
fn try_engines<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, var: Symbol) -> Option<Atom<'a>> {
if let Some(r) = binomial::integrate_binomial(ctx, expr, var)
&& !contains_integral(r)
&& verify_local(ctx, expr, r, var)
{
return Some(r);
}
if let Some(r) = sqrt_quadratic::integrate_sqrt_quadratic(ctx, expr, var)
&& !contains_integral(r)
&& verify_local(ctx, expr, r, var)
{
return Some(r);
}
if let Some(r) = rational::integrate_rational(ctx, expr, var)
&& !contains_integral(r)
&& verify_local(ctx, expr, r, var)
{
return Some(r);
}
None
}
fn try_radical_rationalize<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let mut base: Option<Atom<'a>> = None;
let mut dens: Vec<i64> = Vec::new();
let mut sites = 0usize;
collect_radical_sites(ctx, expr, var, &mut base, &mut dens, &mut sites)?;
let g = base?;
let (beta, alpha) = linear_form(ctx, g, var)?;
if matches!(normalize(ctx, beta).node(), AtomNode::Num(0)) {
return None;
}
let l = dens.iter().try_fold(1i64, |acc, s| lcm_i64(acc, *s))?;
if !(2..=MAX_RADICAL_DEG).contains(&l) {
return None;
}
let w_sym = pick_subst_symbol(expr, var)?;
let w = ctx.var(w_sym.as_str());
let var_w = normalize(
ctx,
ctx.mul(&[
normalize(
ctx,
ctx.add(&[
int_pow(ctx, w, l),
normalize(ctx, ctx.mul(&[ctx.num(-1), alpha])),
]),
),
inv(ctx, beta),
]),
);
let subbed = subst_radical(ctx, expr, var, g, l, w, var_w)?;
let dvar = normalize(
ctx,
ctx.mul(&[ctx.num(l), inv(ctx, beta), int_pow(ctx, w, l - 1)]),
);
let w_form = normalize(ctx, ctx.mul(&[subbed, dvar]));
if node_count(w_form) > MAX_NODES || !is_rational_form(w_form, var) {
return None;
}
let anti = integrate_rational_complete(ctx, w_form, w_sym)?;
let back = replace_symbol(ctx, anti, w_sym, ctx.pow(g, rat_atom(ctx, 1, l)));
Some(normalize(ctx, back))
}
fn collect_radical_sites<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
base: &mut Option<Atom<'a>>,
dens: &mut Vec<i64>,
sites: &mut usize,
) -> Option<()> {
match expr.node() {
AtomNode::Pow(b, e) => {
if let Some((_, s)) = rat_exponent(*e)
&& s > 1
&& !is_constant(*b, var)
{
if s > MAX_RADICAL_DEG {
return None;
}
let nb = normalize(ctx, *b);
match *base {
Some(g0) if nb != g0 => return None,
None => *base = Some(nb),
_ => {}
}
dens.push(s);
*sites += 1;
if *sites > MAX_RADICAL_SITES {
return None;
}
}
collect_radical_sites(ctx, *b, var, base, dens, sites)?;
collect_radical_sites(ctx, *e, var, base, dens, sites)
}
AtomNode::Fun(name, args) if name.as_str() == "sqrt" && args.len() == 1 => {
if !is_constant(args[0], var) {
let nb = normalize(ctx, args[0]);
match *base {
Some(g0) if nb != g0 => return None,
None => *base = Some(nb),
_ => {}
}
dens.push(2);
*sites += 1;
if *sites > MAX_RADICAL_SITES {
return None;
}
}
collect_radical_sites(ctx, args[0], var, base, dens, sites)
}
AtomNode::Add(args) | AtomNode::Mul(args) | AtomNode::Fun(_, args) => {
for a in args.iter() {
collect_radical_sites(ctx, *a, var, base, dens, sites)?;
}
Some(())
}
AtomNode::Num(_) | AtomNode::Var(_) => Some(()),
}
}
fn subst_radical<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
g: Atom<'a>,
l: i64,
w: Atom<'a>,
var_w: Atom<'a>,
) -> Option<Atom<'a>> {
match expr.node() {
AtomNode::Num(_) => Some(expr),
AtomNode::Var(v) => {
if *v == var {
Some(var_w)
} else {
Some(expr)
}
}
AtomNode::Add(args) => {
let mut kids = Vec::with_capacity(args.len());
for a in args.iter() {
kids.push(subst_radical(ctx, *a, var, g, l, w, var_w)?);
}
Some(ctx.add(&kids))
}
AtomNode::Mul(args) => {
let mut kids = Vec::with_capacity(args.len());
for a in args.iter() {
kids.push(subst_radical(ctx, *a, var, g, l, w, var_w)?);
}
Some(ctx.mul(&kids))
}
AtomNode::Pow(b, e) => {
if normalize(ctx, *b) == g
&& let Some((k, s)) = rat_exponent(*e)
&& s > 1
{
return Some(int_pow(ctx, w, k.checked_mul(l / s)?));
}
let nb = subst_radical(ctx, *b, var, g, l, w, var_w)?;
let ne = subst_radical(ctx, *e, var, g, l, w, var_w)?;
Some(normalize(ctx, ctx.pow(nb, ne)))
}
AtomNode::Fun(name, args) if name.as_str() == "sqrt" && args.len() == 1 => {
if normalize(ctx, args[0]) == g {
return Some(int_pow(ctx, w, l / 2));
}
let na = subst_radical(ctx, args[0], var, g, l, w, var_w)?;
Some(normalize(ctx, ctx.pow(na, rat_atom(ctx, 1, 2))))
}
AtomNode::Fun(name, args) => {
let mut kids = Vec::with_capacity(args.len());
for a in args.iter() {
kids.push(subst_radical(ctx, *a, var, g, l, w, var_w)?);
}
Some(ctx.fun(name.as_str(), &kids))
}
}
}
fn is_rational_form<'a>(expr: Atom<'a>, orig_var: Symbol) -> bool {
match expr.node() {
AtomNode::Num(_) => true,
AtomNode::Var(v) => *v != orig_var,
AtomNode::Add(args) | AtomNode::Mul(args) => {
args.iter().all(|a| is_rational_form(*a, orig_var))
}
AtomNode::Pow(b, e) => {
matches!(e.node(), AtomNode::Num(_)) && is_rational_form(*b, orig_var)
}
AtomNode::Fun(_, _) => false,
}
}
fn integrate_rational_complete<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let x = ctx.var(var.as_str());
let rf = atom_to_rational(expr, &[x])?;
let num = sparse_to_dense(&rf.numerator)?;
let den = sparse_to_dense(&rf.denominator)?;
if den.is_zero() || den.degree()? > MAX_RAT_DEG || num.degree().unwrap_or(0) > MAX_RAT_NUM_DEG {
return None;
}
let (quotient, remainder) = num.div_rem(&den)?;
let mut parts: Vec<Atom<'a>> = Vec::new();
if !quotient.is_zero() {
parts.push(poly_atom(ctx, &rational::poly_integrate("ient), x)?);
}
if !remainder.is_zero() {
let lc = den.lcoeff();
let inv_lc = RationalDomain.inv(&lc)?;
let den = den.mul_scalar(&inv_lc);
let remainder = remainder.mul_scalar(&inv_lc);
let (num, den) = reduce_pair(&remainder, &den)?;
if !num.is_zero() {
let even = if den.degree() == Some(4) && is_even_poly(&den) && is_even_poly(&num) {
even_quartic_atom(ctx, &num, &den, x)
} else {
None
};
let piece = match even {
Some(p) => p,
None => partial_fraction_atom(ctx, &num, &den, x, var)?,
};
parts.push(piece);
}
}
match parts.len() {
0 => Some(ctx.num(0)),
1 => Some(parts.remove(0)),
_ => Some(normalize(ctx, ctx.add(&parts))),
}
}
fn reduce_pair(num: &DPoly, den: &DPoly) -> Option<(DPoly, DPoly)> {
let cloned = || (num.clone(), den.clone());
if num.is_zero() {
return Some(cloned());
}
let (g, _, _) = num.extended_gcd_poly(den);
if g.is_zero() || g.degree() == Some(0) {
return Some(cloned());
}
let (nq, nr) = num.div_rem(&g)?;
let (dq, dr) = den.div_rem(&g)?;
if !nr.is_zero() || !dr.is_zero() {
return Some(cloned());
}
Some((nq, dq))
}
fn is_even_poly(p: &DPoly) -> bool {
p.coeffs()
.iter()
.enumerate()
.all(|(i, c)| i % 2 == 0 || RationalDomain.is_zero(c))
}
fn partial_fraction_atom<'a>(
ctx: &'a AtomArena<'a>,
num: &DPoly,
den: &DPoly,
x: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let factors = factor_den(den)?;
if factors.is_empty() || factors.len() > MAX_FACTORS {
return None;
}
let (poly, pieces) = peel(num, den, &factors)?;
if pieces.len() > MAX_PIECES {
return None;
}
let mut parts: Vec<Atom<'a>> = Vec::new();
if !poly.is_zero() {
parts.push(poly_atom(ctx, &rational::poly_integrate(&poly), x)?);
}
for (n, d) in pieces {
let piece = ratio_atom(ctx, &n, &d, x)?;
let integrated = rational::integrate_rational(ctx, piece, var)?;
if contains_integral(integrated) {
return None;
}
parts.push(integrated);
}
match parts.len() {
0 => Some(ctx.num(0)),
1 => Some(parts.remove(0)),
_ => Some(normalize(ctx, ctx.add(&parts))),
}
}
fn factor_den(den: &DPoly) -> Option<Vec<(DPoly, usize)>> {
let mut out: Vec<(DPoly, usize)> = Vec::new();
for (squarefree, mult) in den.square_free_factorization() {
if squarefree.degree() == Some(0) {
continue;
}
let irreducibles = factor_squarefree(&squarefree)?;
if out.len() + irreducibles.len() > MAX_FACTORS {
return None;
}
for f in irreducibles {
out.push((f, mult));
}
}
Some(out)
}
fn factor_squarefree(f: &DPoly) -> Option<Vec<DPoly>> {
let mut lcm: i64 = 1;
for c in f.coeffs() {
lcm = num_integer::lcm(lcm, c.denom().to_i64()?);
}
let mut zcoeffs: Vec<Integer> = Vec::with_capacity(f.coeffs().len());
for c in f.coeffs() {
let scaled = RationalDomain.mul(c, &Rational::new(lcm, 1));
zcoeffs.push(Integer::from(scaled.numer().to_i64()?));
}
let zpoly = DenseUnivariatePolynomial::from_coeffs(IntegerDomain, zcoeffs);
let factors = zpoly.primitive_part().factor();
let mut out: Vec<DPoly> = Vec::with_capacity(factors.len());
for (fac, _mult) in &factors {
let lc = fac.coeffs().last()?.to_i64()?;
if lc == 0 {
return None;
}
let coeffs: Option<Vec<Rational>> = fac
.coeffs()
.iter()
.map(|c| Some(Rational::new(c.to_i64()?, lc)))
.collect();
let dense = DPoly::from_coeffs(RationalDomain, coeffs?);
if dense.degree()? > 2 {
return None;
}
out.push(dense);
}
Some(out)
}
fn peel(
num: &DPoly,
den: &DPoly,
factors: &[(DPoly, usize)],
) -> Option<(DPoly, Vec<(DPoly, DPoly)>)> {
let zero = || DPoly::from_coeffs(RationalDomain, vec![]);
let mut remaining = den.clone();
let mut current = num.clone();
let mut poly = zero();
let mut pieces: Vec<(DPoly, DPoly)> = Vec::new();
for (f, mult) in factors {
let fm = f.pow(*mult as u32);
let (cofactor, rem) = remaining.div_rem(&fm)?;
if !rem.is_zero() {
return None;
}
if cofactor.degree() == Some(0) {
let c = cofactor.coeffs().first().cloned()?;
let inv_c = RationalDomain.inv(&c)?;
current = current.mul_scalar(&inv_c);
let (quotient, proper) = current.div_rem(&fm)?;
poly = poly.add("ient);
push_power_pieces(f, *mult, &proper, &mut pieces)?;
current = zero();
continue;
}
let (gcd, s, t) = fm.extended_gcd_poly(&cofactor);
if gcd.degree() != Some(0) {
return None;
}
let c = gcd.coeffs().first().cloned()?;
let inv_c = RationalDomain.inv(&c)?;
let s = s.mul_scalar(&inv_c);
let t = t.mul_scalar(&inv_c);
let nt = current.mul(&t);
let (quotient, proper) = nt.div_rem(&fm)?;
poly = poly.add("ient);
push_power_pieces(f, *mult, &proper, &mut pieces)?;
current = current.mul(&s);
remaining = cofactor;
}
if !current.is_zero() {
return None;
}
Some((poly, pieces))
}
fn push_power_pieces(
f: &DPoly,
mult: usize,
proper: &DPoly,
pieces: &mut Vec<(DPoly, DPoly)>,
) -> Option<()> {
let mut rest = proper.clone();
for j in (1..=mult).rev() {
let (quotient, rem) = rest.div_rem(f)?;
if !rem.is_zero() {
pieces.push((rem, f.pow(j as u32)));
}
rest = quotient;
}
if !rest.is_zero() {
return None;
}
Some(())
}
fn even_quartic_atom<'a>(
ctx: &'a AtomArena<'a>,
num: &DPoly,
den: &DPoly,
x: Atom<'a>,
) -> Option<Atom<'a>> {
let zero = RationalDomain.zero();
if num.degree().unwrap_or(0) > 2 || den.degree() != Some(4) {
return None;
}
let lc = den.coeffs().get(4).cloned().unwrap_or_else(|| zero.clone());
if lc != Rational::new(1, 1) {
return None;
}
let p = den.coeffs().get(2).cloned().unwrap_or_else(|| zero.clone());
let q = den
.coeffs()
.first()
.cloned()
.unwrap_or_else(|| zero.clone());
let n0 = num
.coeffs()
.first()
.cloned()
.unwrap_or_else(|| zero.clone());
let n1 = num.coeffs().get(2).cloned().unwrap_or_else(|| zero.clone());
if RationalDomain.is_zero(&n1) && RationalDomain.is_zero(&n0) {
return None;
}
let four = Rational::new(4, 1);
let disc = RationalDomain.sub(&RationalDomain.mul(&p, &p), &RationalDomain.mul(&four, &q));
if let Some(root) = sqrt_exact(&disc) {
let two = Rational::new(2, 1);
let neg_p = RationalDomain.neg(&p);
let r1 = RationalDomain.div(&RationalDomain.add(&neg_p, &root), &two)?;
let r2 = RationalDomain.div(&RationalDomain.sub(&neg_p, &root), &two)?;
if r1 == r2 {
return None;
}
let a_coef = RationalDomain.div(
&RationalDomain.sub(&n0, &RationalDomain.mul(&n1, &r1)),
&RationalDomain.sub(&r2, &r1),
)?;
let b_coef = RationalDomain.sub(&n1, &a_coef);
let ta = quadratic_reciprocal_atom(ctx, &a_coef, &r1, x)?;
let tb = quadratic_reciprocal_atom(ctx, &b_coef, &r2, x)?;
return Some(normalize(ctx, ctx.add(&[ta, tb])));
}
if !rat_is_neg(&disc) {
return None;
}
let two = ctx.num(2);
let f = sqrt_pos_atom(ctx, &q)?;
let p_atom = const_atom(ctx, &p)?;
let two_f = normalize(ctx, ctx.mul(&[two, f]));
let e_sq = normalize(
ctx,
ctx.add(&[two_f, normalize(ctx, ctx.mul(&[ctx.num(-1), p_atom]))]),
);
let e = ctx.pow(e_sq, rat_atom(ctx, 1, 2));
let delta = normalize(ctx, ctx.add(&[p_atom, two_f]));
let sq_delta = ctx.pow(delta, rat_atom(ctx, 1, 2));
let n0_atom = const_atom(ctx, &n0)?;
let n1_atom = const_atom(ctx, &n1)?;
let a_num = normalize(
ctx,
ctx.add(&[
ctx.mul(&[n1_atom, f]),
normalize(ctx, ctx.mul(&[ctx.num(-1), n0_atom])),
]),
);
let a_den = normalize(ctx, ctx.mul(&[two, e, f]));
let a_coef = normalize(ctx, ctx.mul(&[a_num, inv(ctx, a_den)]));
let w2 = ctx.pow(x, ctx.num(2));
let ew = normalize(ctx, ctx.mul(&[e, x]));
let top = normalize(
ctx,
ctx.add(&[w2, normalize(ctx, ctx.mul(&[ctx.num(-1), ew])), f]),
);
let bottom = normalize(ctx, ctx.add(&[w2, ew, f]));
let ratio = normalize(ctx, ctx.mul(&[top, inv(ctx, bottom)]));
let log_term = normalize(
ctx,
ctx.mul(&[a_coef, rat_atom(ctx, 1, 2), ctx.fun("log", &[ratio])]),
);
let c_num = normalize(ctx, ctx.add(&[n0_atom, ctx.mul(&[n1_atom, f])]));
let c_den = normalize(ctx, ctx.mul(&[two, f, sq_delta]));
let c_coef = normalize(ctx, ctx.mul(&[c_num, inv(ctx, c_den)]));
let two_w = normalize(ctx, ctx.mul(&[two, x]));
let shift = normalize(ctx, ctx.mul(&[ctx.num(-1), e]));
let arg1 = normalize(
ctx,
ctx.mul(&[normalize(ctx, ctx.add(&[two_w, shift])), inv(ctx, sq_delta)]),
);
let arg2 = normalize(
ctx,
ctx.mul(&[normalize(ctx, ctx.add(&[two_w, e])), inv(ctx, sq_delta)]),
);
let atan_term = normalize(
ctx,
ctx.mul(&[
c_coef,
ctx.add(&[ctx.fun("atan", &[arg1]), ctx.fun("atan", &[arg2])]),
]),
);
Some(normalize(ctx, ctx.add(&[log_term, atan_term])))
}
fn quadratic_reciprocal_atom<'a>(
ctx: &'a AtomArena<'a>,
c: &Rational,
r: &Rational,
x: Atom<'a>,
) -> Option<Atom<'a>> {
if RationalDomain.is_zero(c) {
return Some(ctx.num(0));
}
let c_atom = const_atom(ctx, c)?;
if RationalDomain.is_zero(r) {
return Some(normalize(ctx, ctx.mul(&[ctx.num(-1), c_atom, inv(ctx, x)])));
}
if rat_is_pos(r) {
let root = sqrt_pos_atom(ctx, r)?;
let arg = normalize(ctx, ctx.mul(&[x, inv(ctx, root)]));
return Some(normalize(
ctx,
ctx.mul(&[c_atom, inv(ctx, root), ctx.fun("atan", &[arg])]),
));
}
let minus_r = RationalDomain.neg(r);
let root = sqrt_pos_atom(ctx, &minus_r)?;
let top = normalize(
ctx,
ctx.add(&[x, normalize(ctx, ctx.mul(&[ctx.num(-1), root]))]),
);
let bottom = normalize(ctx, ctx.add(&[x, root]));
let ratio = normalize(ctx, ctx.mul(&[top, inv(ctx, bottom)]));
let log = ctx.fun("log", &[ratio]);
Some(normalize(
ctx,
ctx.mul(&[c_atom, inv(ctx, ctx.num(2)), inv(ctx, root), log]),
))
}
fn const_atom<'a>(ctx: &'a AtomArena<'a>, r: &Rational) -> Option<Atom<'a>> {
Some(rat_atom(ctx, r.numer().to_i64()?, r.denom().to_i64()?))
}
fn sqrt_exact(r: &Rational) -> Option<Rational> {
if rat_is_neg(r) {
return None;
}
let p = r.numer().to_i64()?;
let q = r.denom().to_i64()?;
let n = (p as i128).checked_mul(q as i128)?;
let m = n.checked_isqrt()?;
if m * m != n {
return None;
}
Some(Rational::new(i64::try_from(m).ok()?, q))
}
fn sqrt_pos_atom<'a>(ctx: &'a AtomArena<'a>, r: &Rational) -> Option<Atom<'a>> {
if let Some(exact) = sqrt_exact(r) {
return const_atom(ctx, &exact);
}
if rat_is_neg(r) {
return None;
}
let p = r.numer().to_i64()?;
let q = r.denom().to_i64()?;
if p == 0 {
return Some(ctx.num(0));
}
let n = p.checked_mul(q)?;
let root = ctx.pow(ctx.num(n), rat_atom(ctx, 1, 2));
Some(normalize(ctx, ctx.mul(&[root, inv(ctx, ctx.num(q))])))
}
fn rat_is_pos(r: &Rational) -> bool {
use num_traits::Signed;
r.inner().is_positive()
}
fn rat_is_neg(r: &Rational) -> bool {
use num_traits::Signed;
r.inner().is_negative()
}
fn poly_atom<'a>(ctx: &'a AtomArena<'a>, p: &DPoly, x: Atom<'a>) -> Option<Atom<'a>> {
if p.is_zero() {
return Some(ctx.num(0));
}
let field = GeneratorField::from_polynomial(dense_to_sparse(p));
rational_to_atom(ctx, &field, &[x])
}
fn ratio_atom<'a>(
ctx: &'a AtomArena<'a>,
num: &DPoly,
den: &DPoly,
x: Atom<'a>,
) -> Option<Atom<'a>> {
let field = GeneratorField::from_num_den(dense_to_sparse(num), dense_to_sparse(den));
rational_to_atom(ctx, &field, &[x])
}
fn sparse_to_dense(p: &Sparse) -> Option<DPoly> {
if p.n_vars() != 1 {
return None;
}
let deg = p.degree_in(0);
if deg > MAX_RAT_NUM_DEG {
return None;
}
let mut coeffs = vec![RationalDomain.zero(); deg + 1];
for (exp, coeff) in p.terms_ref() {
let index = *exp.first()?;
*coeffs.get_mut(index)? = coeff.clone();
}
Some(DPoly::from_coeffs(RationalDomain, coeffs))
}
fn dense_to_sparse(p: &DPoly) -> Sparse {
let terms = p
.coeffs()
.iter()
.enumerate()
.filter(|&(_, c)| !RationalDomain.is_zero(c))
.map(|(i, c)| (vec![i], c.clone()))
.collect();
Sparse::from_terms(RationalDomain, 1, terms)
}
#[cfg(test)]
mod tests {
use super::*;
use ocas_core::arena::Arena;
const SAMPLE_POS: [f64; 3] = [0.35, 0.75, 1.15];
const SAMPLE_SMALL: [f64; 3] = [0.4, 0.8, 1.2];
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.abs().ln(),
"sqrt" => v.sqrt(),
"atan" => v.atan(),
"asin" => v.asin(),
"acos" => v.acos(),
"sinh" => v.sinh(),
"cosh" => v.cosh(),
"tanh" => v.tanh(),
"coth" => 1.0 / v.tanh(),
"sech" => 1.0 / v.cosh(),
"csch" => 1.0 / v.sinh(),
"asinh" => v.asinh(),
"acosh" => v.acosh(),
"atanh" => v.atanh(),
_ => return None,
})
}
}
}
fn parse<'a>(ctx: &'a AtomArena<'a>, s: &str) -> Atom<'a> {
ocas_parse::parse(ctx, s).unwrap()
}
fn assert_antiderivative_num<'a>(
ctx: &'a AtomArena<'a>,
integrand: Atom<'a>,
var: Symbol,
consts: &[(Symbol, f64)],
samples: &[f64],
) {
let result = integrate_kernel_subst(ctx, integrand, var)
.unwrap_or_else(|| panic!("mechanism declined: {integrand}"));
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).unwrap_or_else(|| {
panic!(
"eval diff failed\n integrand: {integrand}\n result: {result}\n diff: {d}"
)
});
let rhs = eval_f64(integrand, &env)
.unwrap_or_else(|| panic!("eval integrand failed: {integrand}"));
let tol = 1e-6 * rhs.abs().max(1.0);
assert!(
(lhs - rhs).abs() < tol,
"integrand: {integrand}\nat x={xv}: diff={lhs} integrand={rhs}\nresult: {result}"
);
}
}
#[test]
fn rewritten_integrand_stays_integrated() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
assert_antiderivative_num(
&ctx,
parse(&ctx, "(1 - tanh(x)^2)^(3/2)"),
var,
&[],
&SAMPLE_POS,
);
}
#[test]
fn numeric_tan_family() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
for s in [
"sqrt(tan(x))",
"1/(2 + 3*tan(x))",
"1/(5 + 3*tan(x))^3",
"tan(x)^3/(1 + tan(x)^2)^(3/2)",
"tan(x)^2/sqrt(1 + tan(x)^2)",
"tan(x)^4",
"1/(1 + tan(x)^2)^2",
] {
assert_antiderivative_num(&ctx, parse(&ctx, s), var, &[], &SAMPLE_POS);
}
}
#[test]
fn numeric_cot_family() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
for s in [
"sqrt(cot(x))",
"1/(2 + 3*cot(x))",
"(1 + cot(x)^2)^(3/2)",
"(1 + cot(x)^2)^(5/2)",
"cot(x)^2*(1 + cot(x)^2)^(3/2)",
"cot(x)^4/(1 + cot(x)^2)",
] {
assert_antiderivative_num(&ctx, parse(&ctx, s), var, &[], &SAMPLE_POS);
}
}
#[test]
fn numeric_tanh_family() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
for s in [
"(1 - tanh(x)^2)^(3/2)",
"(1 - tanh(x)^2)^(5/2)",
"tanh(x)^3",
"1/(2 + 3*tanh(x))",
"tanh(x)^2/(1 + tanh(x)^2)",
] {
assert_antiderivative_num(&ctx, parse(&ctx, s), var, &[], &SAMPLE_POS);
}
}
#[test]
fn numeric_coth_family() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
for s in [
"(1 + coth(x))^(7/2)",
"sqrt(coth(x))",
"1/(2 + 3*coth(x))",
"coth(x)/(1 + coth(x)^2)",
"coth(x)^3/(1 + coth(x)^2)",
] {
assert_antiderivative_num(&ctx, parse(&ctx, s), var, &[], &SAMPLE_POS);
}
}
#[test]
fn symbolic_coefficient_corpus_shapes() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
let env = [
(Symbol::new("a"), 1.3),
(Symbol::new("b"), 0.7),
(Symbol::new("c"), 0.15),
(Symbol::new("d"), 0.8),
(Symbol::new("e"), 0.1),
(Symbol::new("f"), 0.7),
(Symbol::new("i"), 1.0),
];
for s in [
"coth(c + d*x)^4*(a + b*tanh(c + d*x)^2)^2",
"(1 + coth(c + d*x))^(7/2)",
"1/(5 + 3*tan(c + d*x))^3",
"(a + i*a*tan(e + f*x))^3/(d*tan(e + f*x))^(7/2)",
"sqrt(d*tan(e + f*x))*(a + i*a*tan(e + f*x))^2",
"a*tan(b*x)^3",
] {
assert_antiderivative_num(&ctx, parse(&ctx, s), var, &env, &SAMPLE_SMALL);
}
}
#[test]
fn declines_outside_the_family() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
for s in [
"sin(x)",
"sqrt(sin(x))",
"cos(tanh(a + b*x))",
"1/(a + b*sec(x))",
"csc(3*x + 1)^3",
"sec(2*x + 1)^4",
"sech(x)^3",
"(a*sin(e + f*x))^(3/2)*sqrt(b*tan(e + f*x))",
"tan(x)*exp(x)",
"1/(a + b*tan(c + d*x^(1/3)))^2",
"(b*tan(c + d*x)^2)^n",
"tan(x)*tanh(x)",
"1/((2+3*t)*(1+t^2))",
"coth(x)^2/(1 + coth(x)^2)^(3/2)",
] {
assert!(
integrate_kernel_subst(&ctx, parse(&ctx, s), var).is_none(),
"expected a decline: {s}"
);
}
}
#[test]
fn pythagorean_rewrites_are_exact_and_idempotent() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
for (input, expected) in [
("1 - tanh(x)^2", "sech(x)^2"),
("1 - sin(x)^2", "cos(x)^2"),
("sec(x)^2 - 1", "tan(x)^2"),
("1 + cot(x)^2", "csc(x)^2"),
("cosh(x)^2 - sinh(x)^2", "1"),
("1 + tan(x)^2", "sec(x)^2"),
("a - a*tanh(x)^2", "a*sech(x)^2"),
("a - a*sin(x)^2", "a*cos(x)^2"),
] {
let expr = parse(&ctx, input);
let out = pythagorean_rewrite(&ctx, expr, var)
.unwrap_or_else(|| panic!("no rewrite for {input}"));
let want = normalize(&ctx, parse(&ctx, expected));
assert_eq!(out, want, "{input} → {out}, want {want}");
assert!(
pythagorean_rewrite(&ctx, out, var).is_none(),
"not idempotent: {input}"
);
}
for s in [
"1 + 2*tanh(x)^2",
"tan(x)",
"1 + x",
"2 - tanh(x)^2",
"1 - tanh(x)",
"1 + sin(x)^2 - 2*sin(x)^2",
] {
assert!(
pythagorean_rewrite(&ctx, parse(&ctx, s), var).is_none(),
"unexpected rewrite: {s}"
);
}
}
#[test]
fn stress_repeated_runs_are_stable_and_bounded() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
let cases = [
"sqrt(tan(x))",
"1/(2 + 3*tan(x))",
"(1 + coth(x))^(7/2)",
"(1 - tanh(x)^2)^(3/2)",
"tanh(x)^3",
"coth(c + d*x)^4*(a + b*tanh(c + d*x)^2)^2",
"(a + i*a*tan(e + f*x))^3/(d*tan(e + f*x))^(7/2)",
"sin(x)",
"tan(x)*exp(x)",
"1/(a + b*sec(x))",
"csc(3*x + 1)^3",
"(1 + cot(x)^2)^(5/2)",
];
let mut baseline: Vec<bool> = vec![false; cases.len()];
for round in 0..25 {
for (index, s) in cases.iter().enumerate() {
let expr = parse(&ctx, s);
let outcome = integrate_kernel_subst(&ctx, expr, var);
let solved = outcome.is_some();
if let Some(r) = outcome {
assert!(r.to_string().len() <= 8000, "answer grew: {s}");
assert!(
node_count(r) <= 1200,
"answer node count grew: {s} ({})",
node_count(r)
);
assert!(
verify_candidate(&ctx, expr, r, var),
"unverified answer emitted for {s}"
);
}
if round == 0 {
baseline[index] = solved;
} else {
assert_eq!(baseline[index], solved, "unstable outcome for {s}");
}
}
}
assert!(baseline.iter().any(|s| *s));
assert!(baseline.iter().any(|s| !*s));
}
}