use ocas_atom::Symbol;
use ocas_atom::normalize::normalize;
use ocas_atom::{Atom, AtomArena, AtomNode};
use super::linear_form;
const MAX_TRIG_FACTORS: usize = 8;
const MAX_OUTPUT_TERMS: usize = 64;
#[derive(Clone, Copy)]
struct TrigLin<'a> {
sin: bool,
a: Atom<'a>,
b: Atom<'a>,
}
pub(crate) fn trig_reduce_products<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let (rest, trig) = extract_trig_factors(ctx, expr, var)?;
if trig.len() < 2 || trig.len() > MAX_TRIG_FACTORS {
return None;
}
let mut leaves: Vec<(i64, u32, Option<TrigLin<'a>>)> = Vec::new();
reduce(ctx, &trig, 1, 0, &mut leaves)?;
if leaves.len() > MAX_OUTPUT_TERMS {
return None;
}
let mut terms = Vec::with_capacity(leaves.len());
for (sign, two_pow, factor) in leaves {
let mut parts: Vec<Atom<'a>> = rest.clone();
if sign < 0 {
parts.push(ctx.num(-1));
}
if two_pow > 0 {
parts.push(ctx.pow(ctx.num(2), ctx.num(-(two_pow as i64))));
}
if let Some(f) = factor {
parts.push(build_trig(ctx, f, var));
}
terms.push(ctx.mul(&parts));
}
let sum = ctx.add(&terms);
if sum == expr { None } else { Some(sum) }
}
fn classify_factor<'a>(
ctx: &'a AtomArena<'a>,
f: Atom<'a>,
var: Symbol,
rest: &mut Vec<Atom<'a>>,
trig: &mut Vec<TrigLin<'a>>,
) -> Option<()> {
let (name, arg, copies) = match f.node() {
AtomNode::Fun(name, args) if args.len() == 1 => (name, args[0], 1usize),
AtomNode::Pow(b, e) => match (b.node(), e.node()) {
(AtomNode::Fun(name, args), AtomNode::Num(k))
if args.len() == 1 && (1..=8i64).contains(k) =>
{
(name, args[0], *k as usize)
}
_ => {
rest.push(f);
return Some(());
}
},
_ => {
rest.push(f);
return Some(());
}
};
if name.as_str() != "sin" && name.as_str() != "cos" {
rest.push(f);
return Some(());
}
let (a, b) = linear_form(ctx, arg, var)?;
for _ in 0..copies {
trig.push(TrigLin {
sin: name.as_str() == "sin",
a,
b,
});
}
Some(())
}
fn extract_trig_factors<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<(Vec<Atom<'a>>, Vec<TrigLin<'a>>)> {
let mut rest = Vec::new();
let mut trig = Vec::new();
match expr.node() {
AtomNode::Mul(args) => {
for a in args.iter() {
classify_factor(ctx, *a, var, &mut rest, &mut trig)?;
}
}
_ => classify_factor(ctx, expr, var, &mut rest, &mut trig)?,
}
if trig.is_empty() {
return None;
}
Some((rest, trig))
}
fn reduce<'a>(
ctx: &'a AtomArena<'a>,
fs: &[TrigLin<'a>],
sign: i64,
two_pow: u32,
out: &mut Vec<(i64, u32, Option<TrigLin<'a>>)>,
) -> Option<()> {
if out.len() > MAX_OUTPUT_TERMS {
return None;
}
if fs.len() <= 1 {
out.push((sign, two_pow, fs.first().copied()));
return Some(());
}
let (f0, f1) = (fs[0], fs[1]);
let rest = &fs[2..];
let (u, v) = if !f0.sin && f1.sin {
(f1, f0)
} else {
(f0, f1)
};
let branches: [(TrigLin<'a>, i64); 2] = if u.sin && v.sin {
[
(
TrigLin {
sin: false,
a: sub(ctx, u.a, v.a),
b: sub(ctx, u.b, v.b),
},
1,
),
(
TrigLin {
sin: false,
a: add(ctx, u.a, v.a),
b: add(ctx, u.b, v.b),
},
-1,
),
]
} else if !u.sin && !v.sin {
[
(
TrigLin {
sin: false,
a: sub(ctx, u.a, v.a),
b: sub(ctx, u.b, v.b),
},
1,
),
(
TrigLin {
sin: false,
a: add(ctx, u.a, v.a),
b: add(ctx, u.b, v.b),
},
1,
),
]
} else {
[
(
TrigLin {
sin: true,
a: add(ctx, u.a, v.a),
b: add(ctx, u.b, v.b),
},
1,
),
(
TrigLin {
sin: true,
a: sub(ctx, u.a, v.a),
b: sub(ctx, u.b, v.b),
},
1,
),
]
};
for (g, branch_sign) in branches {
let zero_angle = is_zero(g.a) && is_zero(g.b);
if zero_angle && g.sin {
continue;
}
let mut next = Vec::with_capacity(rest.len() + 1);
next.extend_from_slice(rest);
if !zero_angle {
next.push(g);
}
reduce(ctx, &next, sign * branch_sign, two_pow + 1, out)?;
}
Some(())
}
fn build_trig<'a>(ctx: &'a AtomArena<'a>, f: TrigLin<'a>, var: Symbol) -> Atom<'a> {
let arg = build_linear(ctx, f.a, f.b, var);
if let AtomNode::Num(n) = f.a.node()
&& *n < 0
{
let pos_arg = build_linear(ctx, ctx.num(-n), ctx.mul(&[ctx.num(-1), f.b]), var);
let fun = ctx.fun(if f.sin { "sin" } else { "cos" }, &[pos_arg]);
return if f.sin {
ctx.mul(&[ctx.num(-1), fun])
} else {
fun
};
}
ctx.fun(if f.sin { "sin" } else { "cos" }, &[arg])
}
fn build_linear<'a>(ctx: &'a AtomArena<'a>, a: Atom<'a>, b: Atom<'a>, var: Symbol) -> Atom<'a> {
let x = ctx.var(var.as_str());
let a = fold(ctx, a);
let b = fold(ctx, b);
let ax = if is_one(a) {
x
} else if is_zero(a) {
ctx.num(0)
} else {
ctx.mul(&[a, x])
};
let sum = if is_zero(b) { ax } else { ctx.add(&[b, ax]) };
fold(ctx, sum)
}
fn add<'a>(ctx: &'a AtomArena<'a>, u: Atom<'a>, v: Atom<'a>) -> Atom<'a> {
fold(ctx, ctx.add(&[u, v]))
}
fn sub<'a>(ctx: &'a AtomArena<'a>, u: Atom<'a>, v: Atom<'a>) -> Atom<'a> {
if u == v {
return ctx.num(0);
}
let folded = fold(ctx, ctx.add(&[u, ctx.mul(&[ctx.num(-1), v])]));
if is_zero(folded) { ctx.num(0) } else { folded }
}
fn fold<'a>(ctx: &'a AtomArena<'a>, e: Atom<'a>) -> Atom<'a> {
normalize(ctx, crate::ode::util::collect_terms(ctx, e))
}
fn is_zero(e: Atom<'_>) -> bool {
matches!(e.node(), AtomNode::Num(0))
}
fn is_one(e: Atom<'_>) -> bool {
matches!(e.node(), AtomNode::Num(1))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::integral::integrate;
use ocas_core::arena::Arena;
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(),
"sec" => v.cos().recip(),
"csc" => v.sin().recip(),
"cot" => v.tan().recip(),
"exp" => v.exp(),
"log" => v.ln(),
"sqrt" => v.sqrt(),
_ => return None,
})
}
}
}
fn assert_antiderivative_num<'a>(
ctx: &'a AtomArena<'a>,
integrand: Atom<'a>,
var: Symbol,
consts: &[(Symbol, f64)],
samples: &[f64],
) {
let result = integrate(ctx, integrand, var);
assert!(
!result.to_string().contains("Integral"),
"fallback: {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})"
);
}
}
#[test]
fn sin_times_cos() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.fun("sin", &[x]), ctx.fun("cos", &[x])]);
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &[], &[0.3, 0.7, 1.1]);
}
#[test]
fn cos_squared() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.pow(ctx.fun("cos", &[x]), ctx.num(2));
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &[], &[0.3, 0.7, 1.1]);
}
#[test]
fn sin_squared_times_cos() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let s2 = ctx.pow(ctx.fun("sin", &[x]), ctx.num(2));
let expr = ctx.mul(&[s2, ctx.fun("cos", &[x])]);
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &[], &[0.3, 0.7, 1.1]);
}
#[test]
fn poly_times_trig_product_symbolic_linear_arg() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let a = ctx.var("a");
let b = ctx.var("b");
let u = ctx.add(&[a, ctx.mul(&[b, x])]);
let s2 = ctx.pow(ctx.fun("sin", &[u]), ctx.num(2));
let expr = ctx.mul(&[ctx.fun("cos", &[u]), s2]);
let env = [(Symbol::new("a"), 0.4), (Symbol::new("b"), 1.3)];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.3, 0.7, 1.1]);
}
#[test]
fn poly_factor_times_reduced_trig() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let poly = ctx.add(&[ctx.mul(&[ctx.num(2), x]), ctx.num(1)]);
let s2 = ctx.pow(ctx.fun("sin", &[x]), ctx.num(2));
let expr = ctx.mul(&[poly, ctx.fun("cos", &[x]), s2]);
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &[], &[0.3, 0.7, 1.1]);
}
#[test]
fn different_linear_args_product_to_sum() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[
ctx.fun("sin", &[x]),
ctx.fun("cos", &[ctx.mul(&[ctx.num(2), x])]),
]);
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &[], &[0.3, 0.7, 1.1]);
}
#[test]
fn declines_nonlinear_argument() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let x2 = ctx.pow(x, ctx.num(2));
let expr = ctx.mul(&[ctx.fun("sin", &[x2]), ctx.fun("cos", &[x2])]);
assert!(trig_reduce_products(&ctx, expr, Symbol::new("x")).is_none());
}
#[test]
fn resonant_equal_slopes_stay_finite_and_correct() {
let inputs = [
"sin(c + d*x)*cos(c + d*x)",
"cos(c + d*x)^2",
"sin(c + d*x)^2",
"sin(c + d*x)^2*(a + b*sin(c + d*x)^2)",
"cos(c + d*x)^3*(a + a*sin(c + d*x))",
"(a*cos(c + d*x) + b*sin(c + d*x))^2",
];
for input in inputs {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, input).expect("parse");
let result = integrate(&ctx, expr, Symbol::new("x"));
let text = result.to_string();
assert!(!text.contains("Integral"), "{input}: fallback {text}");
assert!(
!text.contains("(-1*d)") || !text.contains("d + (-1*d)"),
"{input}: resonant denominator left behind: {text}"
);
let env = [
(Symbol::new("a"), 1.5),
(Symbol::new("b"), 2.0),
(Symbol::new("c"), 0.5),
(Symbol::new("d"), 1.5),
];
assert_antiderivative_num(&ctx, expr, Symbol::new("x"), &env, &[0.3, 0.7, 1.9]);
}
}
}