#![allow(clippy::collapsible_if)]
#![allow(clippy::missing_const_for_thread_local)]
pub(crate) mod heuristic;
pub mod rational;
pub(crate) mod rde;
pub(crate) mod risch;
pub(crate) mod rules;
pub(crate) mod special;
pub(crate) mod symbolic_rational;
pub(crate) mod trig;
pub(crate) mod trig_reduce;
use ocas_atom::normalize::normalize;
use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use ocas_core::error::Result;
use ocas_core::fuel::Fuel;
use ocas_rewrite::rules::default_rules;
use ocas_rewrite::simplify::{simplify, simplify_with_fuel};
use crate::rules::calculus_rules;
const MAX_DEPTH: usize = 8;
const MAX_CHAIN_ENTRIES: u32 = 256;
thread_local! {
static CHAIN_ENTRIES: std::cell::Cell<u32> = const { std::cell::Cell::new(0) };
}
fn reset_chain_budget() {
CHAIN_ENTRIES.with(|c| c.set(0));
}
fn chain_budget_exhausted() -> bool {
CHAIN_ENTRIES.with(|c| {
let v = c.get().saturating_add(1);
c.set(v);
v > MAX_CHAIN_ENTRIES
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct IntegrateOptions {
pub rules: bool,
}
impl Default for IntegrateOptions {
fn default() -> Self {
Self { rules: true }
}
}
pub fn integrate_with_options<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
options: IntegrateOptions,
) -> Atom<'a> {
let normalized = normalize(ctx, expr);
let calc_rules = calculus_rules(ctx, &crate::pattern_alloc::VecAlloc);
let default_rules = default_rules(ctx, &crate::pattern_alloc::VecAlloc);
reset_chain_budget();
let raw = integrate_raw(ctx, normalized, var, 0, options.rules, 0, 0);
let after_default = simplify(ctx, raw, &default_rules, 20);
let after_calc = simplify(ctx, after_default, &calc_rules, 10);
normalize(ctx, after_calc)
}
pub fn integrate<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, var: Symbol) -> Atom<'a> {
integrate_with_options(ctx, expr, var, IntegrateOptions::default())
}
pub fn integrate_with_fuel<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
fuel: &Fuel,
) -> Result<Atom<'a>> {
let normalized = normalize(ctx, expr);
let calc_rules = calculus_rules(ctx, &crate::pattern_alloc::VecAlloc);
let default_rules = default_rules(ctx, &crate::pattern_alloc::VecAlloc);
reset_chain_budget();
let raw = integrate_raw(ctx, normalized, var, 0, true, 0, 0);
let after_default = simplify_with_fuel(ctx, raw, &default_rules, 20, fuel)?;
let after_calc = simplify_with_fuel(ctx, after_default, &calc_rules, 10, fuel)?;
Ok(normalize(ctx, after_calc))
}
pub fn integrate_heuristic<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, var: Symbol) -> Atom<'a> {
let normalized = normalize(ctx, expr);
reset_chain_budget();
if let Some(r) = heuristic::heuristic_integrate(ctx, normalized, var, 0) {
let calc_rules = calculus_rules(ctx, &crate::pattern_alloc::VecAlloc);
let default_rules = default_rules(ctx, &crate::pattern_alloc::VecAlloc);
let after_default = simplify(ctx, r, &default_rules, 20);
let after_calc = simplify(ctx, after_default, &calc_rules, 10);
return normalize(ctx, after_calc);
}
fallback(ctx, expr, var)
}
pub(crate) fn integrate_raw<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
depth: usize,
rules_enabled: bool,
rule_depth: usize,
parts_depth: usize,
) -> Atom<'a> {
if depth > MAX_DEPTH {
return fallback(ctx, expr, var);
}
match expr.node() {
AtomNode::Num(_) => {
let x = ctx.var(var.as_str());
ctx.mul(&[expr, x])
}
AtomNode::Var(v) => {
if *v == var {
ctx.mul(&[
ctx.pow(ctx.var(var.as_str()), ctx.num(2)),
ctx.pow(ctx.num(2), ctx.num(-1)),
])
} else {
ctx.mul(&[expr, ctx.var(var.as_str())])
}
}
AtomNode::Add(args) => {
let mut terms = Vec::with_capacity(args.len());
for a in args.iter() {
terms.push(integrate_raw(
ctx,
*a,
var,
depth,
rules_enabled,
rule_depth,
parts_depth,
));
}
ctx.add(&terms)
}
AtomNode::Mul(args) => {
let r = integrate_product(
ctx,
args,
var,
depth,
rules_enabled,
rule_depth,
parts_depth,
);
if is_fallback(&r) {
try_risch_or_fallback(ctx, expr, var, rules_enabled, rule_depth, parts_depth)
} else {
r
}
}
AtomNode::Pow(base, exp) => {
let r = integrate_power(ctx, *base, *exp, var, depth);
if is_fallback(&r) {
try_risch_or_fallback(ctx, expr, var, rules_enabled, rule_depth, parts_depth)
} else {
r
}
}
AtomNode::Fun(name, args) => {
let r = integrate_function(ctx, *name, args, var, depth);
if is_fallback(&r) {
try_risch_or_fallback(ctx, expr, var, rules_enabled, rule_depth, parts_depth)
} else {
r
}
}
}
}
fn try_risch_or_fallback<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
rules_enabled: bool,
rule_depth: usize,
parts_depth: usize,
) -> Atom<'a> {
if chain_budget_exhausted() {
return fallback(ctx, expr, var);
}
if let Some(r) = rational::integrate_rational(ctx, expr, var) {
return r;
}
if let Some(r) = symbolic_rational::integrate_rational_symbolic(ctx, expr, var) {
return r;
}
if let Some(r) = risch::risch_integrate(ctx, expr, var) {
return r;
}
if trig::trig_args_numeric(ctx, expr, var)
&& let Some(exp_form) = trig::trig_to_exp(ctx, expr)
&& let Some(complex_ans) = risch::risch_integrate(ctx, exp_form, var)
{
return trig::realify(ctx, complex_ans);
}
if let Some(r) = special::special_integrate(ctx, expr, ctx.var(var.as_str())) {
return r;
}
if rules_enabled {
if let Some(table) = rules::build_rule_table(ctx, var)
&& let Some(r) = rules::integrate_rules(ctx, &table, expr, var, rule_depth)
{
return r;
}
}
if let Some(reduced) = trig_reduce::trig_reduce_products(ctx, expr, var) {
let candidate = crate::expand::expand_bounded(ctx, reduced).unwrap_or(reduced);
let folded = crate::ode::util::collect_terms(ctx, candidate);
let r = integrate_raw(ctx, folded, var, 0, rules_enabled, rule_depth, parts_depth);
if !is_fallback(&r) {
return r;
}
}
if let Some(r) = heuristic::heuristic_integrate(ctx, expr, var, parts_depth) {
return r;
}
if let Some(expanded) = crate::expand::expand_bounded(ctx, expr) {
let folded = crate::ode::util::collect_terms(ctx, expanded);
let r = integrate_raw(ctx, folded, var, 0, rules_enabled, rule_depth, parts_depth);
if !is_fallback(&r) {
return r;
}
}
fallback(ctx, expr, var)
}
pub(crate) fn fallback<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, var: Symbol) -> Atom<'a> {
ctx.fun("Integral", &[expr, ctx.var(var.as_str())])
}
pub(crate) fn is_constant<'a>(expr: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Num(_) => true,
AtomNode::Var(v) => *v != var,
AtomNode::Add(args) | AtomNode::Mul(args) | AtomNode::Fun(_, args) => {
args.iter().all(|a| is_constant(*a, var))
}
AtomNode::Pow(base, exp) => is_constant(*base, var) && is_constant(*exp, var),
}
}
fn integrate_product<'a>(
ctx: &'a AtomArena<'a>,
args: &'a [Atom<'a>],
var: Symbol,
depth: usize,
rules_enabled: bool,
rule_depth: usize,
parts_depth: usize,
) -> Atom<'a> {
let mut constants: Vec<Atom<'a>> = Vec::new();
let mut non_constant: Vec<Atom<'a>> = Vec::new();
for a in args.iter() {
if is_constant(*a, var) {
constants.push(*a);
} else {
non_constant.push(*a);
}
}
if non_constant.is_empty() {
return ctx.mul(&[ctx.mul(args), ctx.var(var.as_str())]);
}
let core = if non_constant.len() == 1 {
non_constant[0]
} else {
ctx.mul(&non_constant)
};
let integrated_core = integrate_raw(
ctx,
core,
var,
depth + 1,
rules_enabled,
rule_depth,
parts_depth,
);
if is_fallback(&integrated_core) {
return fallback(ctx, ctx.mul(args), var);
}
let mut result_factors = constants;
result_factors.push(integrated_core);
ctx.mul(&result_factors)
}
pub(crate) fn is_fallback<'a>(atom: &Atom<'a>) -> bool {
matches!(atom.node(), AtomNode::Fun(name, _) if name.as_str() == "Integral")
}
fn integrate_power<'a>(
ctx: &'a AtomArena<'a>,
base: Atom<'a>,
exp: Atom<'a>,
var: Symbol,
_depth: usize,
) -> Atom<'a> {
if matches!(exp.node(), AtomNode::Num(0)) {
return ctx.var(var.as_str());
}
if let AtomNode::Var(v) = base.node()
&& *v == var
{
if let AtomNode::Num(n) = exp.node() {
if *n == -1 {
return ctx.fun("log", &[base]);
}
let new_exp = ctx.num(n + 1);
let denom = ctx.num(n + 1);
return ctx.mul(&[ctx.pow(base, new_exp), ctx.pow(denom, ctx.num(-1))]);
}
if let Some((p, q)) = fraction_exponent(exp) {
if p != -q {
let new_exp = ctx.mul(&[ctx.num(p + q), ctx.pow(ctx.num(q), ctx.num(-1))]);
let denom = ctx.mul(&[ctx.num(p + q), ctx.pow(ctx.num(q), ctx.num(-1))]);
return ctx.mul(&[ctx.pow(base, new_exp), ctx.pow(denom, ctx.num(-1))]);
}
}
}
if let AtomNode::Num(n) = exp.node()
&& let Some((a, _b)) = linear_form(ctx, base, var)
{
if *n == -1 {
return ctx.mul(&[ctx.fun("log", &[base]), ctx.pow(a, ctx.num(-1))]);
}
let new_exp = ctx.num(n + 1);
let denom = ctx.mul(&[a, ctx.num(n + 1)]);
return ctx.mul(&[ctx.pow(base, new_exp), ctx.pow(denom, ctx.num(-1))]);
}
if let Some((p, q)) = fraction_exponent(exp)
&& p != -q
&& let Some((a, _b)) = linear_form(ctx, base, var)
{
let new_exp = ctx.mul(&[ctx.num(p + q), ctx.pow(ctx.num(q), ctx.num(-1))]);
let coeff_num = ctx.num(q);
let coeff_den = ctx.mul(&[a, ctx.num(p + q)]);
return ctx.mul(&[
ctx.pow(base, new_exp),
coeff_num,
ctx.pow(coeff_den, ctx.num(-1)),
]);
}
fallback(ctx, ctx.pow(base, exp), var)
}
fn fraction_exponent<'a>(exp: Atom<'a>) -> Option<(i64, i64)> {
if let AtomNode::Mul(args) = exp.node() {
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
{
if den.is_some() {
return None;
}
den = Some(*bb);
} else {
return None;
}
}
_ => return None,
}
}
if let (Some(p), Some(q)) = (num, den)
&& q > 0
{
return Some((p, q));
}
}
None
}
pub(crate) fn linear_form<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>)> {
match expr.node() {
AtomNode::Var(v) if *v == var => Some((ctx.num(1), ctx.num(0))),
AtomNode::Mul(args) => {
let mut coeff = ctx.num(1);
let mut has_var = false;
for a in args.iter() {
if let AtomNode::Var(v) = a.node()
&& *v == var
{
has_var = true;
continue;
}
if is_constant(*a, var) {
coeff = ctx.mul(&[coeff, *a]);
} else {
return None;
}
}
if has_var {
Some((coeff, ctx.num(0)))
} else {
None
}
}
AtomNode::Add(args) => {
let mut a_part = ctx.num(0);
let mut b_part = ctx.num(0);
for arg in args.iter() {
if let Some((ca, _cb)) = linear_form(ctx, *arg, var) {
a_part = ctx.add(&[a_part, ca]);
} else if is_constant(*arg, var) {
b_part = ctx.add(&[b_part, *arg]);
} else {
return None;
}
}
Some((a_part, b_part))
}
_ => None,
}
}
fn integrate_function<'a>(
ctx: &'a AtomArena<'a>,
name: Symbol,
args: &'a [Atom<'a>],
var: Symbol,
_depth: usize,
) -> Atom<'a> {
if args.is_empty() {
return fallback(ctx, ctx.fun(name.as_str(), args), var);
}
let u = args[0];
if let Some((a, _b)) = linear_form(ctx, u, var)
&& is_constant(a, var)
&& !is_one(a)
{
let inner_integral = match name.as_str() {
"sin" => ctx.mul(&[ctx.num(-1), ctx.fun("cos", &[u])]),
"cos" => ctx.fun("sin", &[u]),
"exp" => ctx.fun("exp", &[u]),
_ => return fallback(ctx, ctx.fun(name.as_str(), args), var),
};
return ctx.mul(&[ctx.pow(a, ctx.num(-1)), inner_integral]);
}
if let AtomNode::Var(v) = u.node()
&& *v == var
{
let antiderivative: Option<Atom<'a>> = match name.as_str() {
"sin" => Some(ctx.mul(&[ctx.num(-1), ctx.fun("cos", &[u])])),
"cos" => Some(ctx.fun("sin", &[u])),
"exp" => Some(ctx.fun("exp", &[u])),
"log" => Some(ctx.mul(&[u, ctx.add(&[ctx.fun("log", &[u]), ctx.num(-1)])])),
_ => None,
};
if let Some(anti) = antiderivative {
return anti;
}
}
fallback(ctx, ctx.fun(name.as_str(), args), var)
}
fn is_one<'a>(expr: Atom<'a>) -> bool {
matches!(expr.node(), AtomNode::Num(1))
}
#[cfg(test)]
mod tests {
use ocas_atom::AtomArena;
use ocas_core::arena::Arena;
use super::*;
#[test]
fn integrate_power() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.pow(x, ctx.num(2));
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "(3^-1)*(x^3)");
}
#[test]
fn integrate_inverse() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.pow(x, ctx.num(-1));
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "log(x)");
}
#[test]
fn integrate_sin() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("sin", &[x]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "-1*(cos(x))");
}
#[test]
fn integrate_cos() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("cos", &[x]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "sin(x)");
}
#[test]
fn integrate_exp() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("exp", &[x]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "exp(x)");
}
#[test]
fn integrate_linear_substitution() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let two_x_plus_one = ctx.add(&[ctx.mul(&[ctx.num(2), x]), ctx.num(1)]);
let expr = ctx.pow(two_x_plus_one, ctx.num(2));
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "(6^-1)*((1 + (2*x))^3)");
}
#[test]
fn integrate_unknown() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("f", &[x]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "Integral(f(x), x)");
}
#[test]
fn integrate_sin_times_cos_via_trig_path() {
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])]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
}
#[test]
fn integrate_cos_squared_via_trig_path() {
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));
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
}
#[test]
fn integrate_exp_neg_x_squared_gives_erf() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("exp", &[ctx.mul(&[ctx.num(-1), ctx.pow(x, ctx.num(2))])]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(result.to_string().contains("erf"), "got {result}");
assert!(!result.to_string().starts_with("Integral"), "got {result}");
}
#[test]
fn integrate_expand_product_retry() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[x, ctx.pow(ctx.add(&[x, ctx.num(1)]), ctx.num(2))]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
let d = crate::diff(&ctx, result, Symbol::new("x"));
let residual = ctx.add(&[d, ctx.mul(&[ctx.num(-1), expr])]);
let folded = crate::ode::util::collect_terms(&ctx, residual);
assert_eq!(folded.to_string(), "0", "residual: {folded}");
}
#[test]
fn integrate_expand_deep_product() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let s1 = ctx.pow(ctx.add(&[x, ctx.num(1)]), ctx.num(2));
let s2 = ctx.pow(ctx.add(&[x, ctx.num(2)]), ctx.num(2));
let expr = ctx.mul(&[s1, s2]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
let d = crate::diff(&ctx, result, Symbol::new("x"));
let residual = ctx.add(&[d, ctx.mul(&[ctx.num(-1), expr])]);
let folded = crate::ode::util::collect_terms(&ctx, residual);
assert_eq!(folded.to_string(), "0", "residual: {folded}");
}
#[test]
fn integrate_expand_budget_keeps_fallback() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let s = ctx.pow(ctx.add(&[x, ctx.num(1)]), ctx.num(8));
let expr = ctx.mul(&[s, ctx.fun("f", &[x])]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "Integral((f(x))*((1 + x)^8), x)");
}
#[test]
fn integrate_weierstrass_cubed_denominator_terminates() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, "1/(-5 + 3*cos(c + d*x))^3").expect("parse");
let result = integrate(&ctx, expr, Symbol::new("x"));
let _ = result; }
#[test]
fn integrate_parts_weierstrass_cycle_terminates() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, "(c + d*x)^2/(a + a*sin(e + f*x))").expect("parse");
let result = integrate(&ctx, expr, Symbol::new("x"));
let _ = result; }
#[test]
fn integrate_exp_x_over_x_gives_ei() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.fun("exp", &[x]), ctx.pow(x, ctx.num(-1))]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "Ei(x)");
}
}