use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use crate::derivative::diff;
use crate::integral::{contains_integral, integrate_raw, is_constant, is_fallback, node_count};
const PARTS_MAX_DEPTH: usize = 2;
const MAX_SUBST_NODES: usize = 100_000;
thread_local! {
static SUBST_NODES: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
fn reset_subst_budget() {
SUBST_NODES.with(|c| c.set(0));
}
fn subst_budget_exhausted() -> bool {
SUBST_NODES.with(|c| {
let v = c.get().saturating_add(1);
c.set(v);
v > MAX_SUBST_NODES
})
}
pub(crate) fn heuristic_integrate<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
parts_depth: usize,
) -> Option<Atom<'a>> {
reset_subst_budget();
if let Some(r) = try_parts(ctx, expr, var, parts_depth) {
if !contains_integral(r) {
return Some(r);
}
}
let ts = try_trig_substitution(ctx, expr, var);
if let Some(r) = ts {
if !contains_integral(r) {
return Some(r);
}
}
if let Some(r) = try_weierstrass(ctx, expr, var) {
if !contains_integral(r) {
return Some(r);
}
}
if let Some(r) = try_euler_substitution(ctx, expr, var) {
if !contains_integral(r) {
return Some(r);
}
}
None
}
fn liate_score(expr: &Atom<'_>) -> u32 {
match expr.node() {
AtomNode::Fun(name, _) => match name.as_str() {
"log" | "ln" => 0,
"asin" | "acos" | "atan" | "acot" | "asec" | "acsc" | "asinh" | "acosh" | "atanh" => 1,
"sin" | "cos" | "tan" | "sec" | "csc" | "cot" | "sinh" | "cosh" | "tanh" => 3,
"exp" => 4,
_ => 5,
},
AtomNode::Pow(base, exp) => {
if matches!(base.node(), AtomNode::Var(_))
&& matches!(exp.node(), AtomNode::Num(n) if *n > 0)
{
2
} else {
liate_score(base).min(5)
}
}
AtomNode::Var(_) => 2, AtomNode::Mul(args) => args.iter().map(liate_score).min().unwrap_or(5),
_ => 5,
}
}
fn try_parts<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
parts_depth: usize,
) -> Option<Atom<'a>> {
if parts_depth >= PARTS_MAX_DEPTH {
return None;
}
let factors: Vec<Atom<'a>> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => return None,
};
if factors.len() < 2 {
return None;
}
let mut constants = Vec::new();
let mut non_constants = Vec::new();
for f in &factors {
if is_constant(*f, var) {
constants.push(*f);
} else {
non_constants.push(*f);
}
}
if non_constants.len() < 2 {
return None;
}
let (u_idx, _) = non_constants
.iter()
.enumerate()
.min_by_key(|(_, e)| liate_score(e))
.unwrap();
let u = non_constants[u_idx];
let v_prime_factors: Vec<Atom<'a>> = non_constants
.iter()
.enumerate()
.filter(|(i, _)| *i != u_idx)
.map(|(_, e)| *e)
.collect();
let v_prime = if v_prime_factors.len() == 1 {
v_prime_factors[0]
} else {
ctx.mul(&v_prime_factors)
};
let v = integrate_raw(ctx, v_prime, var, parts_depth + 2, true, 0, parts_depth + 1);
if contains_integral(v) {
return None;
}
let u_prime = diff(ctx, u, var);
if is_fallback(&u_prime) {
return None;
}
let u_prime_v = ctx.mul(&[u_prime, v]);
let integral_u_prime_v = integrate_raw(
ctx,
u_prime_v,
var,
parts_depth + 2,
true,
0,
parts_depth + 1,
);
let u_times_v = ctx.mul(&[u, v]);
let core_result = if contains_integral(integral_u_prime_v) {
return None;
} else {
ctx.add(&[u_times_v, ctx.mul(&[ctx.num(-1), integral_u_prime_v])])
};
if constants.is_empty() {
Some(core_result)
} else {
let mut result_factors = constants;
result_factors.push(core_result);
Some(ctx.mul(&result_factors))
}
}
#[allow(dead_code)]
fn is_half_exponent(e: &AtomNode) -> bool {
match e {
AtomNode::Pow(b, exp) => {
matches!(b.node(), AtomNode::Num(2)) && matches!(exp.node(), AtomNode::Num(-1))
}
AtomNode::Mul(args) if args.len() == 2 => {
let has_half = args.iter().any(|a| {
matches!(a.node(), AtomNode::Pow(b, e)
if matches!(b.node(), AtomNode::Num(2)) && matches!(e.node(), AtomNode::Num(-1)))
});
let has_one = args.iter().any(|a| matches!(a.node(), AtomNode::Num(1)));
has_half && has_one
}
_ => false,
}
}
fn is_neg_half_exponent(e: &AtomNode) -> bool {
match e {
AtomNode::Mul(args) if args.len() == 2 => {
let has_neg = args.iter().any(|a| matches!(a.node(), AtomNode::Num(-1)));
let has_half = args.iter().any(|a| {
matches!(a.node(), AtomNode::Pow(b, e)
if matches!(b.node(), AtomNode::Num(2)) && matches!(e.node(), AtomNode::Num(-1)))
});
has_neg && has_half
}
_ => false,
}
}
fn match_subtracted_squares<'a>(
ctx: &'a AtomArena<'a>,
positive: Atom<'a>,
negative: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let x_squared = match negative.node() {
AtomNode::Mul(args) if args.len() == 2 => {
if matches!(args[0].node(), AtomNode::Num(-1)) {
args[1]
} else if matches!(args[1].node(), AtomNode::Num(-1)) {
args[0]
} else {
return None;
}
}
_ => return None,
};
if !matches!(x_squared.node(), AtomNode::Pow(b, e)
if matches!(b.node(), AtomNode::Var(v) if *v == var) && matches!(e.node(), AtomNode::Num(2)))
{
return None;
}
extract_square_base(ctx, positive, var)
}
fn match_sum_squares<'a>(
ctx: &'a AtomArena<'a>,
t1: Atom<'a>,
t2: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let is_x2 = |a: &Atom<'a>| {
matches!(a.node(), AtomNode::Pow(b, e)
if matches!(b.node(), AtomNode::Var(v) if *v == var) && matches!(e.node(), AtomNode::Num(2)))
};
if is_x2(&t1) && is_constant(t2, var) {
extract_square_base(ctx, t2, var)
} else if is_x2(&t2) && is_constant(t1, var) {
extract_square_base(ctx, t1, var)
} else {
None
}
}
fn match_x_minus_a_squared<'a>(
ctx: &'a AtomArena<'a>,
positive: Atom<'a>,
negative: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if !matches!(positive.node(), AtomNode::Pow(b, e)
if matches!(b.node(), AtomNode::Var(v) if *v == var) && matches!(e.node(), AtomNode::Num(2)))
{
return None;
}
let a_squared = match negative.node() {
AtomNode::Mul(args) if args.len() == 2 => {
if matches!(args[0].node(), AtomNode::Num(-1)) {
args[1]
} else if matches!(args[1].node(), AtomNode::Num(-1)) {
args[0]
} else {
return None;
}
}
_ => return None,
};
if is_constant(a_squared, var) {
extract_square_base(ctx, a_squared, var)
} else {
None
}
}
fn extract_square_base<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
match expr.node() {
AtomNode::Pow(b, e) => {
if matches!(e.node(), AtomNode::Num(2)) && is_constant(*b, var) {
Some(*b)
} else {
None
}
}
AtomNode::Num(n) => {
if *n > 0 {
let sqrt = (*n as f64).sqrt();
if sqrt == sqrt.floor() {
Some(ctx.num(sqrt as i64))
} else {
None
}
} else {
None
}
}
_ => None,
}
}
fn try_trig_substitution<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
if let Some(antideriv) = match_inv_sqrt_a2_minus_x2(ctx, expr, var) {
return Some(antideriv);
}
if let Some(antideriv) = match_inv_sqrt_a2_plus_x2(ctx, expr, var) {
return Some(antideriv);
}
if let Some(antideriv) = match_inv_sqrt_x2_minus_a2(ctx, expr, var) {
return Some(antideriv);
}
if let Some(antideriv) = match_sqrt_a2_minus_x2(ctx, expr, var) {
return Some(antideriv);
}
if let Some(antideriv) = match_sqrt_a2_plus_x2(ctx, expr, var) {
return Some(antideriv);
}
if let Some(antideriv) = match_sqrt_x2_minus_a2(ctx, expr, var) {
return Some(antideriv);
}
None
}
fn match_inv_sqrt_a2_minus_x2<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let base = get_sqrt_base(expr, var)?;
if let Some(_a) = match_subtracted_squares_pattern(ctx, base, var) {
let x = ctx.var(var.as_str());
let a = _a;
if is_one_atom(a) {
return Some(ctx.fun("asin", &[x]));
}
return Some(ctx.fun("asin", &[ctx.mul(&[x, ctx.pow(a, ctx.num(-1))])]));
}
None
}
fn match_inv_sqrt_a2_plus_x2<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let base = get_sqrt_base(expr, var)?;
if let Some(_a) = match_sum_squares_pattern(ctx, base, var) {
let x = ctx.var(var.as_str());
let a = _a;
if is_one_atom(a) {
return Some(ctx.fun("asinh", &[x]));
}
return Some(ctx.fun("asinh", &[ctx.mul(&[x, ctx.pow(a, ctx.num(-1))])]));
}
None
}
fn match_inv_sqrt_x2_minus_a2<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let base = get_sqrt_base(expr, var)?;
if let Some(_a) = match_x_minus_a_sq_pattern(ctx, base, var) {
let x = ctx.var(var.as_str());
let a = _a;
if is_one_atom(a) {
return Some(ctx.fun("acosh", &[x]));
}
return Some(ctx.fun("acosh", &[ctx.mul(&[x, ctx.pow(a, ctx.num(-1))])]));
}
None
}
fn match_sqrt_a2_minus_x2<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let base = get_sqrt_base_positive(expr, var)?;
if let Some(_a) = match_subtracted_squares_pattern(ctx, base, var) {
let x = ctx.var(var.as_str());
let a = _a;
let sqrt_part = ctx.mul(&[x, expr]); let a_sq = ctx.pow(a, ctx.num(2));
let asin_arg = if is_one_atom(a) {
x
} else {
ctx.mul(&[x, ctx.pow(a, ctx.num(-1))])
};
let asin_part = ctx.mul(&[a_sq, ctx.fun("asin", &[asin_arg])]);
return Some(ctx.mul(&[
ctx.add(&[sqrt_part, asin_part]),
ctx.pow(ctx.num(2), ctx.num(-1)),
]));
}
None
}
fn match_sqrt_a2_plus_x2<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let base = get_sqrt_base_positive(expr, var)?;
if let Some(_a) = match_sum_squares_pattern(ctx, base, var) {
let x = ctx.var(var.as_str());
let a = _a;
let sqrt_part = ctx.mul(&[x, expr]);
let a_sq = ctx.pow(a, ctx.num(2));
let asinh_arg = if is_one_atom(a) {
x
} else {
ctx.mul(&[x, ctx.pow(a, ctx.num(-1))])
};
let asinh_part = ctx.mul(&[a_sq, ctx.fun("asinh", &[asinh_arg])]);
return Some(ctx.mul(&[
ctx.add(&[sqrt_part, asinh_part]),
ctx.pow(ctx.num(2), ctx.num(-1)),
]));
}
None
}
fn match_sqrt_x2_minus_a2<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let base = get_sqrt_base_positive(expr, var)?;
if let Some(_a) = match_x_minus_a_sq_pattern(ctx, base, var) {
let x = ctx.var(var.as_str());
let a = _a;
let sqrt_part = ctx.mul(&[x, expr]);
let a_sq = ctx.pow(a, ctx.num(2));
let acosh_arg = if is_one_atom(a) {
x
} else {
ctx.mul(&[x, ctx.pow(a, ctx.num(-1))])
};
let acosh_part = ctx.mul(&[a_sq, ctx.fun("acosh", &[acosh_arg])]);
return Some(ctx.mul(&[
ctx.add(&[sqrt_part, ctx.mul(&[ctx.num(-1), acosh_part])]),
ctx.pow(ctx.num(2), ctx.num(-1)),
]));
}
None
}
fn get_sqrt_base<'a>(expr: Atom<'a>, var: Symbol) -> Option<Atom<'a>> {
match expr.node() {
AtomNode::Pow(b, e) => {
if is_neg_half_exponent(e.node()) {
return Some(*b);
}
if matches!(e.node(), AtomNode::Num(-1)) {
if let Some(inner) = get_sqrt_base_positive(*b, var) {
return Some(inner);
}
}
None
}
_ => None,
}
}
fn get_sqrt_base_positive<'a>(expr: Atom<'a>, _var: Symbol) -> Option<Atom<'a>> {
match expr.node() {
AtomNode::Pow(b, e) => {
if is_half_exponent(e.node()) {
Some(*b)
} else {
None
}
}
_ => None,
}
}
fn is_one_atom(expr: Atom<'_>) -> bool {
matches!(expr.node(), AtomNode::Num(1))
}
fn match_subtracted_squares_pattern<'a>(
ctx: &'a AtomArena<'a>,
inner: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let args = match inner.node() {
AtomNode::Add(a) if a.len() == 2 => a,
_ => return None,
};
let (t1, t2) = (args[0], args[1]);
match_subtracted_squares(ctx, t1, t2, var)
.or_else(|| match_subtracted_squares(ctx, t2, t1, var))
}
fn match_sum_squares_pattern<'a>(
ctx: &'a AtomArena<'a>,
inner: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let args = match inner.node() {
AtomNode::Add(a) if a.len() == 2 => a,
_ => return None,
};
let (t1, t2) = (args[0], args[1]);
match_sum_squares(ctx, t1, t2, var)
}
fn match_x_minus_a_sq_pattern<'a>(
ctx: &'a AtomArena<'a>,
inner: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let args = match inner.node() {
AtomNode::Add(a) if a.len() == 2 => a,
_ => return None,
};
let (t1, t2) = (args[0], args[1]);
match_x_minus_a_squared(ctx, t1, t2, var).or_else(|| match_x_minus_a_squared(ctx, t2, t1, var))
}
fn substitute_atom<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
replacement: Atom<'a>,
) -> Option<Atom<'a>> {
if subst_budget_exhausted() {
return None;
}
match expr.node() {
AtomNode::Var(v) => {
if *v == var {
Some(replacement)
} else {
Some(expr)
}
}
AtomNode::Num(_) => Some(expr),
AtomNode::Add(args) => {
let new_args: Vec<Atom<'a>> = args
.iter()
.map(|a| substitute_atom(ctx, *a, var, replacement))
.collect::<Option<_>>()?;
Some(ctx.add(&new_args))
}
AtomNode::Mul(args) => {
let new_args: Vec<Atom<'a>> = args
.iter()
.map(|a| substitute_atom(ctx, *a, var, replacement))
.collect::<Option<_>>()?;
Some(ctx.mul(&new_args))
}
AtomNode::Pow(base, exp) => {
let new_base = substitute_atom(ctx, *base, var, replacement)?;
let new_exp = substitute_atom(ctx, *exp, var, replacement)?;
Some(ctx.pow(new_base, new_exp))
}
AtomNode::Fun(name, args) => {
let new_args: Vec<Atom<'a>> = args
.iter()
.map(|a| substitute_atom(ctx, *a, var, replacement))
.collect::<Option<_>>()?;
Some(ctx.fun(name.as_str(), &new_args))
}
}
}
fn is_trig_rational<'a>(expr: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Num(_) => true,
AtomNode::Var(_) => is_constant(expr, var),
AtomNode::Add(args) | AtomNode::Mul(args) => args.iter().all(|a| is_trig_rational(*a, var)),
AtomNode::Pow(base, exp) => {
matches!(exp.node(), AtomNode::Num(_)) && is_trig_rational(*base, var)
}
AtomNode::Fun(name, args) => {
let n = name.as_str();
if n == "sin" || n == "cos" {
args.len() == 1 && is_linear_in(args[0], var)
} else {
is_constant(expr, var)
}
}
}
}
fn is_linear_in<'a>(expr: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Var(v) => *v == var,
AtomNode::Num(_) => true,
AtomNode::Add(args) => args
.iter()
.all(|a| is_linear_in(*a, var) || is_constant(*a, var)),
AtomNode::Mul(args) => {
let mut has_var = false;
for a in args.iter() {
if let AtomNode::Var(v) = a.node() {
if *v == var {
has_var = true;
}
} else if !is_constant(*a, var) {
return false;
}
}
has_var
}
_ => false,
}
}
fn trig_count(expr: Atom<'_>) -> u32 {
match expr.node() {
AtomNode::Fun(name, args) => {
let n = name.as_str();
let base = if n == "sin" || n == "cos" { 1 } else { 0 };
base + args.iter().map(|a| trig_count(*a)).sum::<u32>()
}
AtomNode::Add(args) | AtomNode::Mul(args) => args.iter().map(|a| trig_count(*a)).sum(),
AtomNode::Pow(base, exp) => trig_count(*base) + trig_count(*exp),
AtomNode::Num(_) | AtomNode::Var(_) => 0,
}
}
fn trig_linear_arg(expr: Atom<'_>, var: Symbol) -> Option<Atom<'_>> {
fn visit<'a>(expr: Atom<'a>, var: Symbol, found: &mut Option<Atom<'a>>) -> Option<()> {
match expr.node() {
AtomNode::Fun(name, args) if args.len() == 1 => {
let n = name.as_str();
if n == "sin" || n == "cos" {
let u = args[0];
if !contains_var(u, var) {
return Some(());
}
match found {
Some(prev) if *prev != u => return None,
_ => *found = Some(u),
}
Some(())
} else {
if contains_var(expr, var) {
return None;
}
Some(())
}
}
AtomNode::Fun(_, _) => {
if contains_var(expr, var) {
return None;
}
Some(())
}
AtomNode::Add(args) | AtomNode::Mul(args) => {
for a in args.iter() {
visit(*a, var, found)?;
}
Some(())
}
AtomNode::Pow(base, exp) => {
visit(*base, var, found)?;
visit(*exp, var, found)?;
Some(())
}
AtomNode::Num(_) => Some(()),
AtomNode::Var(v) => {
if *v == var {
return None;
}
Some(())
}
}
}
let mut found: Option<Atom<'_>> = None;
visit(expr, var, &mut found)?;
found
}
fn contains_var(expr: Atom<'_>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Var(v) => *v == var,
AtomNode::Num(_) => false,
AtomNode::Fun(_, args) => args.iter().any(|a| contains_var(*a, var)),
AtomNode::Add(args) | AtomNode::Mul(args) => args.iter().any(|a| contains_var(*a, var)),
AtomNode::Pow(base, exp) => contains_var(*base, var) || contains_var(*exp, var),
}
}
fn substitute_trig_arg<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
u: Atom<'a>,
var: Symbol,
sin_u: Atom<'a>,
cos_u: Atom<'a>,
) -> Option<Atom<'a>> {
if subst_budget_exhausted() {
return None;
}
match expr.node() {
AtomNode::Fun(name, args) if args.len() == 1 => {
let n = name.as_str();
if n == "sin" || n == "cos" {
if args[0] == u {
return Some(if n == "sin" { sin_u } else { cos_u });
}
return None;
}
if contains_var(expr, var) {
return None;
}
Some(expr)
}
AtomNode::Num(_) | AtomNode::Var(_) => {
if contains_var(expr, var) {
return None;
}
Some(expr)
}
AtomNode::Add(args) => {
let mut rebuilt = Vec::with_capacity(args.len());
for a in args.iter() {
rebuilt.push(substitute_trig_arg(ctx, *a, u, var, sin_u, cos_u)?);
}
Some(ctx.add(&rebuilt))
}
AtomNode::Mul(args) => {
let mut rebuilt = Vec::with_capacity(args.len());
for a in args.iter() {
rebuilt.push(substitute_trig_arg(ctx, *a, u, var, sin_u, cos_u)?);
}
Some(ctx.mul(&rebuilt))
}
AtomNode::Pow(base, exp) => {
let b = substitute_trig_arg(ctx, *base, u, var, sin_u, cos_u)?;
let e = substitute_trig_arg(ctx, *exp, u, var, sin_u, cos_u)?;
Some(ctx.pow(b, e))
}
AtomNode::Fun(_, _) => {
if contains_var(expr, var) {
return None;
}
Some(expr)
}
}
}
fn try_weierstrass<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>, var: Symbol) -> Option<Atom<'a>> {
if !is_trig_rational(expr, var) {
return None;
}
if trig_count(expr) == 0 {
return None;
}
let u = trig_linear_arg(expr, var)?;
let (a, _b) = crate::integral::linear_form(ctx, u, var)?;
if matches!(a.node(), AtomNode::Num(0)) {
return None;
}
let t = ctx.var("_t");
let one_plus_t2 = ctx.add(&[ctx.num(1), ctx.pow(t, ctx.num(2))]);
let sin_u = ctx.mul(&[ctx.num(2), t, ctx.pow(one_plus_t2, ctx.num(-1))]);
let cos_u = ctx.mul(&[
ctx.add(&[ctx.num(1), ctx.mul(&[ctx.num(-1), ctx.pow(t, ctx.num(2))])]),
ctx.pow(one_plus_t2, ctx.num(-1)),
]);
let dx_dt = ctx.mul(&[
ctx.num(2),
ctx.pow(a, ctx.num(-1)),
ctx.pow(one_plus_t2, ctx.num(-1)),
]);
let substituted = substitute_trig_arg(ctx, expr, u, var, sin_u, cos_u)?;
let integrand = ctx.mul(&[substituted, dx_dt]);
if node_count(integrand) > MAX_SUBST_NODES {
return None;
}
let t_sym = Symbol::new("_t");
if !crate::integral::symbolic_rational::rational_complexity_ok(ctx, integrand, t_sym) {
return None;
}
let result_t = integrate_raw(ctx, integrand, t_sym, 2, true, 0, 0);
if is_fallback(&result_t) {
return None;
}
let back = ctx.fun("tan", &[ctx.mul(&[u, ctx.pow(ctx.num(2), ctx.num(-1))])]);
substitute_atom(ctx, result_t, t_sym, back)
}
fn replace_sqrt<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
quad_t: Atom<'a>,
sqrt_t: Atom<'a>,
) -> Atom<'a> {
match expr.node() {
AtomNode::Fun(name, args) if name.as_str() == "sqrt" && args.len() == 1 => {
let q = ocas_atom::normalize::normalize(ctx, args[0]);
if q == quad_t { sqrt_t } else { expr }
}
AtomNode::Pow(base, exp) => {
if let (AtomNode::Fun(name, args), AtomNode::Num(-1)) = (base.node(), exp.node())
&& name.as_str() == "sqrt"
&& args.len() == 1
&& ocas_atom::normalize::normalize(ctx, args[0]) == quad_t
{
return ctx.pow(sqrt_t, ctx.num(-1));
}
let b = replace_sqrt(ctx, *base, quad_t, sqrt_t);
let e = replace_sqrt(ctx, *exp, quad_t, sqrt_t);
ctx.pow(b, e)
}
AtomNode::Add(args) => {
let rebuilt: Vec<Atom<'a>> = args
.iter()
.map(|a| replace_sqrt(ctx, *a, quad_t, sqrt_t))
.collect();
ctx.add(&rebuilt)
}
AtomNode::Mul(args) => {
let rebuilt: Vec<Atom<'a>> = args
.iter()
.map(|a| replace_sqrt(ctx, *a, quad_t, sqrt_t))
.collect();
ctx.mul(&rebuilt)
}
AtomNode::Fun(name, args) => {
let rebuilt: Vec<Atom<'a>> = args
.iter()
.map(|a| replace_sqrt(ctx, *a, quad_t, sqrt_t))
.collect();
ctx.fun(name.as_str(), &rebuilt)
}
AtomNode::Num(_) | AtomNode::Var(_) => expr,
}
}
fn try_euler_substitution<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let (a, b, c) = crate::integral::rules::quadratic_coeffs(ctx, expr, var)?;
let (pa, qa) = crate::integral::rules::rat_of(a)?;
let (pc, qc) = crate::integral::rules::rat_of(c)?;
let t = ctx.var("_t");
let x = ctx.var(var.as_str());
let rat = |ctx: &'a AtomArena<'a>, p: i64, q: i64| -> Atom<'a> {
ctx.mul(&[ctx.num(p), ctx.pow(ctx.num(q), ctx.num(-1))])
};
let (x_t, sqrt_t, dx_dt): (Atom<'a>, Atom<'a>, Atom<'a>) =
if let Some((r, s)) = crate::integral::rules::rational_sqrt(pa, qa) {
let sq_a = rat(ctx, r, s);
let denom = ctx.add(&[b, ctx.mul(&[ctx.num(-2), sq_a, t])]);
let x_t = ctx.mul(&[
ctx.add(&[ctx.pow(t, ctx.num(2)), ctx.mul(&[ctx.num(-1), c])]),
ctx.pow(denom, ctx.num(-1)),
]);
let sqrt_t = ctx.add(&[ctx.mul(&[sq_a, x_t]), t]);
let num = ctx.add(&[
ctx.mul(&[ctx.num(2), t, denom]),
ctx.mul(&[
ctx.mul(&[ctx.num(2), sq_a]),
ctx.add(&[ctx.pow(t, ctx.num(2)), ctx.mul(&[ctx.num(-1), c])]),
]),
]);
let dx_dt = ctx.mul(&[num, ctx.pow(ctx.pow(denom, ctx.num(2)), ctx.num(-1))]);
(x_t, sqrt_t, dx_dt)
} else if let Some((r, s)) = crate::integral::rules::rational_sqrt(pc, qc) {
let sq_c = rat(ctx, r, s);
let denom = ctx.add(&[a, ctx.mul(&[ctx.num(-1), ctx.pow(t, ctx.num(2))])]);
let x_t = ctx.mul(&[
ctx.add(&[ctx.mul(&[ctx.num(2), sq_c, t]), ctx.mul(&[ctx.num(-1), b])]),
ctx.pow(denom, ctx.num(-1)),
]);
let sqrt_t = ctx.add(&[sq_c, ctx.mul(&[t, x_t])]);
let num = ctx.add(&[
ctx.mul(&[ctx.num(2), sq_c, denom]),
ctx.mul(&[
ctx.add(&[ctx.mul(&[ctx.num(2), sq_c, t]), ctx.mul(&[ctx.num(-1), b])]),
ctx.mul(&[ctx.num(2), t]),
]),
]);
let dx_dt = ctx.mul(&[num, ctx.pow(ctx.pow(denom, ctx.num(2)), ctx.num(-1))]);
(x_t, sqrt_t, dx_dt)
} else {
return None;
};
let substituted = substitute_atom(ctx, expr, var, x_t)?;
let quad_t = ocas_atom::normalize::normalize(
ctx,
ctx.add(&[
ctx.mul(&[a, ctx.pow(x_t, ctx.num(2))]),
ctx.mul(&[b, x_t]),
c,
]),
);
let substituted = replace_sqrt(ctx, substituted, quad_t, sqrt_t);
let integrand = ctx.mul(&[substituted, dx_dt]);
if node_count(integrand) > MAX_SUBST_NODES {
return None;
}
let t_sym = Symbol::new("_t");
let result_t = crate::integral::rational::integrate_rational(ctx, integrand, t_sym)?;
let x_back = x;
let sqrt_back = if crate::integral::rules::rational_sqrt(pa, qa).is_some() {
let (r, s) = crate::integral::rules::rational_sqrt(pa, qa).unwrap();
let sq_a = rat(ctx, r, s);
let quad = ctx.add(&[
ctx.mul(&[a, ctx.pow(x_back, ctx.num(2))]),
ctx.mul(&[b, x_back]),
c,
]);
let sqrt = ctx.fun("sqrt", &[quad]);
ctx.add(&[sqrt, ctx.mul(&[ctx.num(-1), sq_a, x_back])])
} else {
let (r, s) = crate::integral::rules::rational_sqrt(pc, qc).unwrap();
let sq_c = rat(ctx, r, s);
let quad = ctx.add(&[
ctx.mul(&[a, ctx.pow(x_back, ctx.num(2))]),
ctx.mul(&[b, x_back]),
c,
]);
let sqrt = ctx.fun("sqrt", &[quad]);
ctx.mul(&[
ctx.add(&[sqrt, ctx.mul(&[ctx.num(-1), sq_c])]),
ctx.pow(x_back, ctx.num(-1)),
])
};
let back = substitute_atom(ctx, result_t, t_sym, sqrt_back)?;
Some(ocas_atom::normalize::normalize(ctx, back))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::integrate;
use ocas_atom::AtomArena;
use ocas_core::arena::Arena;
#[test]
fn parts_x_exp() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[x, ctx.fun("exp", &[x])]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
}
#[test]
fn parts_x_sin() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[x, ctx.fun("sin", &[x])]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
}
#[test]
fn parts_x2_sin() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.pow(x, ctx.num(2)), ctx.fun("sin", &[x])]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
}
#[test]
fn parts_log() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("log", &[x]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
}
#[test]
fn parts_x_log() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[x, ctx.fun("log", &[x])]);
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral"), "got {result}");
}
#[test]
fn trig_sub_asin_direct() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let inner = ctx.add(&[ctx.num(1), ctx.mul(&[ctx.num(-1), ctx.pow(x, ctx.num(2))])]);
let half = ctx.pow(ctx.num(2), ctx.num(-1));
let sqrt_expr = ctx.pow(inner, half);
let expr = ctx.pow(sqrt_expr, ctx.num(-1));
let result = integrate(&ctx, expr, Symbol::new("x"));
assert_eq!(result.to_string(), "asin(x)");
}
#[test]
fn trig_sub_sqrt_1_minus_x2() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let inner = ctx.add(&[ctx.num(1), ctx.mul(&[ctx.num(-1), ctx.pow(x, ctx.num(2))])]);
let sqrt_expr = ctx.pow(
inner,
ctx.mul(&[ctx.num(1), ctx.pow(ctx.num(2), ctx.num(-1))]),
);
let result = integrate(&ctx, sqrt_expr, Symbol::new("x"));
let _ = result;
}
#[test]
fn weierstrass_1_over_sin_plus_1() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.pow(ctx.add(&[ctx.fun("sin", &[x]), ctx.num(1)]), ctx.num(-1));
let result = integrate(&ctx, expr, Symbol::new("x"));
let _ = result;
}
#[test]
fn weierstrass_1_over_2_plus_cos() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.pow(ctx.add(&[ctx.num(2), ctx.fun("cos", &[x])]), ctx.num(-1));
let result = integrate(&ctx, expr, Symbol::new("x"));
let _ = result;
}
#[test]
fn liate_scoring() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
assert_eq!(liate_score(&ctx.fun("log", &[x])), 0);
assert_eq!(liate_score(&ctx.fun("asin", &[x])), 1);
assert_eq!(liate_score(&x), 2);
assert_eq!(liate_score(&ctx.fun("sin", &[x])), 3);
assert_eq!(liate_score(&ctx.fun("exp", &[x])), 4);
}
#[test]
fn heuristic_none_for_unknown() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("unknown_func", &[x]);
let result = heuristic_integrate(&ctx, expr, Symbol::new("x"), 0);
assert!(result.is_none());
}
#[test]
fn euler_i_inv_sqrt_quadratic() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, "1/sqrt(x^2+2*x+3)").unwrap();
let result = integrate(&ctx, expr, Symbol::new("x"));
let s = result.to_string();
assert!(!s.contains("Integral("), "got {s}");
assert!(s.contains("asinh") || s.contains("log"), "got {s}");
}
#[test]
fn euler_i_sqrt_x2_plus_1() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, "sqrt(x^2+1)").unwrap();
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral("), "got {result}");
}
#[test]
fn euler_i_inv_x_sqrt_quadratic() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, "1/(x*sqrt(x^2+x+1))").unwrap();
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral("), "got {result}");
}
#[test]
fn euler_ii_inv_sqrt_quadratic() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, "1/sqrt(2*x^2+3*x+1)").unwrap();
let result = integrate(&ctx, expr, Symbol::new("x"));
assert!(!result.to_string().contains("Integral("), "got {result}");
}
fn consts() -> Vec<(Symbol, f64)> {
[
("a", 1.3),
("b", 0.7),
("c", 0.4),
("d", 0.9),
("e", 0.5),
("f", 0.8),
("A", 1.1),
("B", 0.6),
("C", 0.8),
]
.into_iter()
.map(|(n, v)| (Symbol::new(n), v))
.collect()
}
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" => v.tan().recip(),
"sec" => v.cos().recip(),
"csc" => v.sin().recip(),
"sinh" => v.sinh(),
"cosh" => v.cosh(),
"tanh" => v.tanh(),
"coth" => v.tanh().recip(),
"sech" => v.cosh().recip(),
"csch" => v.sinh().recip(),
"log" => v.ln(),
"sqrt" => v.sqrt(),
"atan" => v.atan(),
"atanh" => v.atanh(),
"asin" => v.asin(),
"asinh" => v.asinh(),
_ => return None,
})
}
}
}
fn assert_returns_not_wrong(input: &str) {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let integrand = ocas_parse::parse(&ctx, input).unwrap();
let var = Symbol::new("x");
let result = integrate(&ctx, integrand, var);
let text = result.to_string();
if text.contains("Integral(") {
return;
}
let d = crate::diff(&ctx, result, var);
for &xv in &[0.3f64, 0.7] {
let mut env = consts();
env.push((var, xv));
let lhs = eval_f64(d, &env).expect("eval derivative");
let rhs = eval_f64(integrand, &env).expect("eval integrand");
assert!(
(lhs - rhs).abs() < 1e-5 * rhs.abs().max(1.0),
"{input} at x={xv}: derivative {lhs} vs integrand {rhs} (result: {text})"
);
}
}
const HANG_SHAPES: &[&str] = &[
"cos(c + d*x)^4/(a + b*sin(c + d*x)^3)^2", "sec(c + d*x)^7*(a*cos(c + d*x) + b*sin(c + d*x))^5", "1/(a + b*cos(d + e*x) + c*sin(d + e*x))^3", "sin(c + d*x)^4/(a - b*sin(c + d*x)^4)^3", "(a + a*cos(c + d*x))^4*(A + B*cos(c + d*x) + C*cos(c + d*x)^2)*sec(c + d*x)^7", ];
#[test]
fn corpus_hang_shapes_return() {
for input in HANG_SHAPES {
assert_returns_not_wrong(input);
}
}
#[test]
fn budget_keeps_solving_normal_inputs() {
for (input, expected) in [("1/(2+cos(x))", "atan"), ("1/(2*sin(x)+3*cos(x))", "")] {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let expr = ocas_parse::parse(&ctx, input).unwrap();
let result = heuristic_integrate(&ctx, expr, Symbol::new("x"), 0)
.unwrap_or_else(|| panic!("heuristic declined {input}"));
let s = result.to_string();
assert!(!s.contains("Integral("), "{input} left a residue: {s}");
if !expected.is_empty() {
assert!(s.contains(expected), "{input} -> {s}");
}
}
}
#[test]
fn budget_does_not_leak_across_calls() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let input = "cos(c + d*x)^4/(a + b*sin(c + d*x)^3)^2";
let expr = ocas_parse::parse(&ctx, input).unwrap();
let results: Vec<Option<String>> = (0..8)
.map(|_| heuristic_integrate(&ctx, expr, Symbol::new("x"), 0).map(|a| a.to_string()))
.collect();
let first = &results[0];
for (i, r) in results.iter().enumerate() {
assert_eq!(r, first, "call {i} diverged");
}
}
}