use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use super::is_constant;
pub(crate) fn integrate_half_power<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if super::node_count(expr) > 200 {
return None;
}
let expr = ocas_atom::normalize::normalize(ctx, expr);
let (coeff, core) = split_constant(ctx, expr, var)?;
let base_atom = ctx.var(var.as_str());
let cos_atom = ctx.fun("cos", &[base_atom]);
let (base, p) = match core.node() {
AtomNode::Fun(name, args) if name.as_str() == "sqrt" && args.len() == 1 => (args[0], 1i64),
AtomNode::Pow(b, e) => {
if let AtomNode::Fun(fname, fargs) = b.node()
&& fname.as_str() == "sqrt"
&& fargs.len() == 1
{
let (pe, qe) = exp_fraction(*e)?;
if qe != 1 {
return None;
}
(fargs[0], pe)
} else {
let (pe, qe) = exp_fraction(*e)?;
if qe != 2 {
return None;
}
(*b, pe)
}
}
_ => return None,
};
if p % 2 == 0 || p.unsigned_abs() > 5 {
return None;
}
let s = poly_in_base(ctx, base, cos_atom, var, 2)?;
let c0 = s.first().copied().unwrap_or_else(|| ctx.num(0));
let c1 = s.get(1).copied().unwrap_or_else(|| ctx.num(0));
let c2 = s.get(2).copied().unwrap_or_else(|| ctx.num(0));
let zero = |a: Atom<'a>| is_zero(ctx, a);
let (a_coef, beta, scale, z_repl, sign_branch) = if zero(c2) {
if zero(c1) {
return None;
}
let a = cz(ctx, ctx.add(&[c0, c1]));
if zero(a) {
return None;
}
let beta = cz(ctx, ctx.mul(&[ctx.num(2), c1]));
let half_var = ctx.mul(&[base_atom, ctx.pow(ctx.num(2), ctx.num(-1))]);
(a, beta, ctx.num(2), ctx.fun("sin", &[half_var]), false)
} else {
if !zero(c1) {
return None;
}
let a = cz(ctx, ctx.add(&[c0, c2]));
if zero(a) {
return None;
}
(a, c2, ctx.num(1), ctx.fun("sin", &[base_atom]), true)
};
let z = fresh_symbol(expr, var)?;
let zvar = ctx.var(z.as_str());
let radicand = ctx.add(&[
a_coef,
ctx.mul(&[ctx.num(-1), beta, ctx.pow(zvar, ctx.num(2))]),
]);
let one_minus_z2 = ctx.add(&[
ctx.num(1),
ctx.mul(&[ctx.num(-1), ctx.pow(zvar, ctx.num(2))]),
]);
let rewritten = ctx.mul(&[
coeff,
scale,
ctx.pow(radicand, half_exp(ctx, p)),
ctx.pow(one_minus_z2, half_exp(ctx, -1)),
]);
let g = super::elliptic::integrate_elliptic(ctx, rewritten, z)?;
let back = super::replace_symbol(ctx, g, z, z_repl);
let back = if sign_branch {
let one_minus_sin2 = ctx.add(&[
ctx.num(1),
ctx.mul(&[ctx.num(-1), ctx.pow(z_repl, ctx.num(2))]),
]);
ctx.mul(&[
ctx.fun("cos", &[base_atom]),
ctx.pow(one_minus_sin2, half_exp(ctx, -1)),
back,
])
} else {
back
};
Some(ocas_atom::normalize::normalize(ctx, back))
}
fn half<'a>(ctx: &'a AtomArena<'a>) -> Atom<'a> {
ctx.pow(ctx.num(2), ctx.num(-1))
}
fn half_exp<'a>(ctx: &'a AtomArena<'a>, p: i64) -> Atom<'a> {
ctx.mul(&[ctx.num(p), half(ctx)])
}
fn cz<'a>(ctx: &'a AtomArena<'a>, a: Atom<'a>) -> Atom<'a> {
crate::ode::util::collect_terms(ctx, a)
}
fn is_zero<'a>(ctx: &'a AtomArena<'a>, a: Atom<'a>) -> bool {
matches!(cz(ctx, a).node(), AtomNode::Num(0))
}
fn split_constant<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>)> {
let factors: Vec<Atom<'a>> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![expr],
};
let (c, nc): (Vec<Atom<'a>>, Vec<Atom<'a>>) =
factors.into_iter().partition(|f| is_constant(*f, var));
if nc.is_empty() {
return None;
}
let core = if nc.len() == 1 { nc[0] } else { ctx.mul(&nc) };
let coeff = match c.len() {
0 => ctx.num(1),
1 => c[0],
_ => ctx.mul(&c),
};
Some((coeff, core))
}
fn exp_fraction<'a>(exp: Atom<'a>) -> Option<(i64, i64)> {
match exp.node() {
AtomNode::Num(n) => Some((*n, 1)),
AtomNode::Pow(b, e) => {
if let (AtomNode::Num(bb), AtomNode::Num(ee)) = (b.node(), e.node())
&& *ee == -1
&& *bb > 0
{
return Some((1, *bb));
}
None
}
AtomNode::Mul(args) => {
let mut num: Option<i64> = None;
let mut den: Option<i64> = None;
for a in args.iter() {
match a.node() {
AtomNode::Num(n) => {
if num.is_some() {
return None;
}
num = Some(*n);
}
AtomNode::Pow(b, e) => {
if let (AtomNode::Num(bb), AtomNode::Num(ee)) = (b.node(), e.node())
&& *ee == -1
&& *bb > 0
{
if den.is_some() {
return None;
}
den = Some(*bb);
} else {
return None;
}
}
_ => return None,
}
}
match (num, den) {
(Some(p), Some(q)) => Some((p, q)),
(Some(p), None) => Some((p, 1)),
(None, Some(q)) => Some((1, q)),
(None, None) => None,
}
}
_ => None,
}
}
fn contains_atom<'a>(expr: Atom<'a>, needle: Atom<'a>) -> bool {
if expr == needle {
return true;
}
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => false,
AtomNode::Pow(b, e) => contains_atom(*b, needle) || contains_atom(*e, needle),
AtomNode::Add(args) | AtomNode::Mul(args) | AtomNode::Fun(_, args) => {
args.iter().any(|a| contains_atom(*a, needle))
}
}
}
fn fresh_symbol<'a>(expr: Atom<'a>, var: Symbol) -> Option<Symbol> {
for name in ["_z", "_s", "_w", "_v"] {
let s = Symbol::new(name);
if s != var && !super::contains_symbol(expr, s) {
return Some(s);
}
}
None
}
fn poly_in_base<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
base: Atom<'a>,
var: Symbol,
max_deg: usize,
) -> Option<Vec<Atom<'a>>> {
let expanded = cz(ctx, expr);
let terms: Vec<Atom<'a>> = match expanded.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![expanded],
};
let mut out: Vec<Atom<'a>> = Vec::new();
for t in terms {
let (deg, coeff) = monomial(ctx, t, base, max_deg)?;
if !is_constant(coeff, var) {
return None;
}
if out.len() <= deg {
out.resize(deg + 1, ctx.num(0));
}
out[deg] = cz(ctx, ctx.add(&[out[deg], coeff]));
}
while let Some(last) = out.last().copied() {
if is_zero(ctx, last) {
out.pop();
} else {
break;
}
}
Some(out)
}
fn monomial<'a>(
ctx: &'a AtomArena<'a>,
t: Atom<'a>,
base: Atom<'a>,
max_deg: usize,
) -> Option<(usize, Atom<'a>)> {
if t == base {
return Some((1, ctx.num(1)));
}
match t.node() {
AtomNode::Num(_) => Some((0, t)),
AtomNode::Var(_) => {
if contains_atom(t, base) {
None
} else {
Some((0, t))
}
}
AtomNode::Fun(_, _) => {
if contains_atom(t, base) {
None
} else {
Some((0, t))
}
}
AtomNode::Add(_) => None,
AtomNode::Pow(b, e) => {
if *b == base {
if let AtomNode::Num(n) = e.node() {
let n = usize::try_from(*n).ok()?;
if n <= max_deg {
return Some((n, ctx.num(1)));
}
}
return None;
}
if contains_atom(t, base) {
None
} else {
Some((0, t))
}
}
AtomNode::Mul(args) => {
let mut deg = 0usize;
let mut coeffs: Vec<Atom<'a>> = Vec::with_capacity(args.len());
for a in args.iter() {
let (d, c) = monomial(ctx, *a, base, max_deg)?;
deg = deg.checked_add(d)?;
if deg > max_deg {
return None;
}
coeffs.push(c);
}
Some((deg, ctx.mul(&coeffs)))
}
}
}
#[cfg(test)]
mod tests {
use super::super::elliptic::testnum::{diff_local, eval_f64};
use super::*;
use ocas_core::arena::Arena;
fn parse<'a>(ctx: &'a AtomArena<'a>, s: &str) -> Atom<'a> {
ocas_parse::parse(ctx, s).expect("parse")
}
fn assert_antiderivative_num<'a>(
ctx: &'a AtomArena<'a>,
src: &str,
consts: &[(Symbol, f64)],
samples: &[f64],
) {
let var = Symbol::new("x");
let integrand = parse(ctx, src);
let result =
integrate_half_power(ctx, integrand, var).unwrap_or_else(|| panic!("declined: {src}"));
assert!(
!result.to_string().contains("Integral"),
"residue for {src}: {result}"
);
let d = diff_local(ctx, result, var);
let mut checked = 0usize;
for &xv in samples {
let mut env = consts.to_vec();
env.push((var, xv));
let lhs = match eval_f64(d, &env) {
Some(v) => v,
None => continue,
};
let rhs = eval_f64(integrand, &env).expect("eval integrand");
let tol = 1e-9 * rhs.abs().max(1.0);
assert!(
(lhs - rhs).abs() < tol,
"at x={xv}: diff={lhs} integrand={rhs} (src: {src}, result: {result})"
);
checked += 1;
}
assert!(checked >= 2, "only {checked} usable samples for {src}");
}
fn declines<'a>(ctx: &'a AtomArena<'a>, src: &str) {
let e = parse(ctx, src);
let var = Symbol::new("x");
assert!(
integrate_half_power(ctx, e, var).is_none(),
"expected decline for {src}"
);
}
#[test]
fn linear_in_cos_first_kind() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 2.0), (Symbol::new("b"), 1.0)];
assert_antiderivative_num(
&ctx,
"1/sqrt(a+b*cos(x))",
&env,
&[-2.0, -1.0, 0.4, 1.2, 2.4],
);
}
#[test]
fn linear_in_cos_second_kind() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 2.0), (Symbol::new("b"), 1.0)];
assert_antiderivative_num(&ctx, "sqrt(a+b*cos(x))", &env, &[-2.0, -0.8, 0.5, 1.5, 2.5]);
}
#[test]
fn linear_in_cos_hermite() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 2.0), (Symbol::new("b"), 1.0)];
assert_antiderivative_num(
&ctx,
"1/(a+b*cos(x))^(3/2)",
&env,
&[-2.2, -0.6, 0.3, 1.4, 2.6],
);
}
#[test]
fn quadratic_in_cos_first_kind() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 2.0), (Symbol::new("b"), 1.0)];
assert_antiderivative_num(
&ctx,
"1/sqrt(a+b*cos(x)^2)",
&env,
&[-2.4, -1.2, -0.4, 0.6, 1.9, 2.6],
);
}
#[test]
fn quadratic_in_cos_second_kind() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 2.0), (Symbol::new("b"), 1.0)];
assert_antiderivative_num(
&ctx,
"sqrt(a+b*cos(x)^2)",
&env,
&[-2.4, -1.1, -0.3, 0.7, 1.8, 2.7],
);
}
#[test]
fn quadratic_in_cos_hermite() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let env = [(Symbol::new("a"), 2.0), (Symbol::new("b"), 1.0)];
assert_antiderivative_num(
&ctx,
"1/(a+b*cos(x)^2)^(3/2)",
&env,
&[-2.4, -1.3, -0.3, 0.8, 1.7, 2.8],
);
}
#[test]
fn numeric_coefficients_and_extra_powers() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
assert_antiderivative_num(&ctx, "1/sqrt(2+cos(x))", &[], &[-2.0, -0.5, 0.6, 2.0]);
assert_antiderivative_num(&ctx, "1/(2+cos(x)^2)^(3/2)", &[], &[-2.2, -0.7, 0.5, 2.2]);
assert_antiderivative_num(&ctx, "3/sqrt(2+cos(x))", &[], &[-1.8, -0.4, 0.9, 2.1]);
let env = [(Symbol::new("a"), 3.0), (Symbol::new("b"), 1.0)];
assert_antiderivative_num(&ctx, "sqrt(a+b*cos(x)^2)", &env, &[-2.1, -0.8, 0.6, 2.2]);
}
#[test]
fn harness_parity_for_cos_family() {
const TABLE: [f64; 12] = [
2.0,
3.0,
5.0,
7.0,
0.5,
1.5,
0.25,
11.0,
1.0 / 3.0,
4.0,
6.0,
0.75,
];
const SAMPLES: [f64; 8] = [-1.7, -0.9, -0.37, 0.31, 0.77, 1.23, 1.91, 2.53];
fn param_value(name: &str) -> f64 {
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for b in name.bytes() {
h ^= u64::from(b);
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
let v = TABLE[(h % 12) as usize];
if (h >> 8) & 1 == 0 { v } else { -v }
}
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
let cases = [
"1/(a + b*cos(x)^2)^(3/2)",
"1/sqrt(a + b*cos(x))",
"1/sqrt(a + b*cos(x)^2)",
"sqrt(a + b*cos(x))",
"sqrt(a + b*cos(x)^2)",
];
for mode in ["harness dummy parameters", "real regime a=2 b=1"] {
for src in cases {
let env = if mode.starts_with("harness") {
[
(Symbol::new("a"), param_value("a")),
(Symbol::new("b"), param_value("b")),
]
} else {
[(Symbol::new("a"), 2.0), (Symbol::new("b"), 1.0)]
};
let integrand = parse(&ctx, src);
let result = crate::integrate(&ctx, integrand, var);
let text = result.to_string();
assert!(!text.contains("Integral("), "{src} fell back: {text}");
let (mut checked, mut worst) = (0usize, 0.0f64);
for &x in SAMPLES.iter() {
let h = 1e-4 * x.abs().max(1.0);
let eval_at = |t: f64| {
let mut e = env.to_vec();
e.push((var, t));
eval_f64(result, &e)
};
let (fm2, fm1, fp1, fp2) = (
eval_at(x - 2.0 * h),
eval_at(x - h),
eval_at(x + h),
eval_at(x + 2.0 * h),
);
let mut e = env.to_vec();
e.push((var, x));
let rhs = eval_f64(integrand, &e);
if let (Some(a), Some(b), Some(c), Some(d), Some(f)) = (fm2, fm1, fp1, fp2, rhs)
{
let deriv = (-d + 8.0 * c - 8.0 * b + a) / (12.0 * h);
if deriv.is_finite() {
checked += 1;
worst = worst.max((deriv - f).abs() / f.abs().max(1.0));
}
}
}
if mode.starts_with("harness") {
assert!(worst <= 1e-5, "{src} [{mode}]: worst rel {worst:e}");
} else {
assert!(
checked >= 2,
"{src} [{mode}]: only {checked} usable samples"
);
assert!(
worst <= 1e-5,
"{src} [{mode}]: worst rel {worst:e} ({text})"
);
println!("{src}: checked={checked} worst_rel={worst:e}");
}
}
}
}
#[test]
fn declines_unsupported_shapes() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
declines(&ctx, "sqrt(sin(x))");
declines(&ctx, "1/sqrt(a+b*sin(x))");
declines(&ctx, "1/sqrt(a+b*cos(x)+c*cos(x)^2)");
declines(&ctx, "1/(a+b*sec(x))^(5/2)");
declines(&ctx, "sqrt(tan(x))");
declines(&ctx, "sqrt(a+b*sinh(x))");
declines(&ctx, "exp(x)*sqrt(cos(x))");
declines(&ctx, "sin(x)*sqrt(a+b*cos(x))");
declines(&ctx, "x/sqrt(a+b*cos(x))");
declines(&ctx, "1/sqrt(x^5+1)");
declines(&ctx, "1/(1+cos(x))");
declines(&ctx, "cos(x)");
declines(&ctx, "sqrt(1+x^2)");
declines(&ctx, "1/(a+b*cos(x))^(7/2)");
declines(&ctx, "(a+b*cos(x))^(3/2)");
declines(&ctx, "(a+b*cos(x)^2)^(5/2)");
declines(&ctx, "1/sqrt(1-cos(x))");
declines(&ctx, "1/sqrt(a+b*cos(2*x))");
}
#[test]
fn stress_stable_outcomes() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let var = Symbol::new("x");
let solid = parse(&ctx, "1/sqrt(a+b*cos(x))");
let mut first: Option<String> = None;
for i in 0..300 {
let r = integrate_half_power(&ctx, solid, var);
let s = r.map(|a| a.to_string());
if i == 0 {
first = s.clone();
}
assert_eq!(s, first, "non-deterministic outcome at iteration {i}");
}
assert!(first.is_some(), "stable outcome must be a solve");
let none = parse(&ctx, "sqrt(sin(x))");
for i in 0..300 {
assert!(
integrate_half_power(&ctx, none, var).is_none(),
"decline not stable at iteration {i}"
);
}
}
}