use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use ocas_rewrite::matcher::{Bindings, MatchValue, match_pattern};
use ocas_rewrite::pattern::Pattern;
use ocas_rewrite::rules::{HeadKey, Rule, RuleTable, head_of};
pub(crate) const MAX_RULE_DEPTH: usize = 64;
pub(crate) type Pred = for<'a> fn(&Bindings<'a>, Symbol) -> bool;
#[derive(Clone, Copy)]
pub(crate) enum RuleSpec {
Template {
pat: &'static str,
tmpl: &'static str,
cond: Option<Pred>,
},
Closure {
pat: &'static str,
f: for<'a> fn(&'a AtomArena<'a>, &Bindings<'a>, Symbol) -> Option<Atom<'a>>,
cond: Option<Pred>,
},
}
struct ClosureRule<'a> {
head: HeadKey,
pattern: Pattern<'a>,
cond: Option<Pred>,
f: for<'b> fn(&'b AtomArena<'b>, &Bindings<'b>, Symbol) -> Option<Atom<'b>>,
}
pub(crate) struct IntegralRuleTable<'a> {
table: RuleTable<'a>,
closures: Vec<ClosureRule<'a>>,
}
fn sqrt_quadratic_base<'a>(expr: Atom<'a>) -> Option<Atom<'a>> {
match expr.node() {
AtomNode::Fun(name, args) if name.as_str() == "sqrt" && args.len() == 1 => Some(args[0]),
AtomNode::Pow(base, exp) => match (base.node(), exp.node()) {
(AtomNode::Fun(name, args), AtomNode::Num(-1))
if name.as_str() == "sqrt" && args.len() == 1 =>
{
Some(args[0])
}
(_, AtomNode::Pow(e2, e3)) => {
if matches!(e2.node(), AtomNode::Num(2)) && matches!(e3.node(), AtomNode::Num(-1)) {
Some(*base)
} else {
None
}
}
(_, AtomNode::Mul(args)) => {
let has_neg_half = args.iter().any(|a| matches!(a.node(), AtomNode::Num(-1)))
&& args.iter().any(|a| {
matches!(a.node(), AtomNode::Pow(b, e)
if matches!(b.node(), AtomNode::Num(2))
&& matches!(e.node(), AtomNode::Num(-1)))
});
if has_neg_half && args.len() == 2 {
Some(*base)
} else {
None
}
}
_ => None,
},
_ => None,
}
}
pub(crate) fn quadratic_coeffs<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>, Atom<'a>)> {
if let Some(inner) = sqrt_quadratic_base(expr) {
if let Some(coeffs) = quadratic_from_inner(ctx, inner, var) {
return Some(coeffs);
}
}
match expr.node() {
AtomNode::Add(args) | AtomNode::Mul(args) | AtomNode::Fun(_, args) => {
for a in args.iter() {
if let Some(coeffs) = quadratic_coeffs(ctx, *a, var) {
return Some(coeffs);
}
}
}
AtomNode::Pow(base, exp) => {
if let Some(coeffs) = quadratic_coeffs(ctx, *base, var) {
return Some(coeffs);
}
if let Some(coeffs) = quadratic_coeffs(ctx, *exp, var) {
return Some(coeffs);
}
}
AtomNode::Num(_) | AtomNode::Var(_) => {}
}
None
}
fn quadratic_from_inner<'a>(
ctx: &'a AtomArena<'a>,
inner: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>, Atom<'a>)> {
let (mut a, mut b, mut c) = (ctx.num(0), ctx.num(0), ctx.num(0));
let mut found = false;
match inner.node() {
AtomNode::Add(args) => {
for term in args.iter() {
match term.node() {
AtomNode::Pow(base, exp) => {
if matches!(base.node(), AtomNode::Var(v) if *v == var)
&& matches!(exp.node(), AtomNode::Num(2))
{
a = ctx.num(1); found = true;
} else if crate::integral::is_constant(*term, var) {
c = *term;
found = true;
} else {
return None;
}
}
AtomNode::Var(v) if *v == var => {
b = ctx.num(1); found = true;
}
AtomNode::Mul(args) if args.len() == 2 => {
let (coeff, rest) = (args[0], args[1]);
match rest.node() {
AtomNode::Pow(base, exp)
if matches!(base.node(), AtomNode::Var(v) if *v == var)
&& matches!(exp.node(), AtomNode::Num(2)) =>
{
if !crate::integral::is_constant(coeff, var) {
return None;
}
a = coeff;
found = true;
}
AtomNode::Var(v) if *v == var => {
if !crate::integral::is_constant(coeff, var) {
return None;
}
b = coeff;
found = true;
}
_ => {
if crate::integral::is_constant(*term, var) {
c = *term;
found = true;
} else {
return None;
}
}
}
}
AtomNode::Num(_) => {
c = *term;
found = true;
}
_ => {
if crate::integral::is_constant(*term, var) {
c = *term;
found = true;
} else {
return None;
}
}
}
}
}
AtomNode::Pow(base, exp) => {
if matches!(base.node(), AtomNode::Var(v) if *v == var)
&& matches!(exp.node(), AtomNode::Num(2))
{
a = ctx.num(1);
found = true;
}
}
AtomNode::Var(v) if *v == var => {
b = ctx.num(1);
found = true;
}
_ => {
if crate::integral::is_constant(inner, var) {
c = inner;
found = true;
}
}
}
if !found {
return None;
}
Some((a, b, c))
}
pub(crate) fn rat_of(atom: Atom<'_>) -> Option<(i64, i64)> {
match atom.node() {
AtomNode::Num(n) => Some((*n, 1)),
_ => crate::integral::fraction_exponent(atom),
}
}
pub(crate) fn rational_sqrt(p: i64, q: i64) -> Option<(i64, i64)> {
if q <= 0 || p < 0 {
return None;
}
let g = gcd_i64(p, q);
let (p, q) = (p / g, q / g);
let r = isqrt(p)?;
let s = isqrt(q)?;
Some((r, s))
}
fn gcd_i64(mut a: i64, mut b: i64) -> i64 {
while b != 0 {
let t = a % b;
a = b;
b = t;
}
a.abs().max(1)
}
fn isqrt(n: i64) -> Option<i64> {
if n < 0 {
return None;
}
let r = (n as f64).sqrt() as i64;
[r - 1, r, r + 1]
.into_iter()
.find(|&c| c >= 0 && c * c == n)
}
pub(crate) mod pred {
use super::*;
pub(crate) fn bound<'a>(bindings: &Bindings<'a>, name: &str) -> Option<Atom<'a>> {
match bindings.get(Symbol::new(name))? {
MatchValue::Single(a) => Some(*a),
MatchValue::Sequence(_) => None,
}
}
pub(crate) fn free_q(bindings: &Bindings<'_>, var: Symbol) -> bool {
["a", "b", "c", "d", "m", "n", "p"]
.iter()
.all(|n| match bound(bindings, n) {
Some(a) => crate::integral::is_constant(a, var),
None => true,
})
}
pub(crate) fn int_val(bindings: &Bindings<'_>, name: &str) -> Option<i64> {
match bound(bindings, name) {
Some(a) => match a.node() {
AtomNode::Num(n) => Some(*n),
_ => None,
},
None => None,
}
}
pub(crate) fn pos_int_q(bindings: &Bindings<'_>, name: &str) -> bool {
int_val(bindings, name).is_some_and(|n| n >= 0)
}
pub(crate) fn int_ge_2(bindings: &Bindings<'_>, name: &str) -> bool {
int_val(bindings, name).is_some_and(|n| n >= 2)
}
pub(crate) fn not_minus_one(bindings: &Bindings<'_>, name: &str) -> bool {
int_val(bindings, name).is_none_or(|n| n != -1)
}
pub(crate) fn nonzero(bindings: &Bindings<'_>, name: &str) -> bool {
!matches!(bound(bindings, name), Some(a) if matches!(a.node(), AtomNode::Num(0)))
}
pub(crate) fn int_in_range(bindings: &Bindings<'_>, name: &str, lo: i64, hi: i64) -> bool {
int_val(bindings, name).is_some_and(|n| (lo..=hi).contains(&n))
}
pub(crate) fn odd_le_9(bindings: &Bindings<'_>, name: &str) -> bool {
int_val(bindings, name).is_some_and(|n| (1..=9).contains(&n) && n % 2 == 1)
}
fn is_minus_one(a: Atom<'_>) -> bool {
matches!(a.node(), AtomNode::Num(n) if *n == -1)
}
fn same_or_opposite<'a>(a: Atom<'a>, b: Atom<'a>) -> bool {
if a == b {
return true;
}
let (xs, ys) = match (a.node(), b.node()) {
(AtomNode::Num(m), AtomNode::Num(n)) => return *m == -*n,
(AtomNode::Mul(xs), AtomNode::Mul(ys)) => (xs, ys),
(AtomNode::Mul(xs), _) => {
return matches!(xs, [f, g] if is_minus_one(*f) && *g == b);
}
(_, AtomNode::Mul(ys)) => {
return matches!(ys, [f, g] if is_minus_one(*f) && *g == a);
}
_ => return false,
};
if xs.len() != ys.len() {
return false;
}
let mut flips = 0usize;
for (x, y) in xs.iter().zip(ys.iter()) {
if x == y {
continue;
}
match (x.node(), y.node()) {
(AtomNode::Num(m), AtomNode::Num(n)) if *m == -*n => flips += 1,
_ => return false,
}
}
flips == 1
}
pub(crate) fn nonresonant_q(bindings: &Bindings<'_>, s1: &str, s2: &str) -> bool {
match (bound(bindings, s1), bound(bindings, s2)) {
(Some(p), Some(q)) => !same_or_opposite(p, q),
_ => true,
}
}
pub(crate) fn nonresonant_unit_q(bindings: &Bindings<'_>, name: &str) -> bool {
match bound(bindings, name) {
Some(a) => !matches!(a.node(), AtomNode::Num(n) if *n == 1 || *n == -1),
None => true,
}
}
}
fn rule_specs() -> &'static [RuleSpec] {
use pred::*;
&[
RuleSpec::Template {
pat: "x^n_",
tmpl: "x^(n_+1)*(n_+1)^-1",
cond: Some(|b, var| free_q(b, var) && not_minus_one(b, "n")),
},
RuleSpec::Template {
pat: "(a_+b_*x)^-1",
tmpl: "log(a_+b_*x)*b_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "(a_+b_*x)^n_",
tmpl: "(a_+b_*x)^(n_+1)*(b_*(n_+1))^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b") && not_minus_one(b, "n")),
},
RuleSpec::Closure {
pat: "c___*x^m_*(a_+b_*x)^n_",
f: binomial_integrate,
cond: Some(|b, _var| {
int_in_range(b, "m", 0, 8)
&& int_in_range(b, "n", 0, 8)
&& bound(b, "a").is_some()
&& bound(b, "b").is_some()
}),
},
RuleSpec::Template {
pat: "exp(a_*x+b_)",
tmpl: "exp(a_*x+b_)*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "x^n_*exp(a_*x)",
tmpl: "x^n_*exp(a_*x)*a_^-1 + (-1)*n_*a_^-1*Integral(x^(n_ - 1)*exp(a_*x), x)",
cond: Some(|b, var| free_q(b, var) && int_in_range(b, "n", 0, 12) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "x^n_*exp(a_*x+b_)",
tmpl: "x^n_*exp(a_*x+b_)*a_^-1 + (-1)*n_*a_^-1*Integral(x^(n_ - 1)*exp(a_*x+b_), x)",
cond: Some(|b, var| free_q(b, var) && int_in_range(b, "n", 0, 12) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "exp(a_*x)*sin(b_*x)",
tmpl: "exp(a_*x)*(a_*sin(b_*x) + (-1)*b_*cos(b_*x))*((a_^2+b_^2))^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "exp(a_*x)*cos(b_*x)",
tmpl: "exp(a_*x)*(a_*cos(b_*x) + b_*sin(b_*x))*((a_^2+b_^2))^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "exp(x)*sin(b_*x)",
tmpl: "exp(x)*(sin(b_*x) + (-1)*b_*cos(b_*x))*((1+b_^2))^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "exp(x)*cos(b_*x)",
tmpl: "exp(x)*(cos(b_*x) + b_*sin(b_*x))*((1+b_^2))^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "exp(a_*x)*sin(x)",
tmpl: "exp(a_*x)*(a_*sin(x) + (-1)*cos(x))*((a_^2+1))^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "exp(a_*x)*cos(x)",
tmpl: "exp(a_*x)*(a_*cos(x) + sin(x))*((a_^2+1))^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "log(x)^n_",
tmpl: "x*log(x)^n_ + (-1)*n_*Integral(log(x)^(n_ - 1), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "x^m_*log(x)^n_",
tmpl: "x^(m_+1)*log(x)^n_*(m_+1)^-1 + (-1)*n_*(m_+1)^-1*Integral(x^m_*log(x)^(n_ - 1), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n") && not_minus_one(b, "m")),
},
RuleSpec::Template {
pat: "x^-1*log(x)^-1",
tmpl: "log(log(x))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "x*exp(a_*x^2)",
tmpl: "exp(a_*x^2)*(2*a_)^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sin(x)^n_",
tmpl: "(-1)*sin(x)^(n_ - 1)*cos(x)*n_^-1 + (n_ - 1)*n_^-1*Integral(sin(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "cos(x)^n_",
tmpl: "cos(x)^(n_ - 1)*sin(x)*n_^-1 + (n_ - 1)*n_^-1*Integral(cos(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "tan(x)^n_",
tmpl: "tan(x)^(n_ - 1)*(n_ - 1)^-1 + (-1)*Integral(tan(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "tan(x)",
tmpl: "(-1)*log(cos(x))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "cot(x)",
tmpl: "log(sin(x))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "sec(x)^n_",
tmpl: "sec(x)^(n_ - 2)*tan(x)*(n_ - 1)^-1 + (n_ - 2)*(n_ - 1)^-1*Integral(sec(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "sec(x)",
tmpl: "log(sec(x)+tan(x))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "csc(x)",
tmpl: "log(tan(x*2^-1))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "tan(a_*x+b_)",
tmpl: "(-1)*log(cos(a_*x+b_))*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "cot(a_*x+b_)",
tmpl: "log(sin(a_*x+b_))*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sec(a_*x+b_)",
tmpl: "log(sec(a_*x+b_)+tan(a_*x+b_))*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "csc(a_*x+b_)",
tmpl: "log(tan((a_*x+b_)*2^-1))*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sin(a_*x+b_)^n_",
tmpl: "(-1)*sin(a_*x+b_)^(n_ - 1)*cos(a_*x+b_)*(n_*a_)^-1 + (n_ - 1)*(n_)^-1*Integral(sin(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "cos(a_*x+b_)^n_",
tmpl: "cos(a_*x+b_)^(n_ - 1)*sin(a_*x+b_)*(n_*a_)^-1 + (n_ - 1)*(n_)^-1*Integral(cos(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "tan(a_*x+b_)^n_",
tmpl: "tan(a_*x+b_)^(n_ - 1)*((n_ - 1)*a_)^-1 + (-1)*Integral(tan(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "cot(a_*x+b_)^n_",
tmpl: "(-1)*cot(a_*x+b_)^(n_ - 1)*((n_ - 1)*a_)^-1 + (-1)*Integral(cot(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "sec(a_*x+b_)^n_",
tmpl: "sec(a_*x+b_)^(n_ - 2)*tan(a_*x+b_)*((n_ - 1)*a_)^-1 + (n_ - 2)*(n_ - 1)^-1*Integral(sec(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "csc(a_*x+b_)^n_",
tmpl: "(-1)*csc(a_*x+b_)^(n_ - 2)*cot(a_*x+b_)*((n_ - 1)*a_)^-1 + (n_ - 2)*(n_ - 1)^-1*Integral(csc(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "sin(a_*x+b_)*sin(c_*x+d_)",
tmpl: "sin((a_+(-1)*c_)*x+(b_+(-1)*d_))*((a_+(-1)*c_)*2)^-1 + (-1)*sin((a_+c_)*x+(b_+d_))*((a_+c_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_q(b, "a", "c")),
},
RuleSpec::Template {
pat: "cos(a_*x+b_)*cos(c_*x+d_)",
tmpl: "sin((a_+(-1)*c_)*x+(b_+(-1)*d_))*((a_+(-1)*c_)*2)^-1 + sin((a_+c_)*x+(b_+d_))*((a_+c_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_q(b, "a", "c")),
},
RuleSpec::Template {
pat: "sin(a_*x+b_)*cos(a_*x+b_)",
tmpl: "(-1)*cos(a_*x+b_)^2*((2*a_)^-1)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sin(a_*x+b_)*cos(c_*x+d_)",
tmpl: "(-1)*cos((a_+c_)*x+(b_+d_))*((a_+c_)*2)^-1 + (-1)*cos((a_+(-1)*c_)*x+(b_+(-1)*d_))*((a_+(-1)*c_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_q(b, "a", "c")),
},
RuleSpec::Template {
pat: "cot(x)^n_",
tmpl: "(-1)*cot(x)^(n_ - 1)*(n_ - 1)^-1 + (-1)*Integral(cot(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "csc(x)^n_",
tmpl: "(-1)*csc(x)^(n_ - 2)*cot(x)*(n_ - 1)^-1 + (n_ - 2)*(n_ - 1)^-1*Integral(csc(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "sin(a_*x)*sin(b_*x)",
tmpl: "sin((a_+(-1)*b_)*x)*((a_+(-1)*b_)*2)^-1 + (-1)*sin((a_+b_)*x)*((a_+b_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_q(b, "a", "b")),
},
RuleSpec::Template {
pat: "cos(a_*x)*cos(b_*x)",
tmpl: "sin((a_+(-1)*b_)*x)*((a_+(-1)*b_)*2)^-1 + sin((a_+b_)*x)*((a_+b_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_q(b, "a", "b")),
},
RuleSpec::Template {
pat: "sin(a_*x)*cos(b_*x)",
tmpl: "(-1)*cos((a_+b_)*x)*((a_+b_)*2)^-1 + (-1)*cos((a_+(-1)*b_)*x)*((a_+(-1)*b_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_q(b, "a", "b")),
},
RuleSpec::Template {
pat: "sin(x)*sin(b_*x)",
tmpl: "sin((1+(-1)*b_)*x)*((1+(-1)*b_)*2)^-1 + (-1)*sin((1+b_)*x)*((1+b_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_unit_q(b, "b")),
},
RuleSpec::Template {
pat: "sin(a_*x)*sin(x)",
tmpl: "sin((a_+(-1)*1)*x)*((a_+(-1)*1)*2)^-1 + (-1)*sin((a_+1)*x)*((a_+1)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_unit_q(b, "a")),
},
RuleSpec::Template {
pat: "cos(x)*cos(b_*x)",
tmpl: "sin((1+(-1)*b_)*x)*((1+(-1)*b_)*2)^-1 + sin((1+b_)*x)*((1+b_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_unit_q(b, "b")),
},
RuleSpec::Template {
pat: "cos(a_*x)*cos(x)",
tmpl: "sin((a_+(-1)*1)*x)*((a_+(-1)*1)*2)^-1 + sin((a_+1)*x)*((a_+1)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_unit_q(b, "a")),
},
RuleSpec::Template {
pat: "sin(x)*cos(b_*x)",
tmpl: "(-1)*cos((1+b_)*x)*((1+b_)*2)^-1 + (-1)*cos((1+(-1)*b_)*x)*((1+(-1)*b_)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_unit_q(b, "b")),
},
RuleSpec::Template {
pat: "sin(a_*x)*cos(x)",
tmpl: "(-1)*cos((a_+1)*x)*((a_+1)*2)^-1 + (-1)*cos((a_+(-1)*1)*x)*((a_+(-1)*1)*2)^-1",
cond: Some(|b, var| free_q(b, var) && nonresonant_unit_q(b, "a")),
},
RuleSpec::Template {
pat: "sin(x)*cos(x)^n_",
tmpl: "(-1)*cos(x)^(n_+1)*(n_+1)^-1",
cond: Some(|b, var| free_q(b, var) && pos_int_q(b, "n")),
},
RuleSpec::Template {
pat: "sin(x)^m_*cos(x)",
tmpl: "sin(x)^(m_+1)*(m_+1)^-1",
cond: Some(|b, var| free_q(b, var) && pos_int_q(b, "m")),
},
RuleSpec::Closure {
pat: "sin(x)^m_*cos(x)^n_",
f: trig_odd_power,
cond: Some(|b, _var| {
(odd_le_9(b, "m") && int_in_range(b, "n", 2, 9))
|| (int_in_range(b, "m", 2, 9) && odd_le_9(b, "n"))
}),
},
RuleSpec::Template {
pat: "sin(a_*x+b_)*cos(a_*x+b_)^n_",
tmpl: "(-1)*cos(a_*x+b_)^(n_ + 1)*((n_ + 1)*a_)^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && pos_int_q(b, "n")),
},
RuleSpec::Template {
pat: "sin(a_*x+b_)^m_*cos(a_*x+b_)",
tmpl: "sin(a_*x+b_)^(m_ + 1)*((m_ + 1)*a_)^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && pos_int_q(b, "m")),
},
RuleSpec::Closure {
pat: "sin(a_*x+b_)^m_*cos(a_*x+b_)^n_",
f: trig_odd_power_linear,
cond: Some(|b, _var| {
(odd_le_9(b, "m") && int_in_range(b, "n", 2, 9))
|| (int_in_range(b, "m", 2, 9) && odd_le_9(b, "n"))
}),
},
RuleSpec::Template {
pat: "sinh(x)^n_",
tmpl: "sinh(x)^(n_ - 1)*cosh(x)*n_^-1 + (-1)*(n_ - 1)*n_^-1*Integral(sinh(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "cosh(x)^n_",
tmpl: "cosh(x)^(n_ - 1)*sinh(x)*n_^-1 + (n_ - 1)*n_^-1*Integral(cosh(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "tanh(x)^n_",
tmpl: "(-1)*tanh(x)^(n_ - 1)*(n_ - 1)^-1 + Integral(tanh(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "coth(x)^n_",
tmpl: "(-1)*coth(x)^(n_ - 1)*(n_ - 1)^-1 + Integral(coth(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "sech(x)^n_",
tmpl: "sech(x)^(n_ - 2)*tanh(x)*(n_ - 1)^-1 + (n_ - 2)*(n_ - 1)^-1*Integral(sech(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "csch(x)^n_",
tmpl: "(-1)*csch(x)^(n_ - 2)*coth(x)*(n_ - 1)^-1 + (-1)*(n_ - 2)*(n_ - 1)^-1*Integral(csch(x)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "sinh(x)",
tmpl: "cosh(x)",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "cosh(x)",
tmpl: "sinh(x)",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "tanh(x)",
tmpl: "log(cosh(x))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "coth(x)",
tmpl: "log(sinh(x))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "sech(x)",
tmpl: "atan(sinh(x))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "csch(x)",
tmpl: "log(tanh(x*2^-1))",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "sinh(a_*x+b_)^n_",
tmpl: "sinh(a_*x+b_)^(n_ - 1)*cosh(a_*x+b_)*(n_*a_)^-1 + (-1)*(n_ - 1)*(n_)^-1*Integral(sinh(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "cosh(a_*x+b_)^n_",
tmpl: "cosh(a_*x+b_)^(n_ - 1)*sinh(a_*x+b_)*(n_*a_)^-1 + (n_ - 1)*(n_)^-1*Integral(cosh(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "tanh(a_*x+b_)^n_",
tmpl: "(-1)*tanh(a_*x+b_)^(n_ - 1)*((n_ - 1)*a_)^-1 + Integral(tanh(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "coth(a_*x+b_)^n_",
tmpl: "(-1)*coth(a_*x+b_)^(n_ - 1)*((n_ - 1)*a_)^-1 + Integral(coth(a_*x+b_)^(n_ - 2), x)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a") && int_ge_2(b, "n")),
},
RuleSpec::Template {
pat: "sinh(a_*x+b_)",
tmpl: "cosh(a_*x+b_)*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "cosh(a_*x+b_)",
tmpl: "sinh(a_*x+b_)*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "tanh(a_*x+b_)",
tmpl: "log(cosh(a_*x+b_))*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "coth(a_*x+b_)",
tmpl: "log(sinh(a_*x+b_))*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "asin(x)",
tmpl: "x*asin(x) + sqrt(1-x^2)",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "acos(x)",
tmpl: "x*acos(x) + (-1)*sqrt(1-x^2)",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "atan(x)",
tmpl: "x*atan(x) + (-1)*log(x^2+1)*2^-1",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "asinh(x)",
tmpl: "x*asinh(x) + (-1)*sqrt(x^2+1)",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "acosh(x)",
tmpl: "x*acosh(x) + (-1)*sqrt(x^2 - 1)",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "atanh(x)",
tmpl: "x*atanh(x) + log(1-x^2)*2^-1",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "x*asin(x)",
tmpl: "x^2*asin(x)*2^-1 + (-1)*asin(x)*4^-1 + x*(1-x^2)^(2^-1)*4^-1",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "x*atan(x)",
tmpl: "x^2*atan(x)*2^-1 + (-1)*x*2^-1 + atan(x)*2^-1",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "(a_^2+x^2)^-1",
tmpl: "atan(x*a_^-1)*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "(a_^2+(-1)*x^2)^-1",
tmpl: "atanh(x*a_^-1)*a_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "x*(a_+b_*x^2)^-1",
tmpl: "log(a_+b_*x^2)*(2*b_)^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "sqrt(a_+b_*x)",
tmpl: "2*(a_+b_*x)^(3*2^-1)*(3*b_)^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "(a_+b_*x)^(2^-1)",
tmpl: "2*(a_+b_*x)^(3*2^-1)*(3*b_)^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "sqrt(a_+b_*x)^-1",
tmpl: "2*sqrt(a_+b_*x)*b_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "(a_+b_*x)^(-2^-1)",
tmpl: "2*(a_+b_*x)^(2^-1)*b_^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Closure {
pat: "x*sqrt(a_+b_*x)^-1",
f: x_over_sqrt_linear,
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Closure {
pat: "x*(a_+b_*x)^(-2^-1)",
f: x_over_sqrt_linear,
cond: Some(|b, var| free_q(b, var) && nonzero(b, "b")),
},
RuleSpec::Template {
pat: "sqrt(a_^2+(-1)*x^2)",
tmpl: "x*sqrt(a_^2+(-1)*x^2)*2^-1 + a_^2*asin(x*a_^-1)*2^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "(a_^2+(-1)*x^2)^(2^-1)",
tmpl: "x*(a_^2+(-1)*x^2)^(2^-1)*2^-1 + a_^2*asin(x*a_^-1)*2^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "(a_^2+(-1)*x^2)^(-2^-1)",
tmpl: "asin(x*a_^-1)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sqrt(a_^2+(-1)*x^2)^-1",
tmpl: "asin(x*a_^-1)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "(x^2+a_^2)^(-2^-1)",
tmpl: "asinh(x*a_^-1)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sqrt(x^2+a_^2)^-1",
tmpl: "asinh(x*a_^-1)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "(x^2+(-1)*a_^2)^(-2^-1)",
tmpl: "acosh(x*a_^-1)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sqrt(x^2+(-1)*a_^2)^-1",
tmpl: "acosh(x*a_^-1)",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sqrt(x^2+a_^2)",
tmpl: "x*sqrt(x^2+a_^2)*2^-1 + a_^2*asinh(x*a_^-1)*2^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "(x^2+a_^2)^(2^-1)",
tmpl: "x*(x^2+a_^2)^(2^-1)*2^-1 + a_^2*asinh(x*a_^-1)*2^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "sqrt(x^2+(-1)*a_^2)",
tmpl: "x*sqrt(x^2+(-1)*a_^2)*2^-1 + (-1)*a_^2*acosh(x*a_^-1)*2^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "(x^2+(-1)*a_^2)^(2^-1)",
tmpl: "x*(x^2+(-1)*a_^2)^(2^-1)*2^-1 + (-1)*a_^2*acosh(x*a_^-1)*2^-1",
cond: Some(|b, var| free_q(b, var) && nonzero(b, "a")),
},
RuleSpec::Template {
pat: "x*sin(x^2)",
tmpl: "(-1)*cos(x^2)*2^-1",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "x*cos(x^2)",
tmpl: "sin(x^2)*2^-1",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "x*sinh(x^2)",
tmpl: "cosh(x^2)*2^-1",
cond: Some(free_q),
},
RuleSpec::Template {
pat: "x*cosh(x^2)",
tmpl: "sinh(x^2)*2^-1",
cond: Some(free_q),
},
]
}
fn binomial_integrate<'a>(
ctx: &'a AtomArena<'a>,
bindings: &Bindings<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let m = pred::int_val(bindings, "m")?;
let n = pred::int_val(bindings, "n")?;
let a = pred::bound(bindings, "a")?;
let b = pred::bound(bindings, "b")?;
if !(0..=8).contains(&m) || !(0..=8).contains(&n) {
return None;
}
let coeff: Atom<'a> = match bindings.get(Symbol::new("c")) {
Some(MatchValue::Sequence(slice)) => {
if slice.iter().any(|f| !crate::integral::is_constant(*f, var)) {
return None;
}
if slice.is_empty() {
ctx.num(1)
} else {
ctx.mul(slice)
}
}
_ => ctx.num(1),
};
let x = ctx.var(var.as_str());
let mut terms: Vec<Atom<'a>> = Vec::new();
for k in 0..=n {
let binom = binomial_coeff(n, k);
let term = ctx.mul(&[
ctx.num(binom),
ctx.pow(a, ctx.num(n - k)),
ctx.pow(b, ctx.num(k)),
ctx.pow(x, ctx.num(m + k + 1)),
ctx.pow(ctx.num(m + k + 1), ctx.num(-1)),
]);
terms.push(term);
}
let sum = ctx.add(&terms);
Some(if matches!(coeff.node(), AtomNode::Num(1)) {
sum
} else {
ctx.mul(&[coeff, sum])
})
}
fn binomial_coeff(n: i64, k: i64) -> i64 {
let mut r = 1i64;
for i in 0..k {
r = r * (n - i) / (i + 1);
}
r
}
fn trig_odd_power<'a>(
ctx: &'a AtomArena<'a>,
bindings: &Bindings<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let m = pred::int_val(bindings, "m")?;
let n = pred::int_val(bindings, "n")?;
let x = ctx.var(var.as_str());
let sin = ctx.fun("sin", &[x]);
let cos = ctx.fun("cos", &[x]);
if (1..=9).contains(&m) && m % 2 == 1 && (2..=9).contains(&n) {
let h = (m - 1) / 2; let mut terms: Vec<Atom<'a>> = Vec::new();
for k in 0..=h {
let sign = if k % 2 == 0 { -1i64 } else { 1i64 };
let binom = binomial_coeff(h, k);
let denom = n + 2 * k + 1;
terms.push(ctx.mul(&[
ctx.num(sign * binom),
ctx.pow(cos, ctx.num(denom)),
ctx.pow(ctx.num(denom), ctx.num(-1)),
]));
}
return Some(ctx.add(&terms));
}
if (2..=9).contains(&m) && (1..=9).contains(&n) && n % 2 == 1 {
let h = (n - 1) / 2;
let mut terms: Vec<Atom<'a>> = Vec::new();
for k in 0..=h {
let sign = if k % 2 == 0 { 1i64 } else { -1i64 };
let binom = binomial_coeff(h, k);
let denom = m + 2 * k + 1;
terms.push(ctx.mul(&[
ctx.num(sign * binom),
ctx.pow(sin, ctx.num(denom)),
ctx.pow(ctx.num(denom), ctx.num(-1)),
]));
}
return Some(ctx.add(&terms));
}
None
}
fn trig_odd_power_linear<'a>(
ctx: &'a AtomArena<'a>,
bindings: &Bindings<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let m = pred::int_val(bindings, "m")?;
let n = pred::int_val(bindings, "n")?;
let a = pred::bound(bindings, "a")?;
let b = pred::bound(bindings, "b")?;
let x = ctx.var(var.as_str());
let u = ctx.add(&[ctx.mul(&[a, x]), b]);
let sin = ctx.fun("sin", &[u]);
let cos = ctx.fun("cos", &[u]);
let a_inv = ctx.pow(a, ctx.num(-1));
if (1..=9).contains(&m) && m % 2 == 1 && (2..=9).contains(&n) {
let h = (m - 1) / 2;
let mut terms: Vec<Atom<'a>> = Vec::new();
for k in 0..=h {
let sign = if k % 2 == 0 { -1i64 } else { 1i64 };
let binom = binomial_coeff(h, k);
let denom = n + 2 * k + 1;
terms.push(ctx.mul(&[
ctx.num(sign * binom),
a_inv,
ctx.pow(cos, ctx.num(denom)),
ctx.pow(ctx.num(denom), ctx.num(-1)),
]));
}
return Some(ctx.add(&terms));
}
if (2..=9).contains(&m) && (1..=9).contains(&n) && n % 2 == 1 {
let h = (n - 1) / 2;
let mut terms: Vec<Atom<'a>> = Vec::new();
for k in 0..=h {
let sign = if k % 2 == 0 { 1i64 } else { -1i64 };
let binom = binomial_coeff(h, k);
let denom = m + 2 * k + 1;
terms.push(ctx.mul(&[
ctx.num(sign * binom),
a_inv,
ctx.pow(sin, ctx.num(denom)),
ctx.pow(ctx.num(denom), ctx.num(-1)),
]));
}
return Some(ctx.add(&terms));
}
None
}
fn x_over_sqrt_linear<'a>(
ctx: &'a AtomArena<'a>,
bindings: &Bindings<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let a = pred::bound(bindings, "a")?;
let b = pred::bound(bindings, "b")?;
let x = ctx.var(var.as_str());
let sqrt = ctx.fun("sqrt", &[ctx.add(&[a, ctx.mul(&[b, x])])]);
let inner = ctx.add(&[ctx.mul(&[b, x]), ctx.mul(&[ctx.num(-2), a])]);
let denom = ctx.mul(&[ctx.num(3), b, b]);
Some(ctx.mul(&[ctx.num(2), inner, sqrt, ctx.pow(denom, ctx.num(-1))]))
}
fn is_identifier(name: &str) -> bool {
let mut chars = name.chars();
match chars.next() {
Some(c) if c.is_ascii_alphabetic() => {}
_ => return false,
}
chars.all(|c| c.is_ascii_alphanumeric())
}
fn bake_var(s: &str, var: &str) -> String {
let mut out = String::with_capacity(s.len());
let bytes = s.as_bytes();
let mut i = 0;
while i < bytes.len() {
if bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_' {
let start = i;
while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') {
i += 1;
}
let word = &s[start..i];
if word == "x" {
out.push_str(var);
} else {
out.push_str(word);
}
} else {
out.push(bytes[i] as char);
i += 1;
}
}
out
}
fn wildcard_names(pat: &Pattern<'_>, out: &mut Vec<String>) {
match pat {
Pattern::Wildcard(name, _) => out.push(name.as_str().to_string()),
Pattern::Add(pats) | Pattern::Mul(pats) | Pattern::Fun(_, pats) => {
for p in pats {
wildcard_names(p, out);
}
}
Pattern::Pow(p) => {
wildcard_names(&p.0, out);
wildcard_names(&p.1, out);
}
Pattern::Literal(_) => {}
}
}
pub(crate) fn build_rule_table<'a>(
ctx: &'a AtomArena<'a>,
var: Symbol,
) -> Option<IntegralRuleTable<'a>> {
let var_name = var.as_str();
if !is_identifier(var_name) {
return None;
}
let mut rules: Vec<Rule<'a>> = Vec::new();
let mut closures: Vec<ClosureRule<'a>> = Vec::new();
for spec in rule_specs()
.iter()
.copied()
.chain(crate::integral::rules_ext::specs())
{
let baked_pat = bake_var(spec_pat(&spec), var_name);
let parsed = ocas_parse::parse(ctx, &baked_pat).ok()?;
let pattern = Pattern::from_atom(&crate::pattern_alloc::VecAlloc, parsed);
let mut names = Vec::new();
wildcard_names(&pattern, &mut names);
if names.iter().any(|n| n == var_name) {
return None;
}
let head = pattern_head_key(&pattern);
match spec {
RuleSpec::Template { tmpl, cond, .. } => {
let baked_tmpl = bake_var(tmpl, var_name);
let mut rule = Rule::from_template(
ctx,
&crate::pattern_alloc::VecAlloc,
&baked_pat,
&baked_tmpl,
);
if let Some(pred) = cond {
rule = rule.with_condition(move |b: &Bindings<'a>| pred(b, var));
}
rules.push(rule);
}
RuleSpec::Closure { f, cond, .. } => {
closures.push(ClosureRule {
head,
pattern,
cond,
f,
});
}
}
}
Some(IntegralRuleTable {
table: RuleTable::from_rules(rules),
closures,
})
}
fn spec_pat(spec: &RuleSpec) -> &'static str {
match spec {
RuleSpec::Template { pat, .. } | RuleSpec::Closure { pat, .. } => pat,
}
}
fn pattern_head_key(pattern: &Pattern<'_>) -> HeadKey {
match pattern {
Pattern::Fun(name, _) => HeadKey::Fun(*name),
Pattern::Add(_) => HeadKey::Add,
Pattern::Mul(_) => HeadKey::Mul,
Pattern::Pow(_) => HeadKey::Pow,
Pattern::Literal(_) | Pattern::Wildcard(_, _) => HeadKey::Any,
}
}
impl<'a> IntegralRuleTable<'a> {
fn apply(&self, ctx: &'a AtomArena<'a>, atom: Atom<'a>, var: Symbol) -> Option<Atom<'a>> {
if let Some(r) = self.table.apply(ctx, atom)
&& !has_zero_denominator(ctx, r)
{
return Some(r);
}
let key = head_of(atom);
for rule in &self.closures {
if rule.head != key && rule.head != HeadKey::Any {
continue;
}
let Ok(bindings) = match_pattern(rule.pattern.clone(), atom) else {
continue;
};
if let Some(pred) = rule.cond
&& !pred(&bindings, var)
{
continue;
}
if let Some(r) = (rule.f)(ctx, &bindings, var)
&& !has_zero_denominator(ctx, r)
{
return Some(r);
}
}
None
}
}
fn has_zero_denominator<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>) -> bool {
match expr.node() {
AtomNode::Pow(b, e) => {
if matches!(e.node(), AtomNode::Num(n) if *n < 0) && is_identically_zero(ctx, *b) {
return true;
}
has_zero_denominator(ctx, *b) || has_zero_denominator(ctx, *e)
}
AtomNode::Add(args) | AtomNode::Mul(args) | AtomNode::Fun(_, args) => {
args.iter().any(|a| has_zero_denominator(ctx, *a))
}
AtomNode::Num(_) | AtomNode::Var(_) => false,
}
}
fn is_identically_zero<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>) -> bool {
let collected =
ocas_atom::normalize::normalize(ctx, crate::ode::util::collect_terms(ctx, expr));
matches!(collected.node(), AtomNode::Num(n) if *n == 0)
}
fn fold_trivial_powers<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>) -> Atom<'a> {
match expr.node() {
AtomNode::Pow(base, exp) => match exp.node() {
AtomNode::Num(1) => fold_trivial_powers(ctx, *base),
AtomNode::Num(0) => ctx.num(1),
_ => {
let b = fold_trivial_powers(ctx, *base);
let e = fold_trivial_powers(ctx, *exp);
ctx.pow(b, e)
}
},
AtomNode::Fun(name, args) => {
let rebuilt: Vec<Atom<'a>> =
args.iter().map(|a| fold_trivial_powers(ctx, *a)).collect();
let rebuilt = ctx.fun(name.as_str(), &rebuilt);
if rebuilt == expr { expr } else { rebuilt }
}
AtomNode::Add(args) => {
let rebuilt: Vec<Atom<'a>> =
args.iter().map(|a| fold_trivial_powers(ctx, *a)).collect();
let rebuilt = ctx.add(&rebuilt);
if rebuilt == expr { expr } else { rebuilt }
}
AtomNode::Mul(args) => {
let rebuilt: Vec<Atom<'a>> =
args.iter().map(|a| fold_trivial_powers(ctx, *a)).collect();
let rebuilt = ctx.mul(&rebuilt);
if rebuilt == expr { expr } else { rebuilt }
}
AtomNode::Num(_) | AtomNode::Var(_) => expr,
}
}
pub(crate) fn integrate_rules<'a>(
ctx: &'a AtomArena<'a>,
rules: &IntegralRuleTable<'a>,
expr: Atom<'a>,
var: Symbol,
rule_depth: usize,
) -> Option<Atom<'a>> {
if rule_depth >= MAX_RULE_DEPTH {
return None;
}
let expr = fold_trivial_powers(ctx, expr);
let applied = rules.apply(ctx, expr, var)?;
let applied = fold_trivial_powers(ctx, applied);
Some(resolve_residuals(ctx, applied, var, rule_depth))
}
pub(crate) fn resolve_residuals<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
rule_depth: usize,
) -> Atom<'a> {
match expr.node() {
AtomNode::Fun(name, args) if name.as_str() == "Integral" && args.len() == 2 => {
let v = args[1];
if matches!(v.node(), AtomNode::Var(s) if *s == var) {
let g = fold_trivial_powers(ctx, args[0]);
let resolved =
crate::integral::integrate_raw(ctx, g, var, 0, true, rule_depth + 1, 0);
if resolved == expr {
return expr;
}
return resolve_residuals(ctx, resolved, var, rule_depth);
}
expr
}
AtomNode::Fun(name, args) => {
let rebuilt: Vec<Atom<'a>> = args
.iter()
.map(|a| resolve_residuals(ctx, *a, var, rule_depth))
.collect();
let rebuilt = ctx.fun(name.as_str(), &rebuilt);
if rebuilt == expr { expr } else { rebuilt }
}
AtomNode::Add(args) => {
let rebuilt: Vec<Atom<'a>> = args
.iter()
.map(|a| resolve_residuals(ctx, *a, var, rule_depth))
.collect();
let rebuilt = ctx.add(&rebuilt);
if rebuilt == expr { expr } else { rebuilt }
}
AtomNode::Mul(args) => {
let rebuilt: Vec<Atom<'a>> = args
.iter()
.map(|a| resolve_residuals(ctx, *a, var, rule_depth))
.collect();
let rebuilt = ctx.mul(&rebuilt);
if rebuilt == expr { expr } else { rebuilt }
}
AtomNode::Pow(base, exp) => {
let b = resolve_residuals(ctx, *base, var, rule_depth);
let e = resolve_residuals(ctx, *exp, var, rule_depth);
let rebuilt = ctx.pow(b, e);
if rebuilt == expr { expr } else { rebuilt }
}
AtomNode::Num(_) | AtomNode::Var(_) => expr,
}
}
#[cfg(test)]
mod tests {
use super::*;
use ocas_core::arena::Arena;
#[test]
fn bake_var_replaces_whole_words_only() {
assert_eq!(bake_var("sin(x)^n_ + exp(x)", "t"), "sin(t)^n_ + exp(t)");
assert_eq!(bake_var("x_*exp(x)", "y"), "x_*exp(y)");
assert_eq!(bake_var("exp(x)", "t"), "exp(t)");
assert_eq!(bake_var("x", "t"), "t");
assert_eq!(bake_var("exp(x)+x_2", "t"), "exp(t)+x_2");
}
#[test]
fn build_rule_table_guards_bad_variable_names() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
assert!(build_rule_table(&ctx, Symbol::new("1 x")).is_none());
assert!(build_rule_table(&ctx, Symbol::new("x")).is_some());
assert!(build_rule_table(&ctx, Symbol::new("t")).is_some());
assert!(build_rule_table(&ctx, Symbol::new("m")).is_none());
}
fn int_str(input: &str, var: &str) -> String {
use ocas_core::arena::Arena;
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, input).unwrap();
crate::integrate(&ctx, expr, Symbol::new(var)).to_string()
}
fn assert_solved(input: &str) {
let r = int_str(input, "x");
assert!(
!r.contains("Integral("),
"integrate({input}) left a residue: {r}"
);
}
#[test]
fn family_a_powers() {
assert_solved("x^n");
assert_solved("(a+b*x)^n");
assert_solved("x^2*(a+b*x)^3");
assert_solved("2*x^3*(a+b*x)^2");
assert_solved("x^5*(a+b*x)^4");
}
#[test]
fn family_b_exp_log() {
assert_solved("exp(3*x+1)");
assert_solved("x^3*exp(2*x)");
assert_solved("x*exp(-x^2)");
assert_solved("exp(2*x)*sin(3*x)");
assert_solved("log(x)^3");
assert_solved("x^2*log(x)^2");
assert_solved("1/(x*log(x))");
}
#[test]
fn family_c_trig() {
assert_solved("sin(x)^5");
assert_solved("cos(x)^4");
assert_solved("tan(x)^4");
assert_solved("sec(x)^4");
assert_solved("cot(x)^3");
assert_solved("csc(x)^3");
assert_solved("sin(x)^3*cos(x)^2");
assert_solved("sin(x)^2*cos(x)^5");
assert_solved("sin(x)*cos(x)^3");
assert_solved("sin(x)^3*cos(x)");
assert_solved("sin(2*x)*sin(3*x)");
assert_solved("cos(2*x)*cos(3*x)");
assert_solved("sin(2*x)*cos(3*x)");
let r = int_str("tan(x)^4", "x");
assert!(!r.contains("Integral("), "tan^4 left a residue: {r}");
assert!(r.contains("tan(x)") && r.contains("x"), "tan^4 result: {r}");
assert!(r.contains("^3"), "tan^4 result lacks the cubic term: {r}");
}
#[test]
fn family_d_hyperbolic() {
assert_solved("sinh(x)^4");
assert_solved("cosh(x)^4");
assert_solved("tanh(x)^3");
assert_solved("coth(x)^3");
assert_solved("sech(x)^3");
assert_solved("csch(x)^3");
assert_solved("sinh(2*x+1)");
assert_solved("cosh(2*x+1)");
assert_solved("tanh(2*x+1)");
assert_solved("coth(2*x+1)");
}
#[test]
fn family_e_inverse_trig() {
assert_solved("asin(x)");
assert_solved("acos(x)");
assert_solved("atan(x)");
assert_solved("asinh(x)");
assert_solved("acosh(x)");
assert_solved("atanh(x)");
}
#[test]
fn family_f_rational_intercepts() {
assert_solved("1/(a^2+x^2)");
assert_solved("1/(a^2-x^2)");
assert_solved("x/(a+b*x^2)");
}
#[test]
fn family_g_radicals() {
assert_solved("sqrt(a+b*x)");
assert_solved("1/sqrt(a+b*x)");
assert_solved("x/sqrt(a+b*x)");
assert_solved("sqrt(a^2-x^2)");
assert_solved("1/sqrt(a^2-x^2)");
assert_solved("1/sqrt(x^2+a^2)");
assert_solved("1/sqrt(x^2-a^2)");
assert_solved("sqrt(x^2+a^2)");
assert_solved("sqrt(x^2-a^2)");
}
#[test]
fn family_h_special_forms() {
assert_solved("x*sin(x^2)");
assert_solved("x*cos(x^2)");
assert_solved("x*sinh(x^2)");
assert_solved("x*cosh(x^2)");
}
#[test]
fn rules_off_returns_unevaluated() {
use ocas_core::arena::Arena;
let probes = ["csc(x)^5", "sec(x)^6", "csc(x)^7", "sec(x)^3", "csc(x)^3"];
let mut gated = 0usize;
let mut details = Vec::new();
for input in probes {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, input).unwrap();
let var = Symbol::new("x");
let on = crate::integrate_with_options(
&ctx,
expr,
var,
crate::IntegrateOptions { rules: true },
);
let off = crate::integrate_with_options(
&ctx,
expr,
var,
crate::IntegrateOptions { rules: false },
);
let solved_on = !on.to_string().contains("Integral(");
let solved_off = !off.to_string().contains("Integral(");
if solved_on && !solved_off {
gated += 1;
}
details.push(format!(
"{input}: rules=on solved={solved_on}, rules=off solved={solved_off}"
));
}
assert!(
gated > 0,
"no probe distinguishes `rules: true` from `rules: false`, so the rule \
table is either always or never consulted:\n{}",
details.join("\n")
);
}
#[test]
fn a4_declines_when_the_coefficient_sequence_depends_on_the_variable() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let input = "x^2*(d + e*x)^3*(a + b*log(c*x^n))";
let integrand = ocas_parse::parse(&ctx, input).unwrap();
let result = crate::integrate(&ctx, integrand, Symbol::new("x"));
let text = result.to_string();
if text.contains("Integral(") {
return;
}
let derivative = crate::derivative::diff(&ctx, result, Symbol::new("x"));
assert!(
!derivative.to_string().contains("Derivative("),
"antiderivative contains a head the derivative table does not know: {text}"
);
let base: Vec<(Symbol, f64)> = [
("a", 1.3),
("b", 0.7),
("c", 1.1),
("d", 0.9),
("e", 1.4),
("n", 1.6),
]
.iter()
.map(|(s, v)| (Symbol::new(s), *v))
.collect();
for xv in [0.4, 0.8, 1.3, 1.9] {
let mut env = base.clone();
env.push((Symbol::new("x"), xv));
let lhs = eval_env(derivative, &env).expect("derivative evaluates");
let rhs = eval_env(integrand, &env).expect("integrand evaluates");
assert!(
(lhs - rhs).abs() < 1e-6 * rhs.abs().max(1.0),
"at x={xv}: d/dx = {lhs}, integrand = {rhs}\n {text}"
);
}
assert_solved("x^2*(a+b*x)^3");
assert_solved("2*x^3*(a+b*x)^2");
}
fn eval_env(atom: Atom<'_>, env: &[(Symbol, f64)]) -> Option<f64> {
match atom.node() {
AtomNode::Num(n) => Some(*n as f64),
AtomNode::Var(v) => env.iter().find(|(s, _)| s == v).map(|(_, x)| *x),
AtomNode::Add(args) => args
.iter()
.try_fold(0.0, |acc, a| Some(acc + eval_env(*a, env)?)),
AtomNode::Mul(args) => args
.iter()
.try_fold(1.0, |acc, a| Some(acc * eval_env(*a, env)?)),
AtomNode::Pow(b, e) => Some(eval_env(*b, env)?.powf(eval_env(*e, env)?)),
AtomNode::Fun(name, args) => {
let v = eval_env(*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(),
"log" => v.ln(),
_ => return None,
})
}
}
}
fn resonant_env() -> Vec<(Symbol, f64)> {
vec![
(Symbol::new("a"), 1.0),
(Symbol::new("b"), 2.0),
(Symbol::new("c"), 0.5),
(Symbol::new("d"), 1.5),
]
}
#[test]
fn product_to_sum_never_emits_a_resonant_denominator() {
let shapes = [
"cos(c + d*x)^2",
"sin(c + d*x)^2",
"sin(c + d*x)*cos(c + d*x)",
"sin(c + d*x)*sin(c + d*x)",
"cos(c + d*x)*cos(c + d*x)",
"(a*cos(c + d*x) + b*sin(c + d*x))^2",
"cos(c + d*x)^3*(a + a*sin(c + d*x))",
"sin(c + d*x)^2*(a + b*sin(c + d*x)^2)",
"sin(2*x)*cos(3*x)",
"sin(x)*sin(2*x)",
"cos(2*x)*cos(3*x)",
];
for input in shapes {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr =
ocas_atom::normalize::normalize(&ctx, ocas_parse::parse(&ctx, input).unwrap());
let table = build_rule_table(&ctx, Symbol::new("x")).unwrap();
let Some(r) = integrate_rules(&ctx, &table, expr, Symbol::new("x"), 0) else {
continue; };
assert!(
!has_zero_denominator(&ctx, r),
"identically zero denominator for {input}: {r}"
);
let d = crate::diff(&ctx, r, Symbol::new("x"));
for &xv in &[-1.3, -0.4, 0.7, 1.9] {
let mut env = resonant_env();
env.push((Symbol::new("x"), xv));
let lhs = eval_env(d, &env).expect("eval derivative");
let rhs = eval_env(expr, &env).expect("eval integrand");
assert!(
lhs.is_finite() && rhs.is_finite(),
"{input} at x={xv}: non-finite (derivative {lhs}, integrand {rhs}) in {r}"
);
let tol = 1e-6 * rhs.abs().max(1.0);
assert!(
(lhs - rhs).abs() < tol,
"{input} at x={xv}: derivative {lhs} != integrand {rhs} (result {r})"
);
}
}
}
#[test]
fn resonant_shapes_stay_solved() {
for input in [
"cos(c + d*x)^2",
"sin(c + d*x)^2",
"sin(c + d*x)*cos(c + d*x)",
"(a*cos(c + d*x) + b*sin(c + d*x))^2",
] {
let r = int_str(input, "x");
assert!(
!r.contains("Integral("),
"resonant shape {input} left a residue: {r}"
);
}
}
fn eval_num(atom: Atom<'_>, x: f64) -> f64 {
match atom.node() {
AtomNode::Num(n) => *n as f64,
AtomNode::Var(_) => x,
AtomNode::Add(args) => args.iter().map(|a| eval_num(*a, x)).sum(),
AtomNode::Mul(args) => args.iter().map(|a| eval_num(*a, x)).product(),
AtomNode::Pow(b, e) => eval_num(*b, x).powf(eval_num(*e, x)),
AtomNode::Fun(name, args) => {
let v = eval_num(args[0], x);
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(),
"sinh" => v.sinh(),
"cosh" => v.cosh(),
"tanh" => v.tanh(),
"sech" => v.cosh().recip(),
"csch" => v.sinh().recip(),
"coth" => v.tanh().recip(),
"log" => v.ln(),
"sqrt" => v.sqrt(),
other => panic!("eval_num: unsupported {other}"),
}
}
}
}
fn assert_numeric_antiderivative(input: &str, samples: &[f64]) {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, input).unwrap();
let r = crate::integrate(&ctx, expr, Symbol::new("x"));
assert!(!r.to_string().contains("Integral("), "residue: {r}");
let d = crate::diff(&ctx, r, Symbol::new("x"));
for &xv in samples {
let lhs = eval_num(d, xv);
let rhs = eval_num(expr, xv);
let tol = 1e-6 * rhs.abs().max(1.0);
assert!(
(lhs - rhs).abs() < tol,
"{input} at x={xv}: diff={lhs} integrand={rhs}"
);
}
}
#[test]
fn linear_arg_power_reductions_numeric() {
assert_numeric_antiderivative("sin(2*x+1)^3", &[0.2, 0.6, 1.1]);
assert_numeric_antiderivative("cos(3*x+1)^4", &[0.2, 0.5, 0.9]);
assert_numeric_antiderivative("tan(2*x+1)^3", &[0.3, 0.5, 0.7]);
assert_numeric_antiderivative("sec(2*x+1)^4", &[0.2, 0.4, 0.6]);
assert_numeric_antiderivative("csc(3*x+1)^3", &[0.4, 0.7, 1.0]);
assert_numeric_antiderivative("cot(2*x+1)^3", &[0.3, 0.6, 0.9]);
assert_numeric_antiderivative("sinh(2*x+1)^3", &[0.2, 0.5, 0.8]);
assert_numeric_antiderivative("cosh(3*x+1)^4", &[0.1, 0.3, 0.5]);
assert_numeric_antiderivative("tanh(2*x+1)^3", &[0.2, 0.6, 1.0]);
}
}