use ocas_atom::{Atom, AtomArena, AtomNode};
pub(crate) fn special_integrate<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
x: Atom<'a>,
) -> Option<Atom<'a>> {
let factors: Vec<Atom> = match expr.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![expr],
};
erf_family(ctx, &factors, x)
.or_else(|| ei_family(ctx, &factors, x))
.or_else(|| trig_integral_family(ctx, &factors, x))
.or_else(|| fresnel_family(ctx, &factors, x))
}
fn as_exp(f: Atom) -> Option<Atom> {
if let AtomNode::Fun(name, args) = f.node() {
if name.as_str() == "exp" && args.len() == 1 {
return Some(args[0]);
}
}
None
}
fn is_x_inv<'a>(f: Atom<'a>, x: Atom<'a>) -> bool {
matches!(f.node(), AtomNode::Pow(b, e) if *b == x && matches!(e.node(), AtomNode::Num(-1)))
}
fn as_quadratic<'a>(u: Atom<'a>, x: Atom<'a>) -> Option<Atom<'a>> {
if matches!(u.node(), AtomNode::Pow(b, e) if *b == x && matches!(e.node(), AtomNode::Num(2))) {
return None; }
if let AtomNode::Mul(factors) = u.node() {
if factors.len() == 2 {
if matches!(factors[1].node(), AtomNode::Pow(b, e)
if *b == x && matches!(e.node(), AtomNode::Num(2)))
{
return Some(factors[0]);
}
}
}
None
}
fn is_x_squared<'a>(u: Atom<'a>, x: Atom<'a>) -> bool {
if matches!(u.node(), AtomNode::Pow(b, e) if *b == x && matches!(e.node(), AtomNode::Num(2))) {
return true;
}
if let AtomNode::Pow(b, e) = u.node() {
if matches!(e.node(), AtomNode::Num(2)) {
if let AtomNode::Mul(factors) = b.node() {
let all_neg_one_or_x = factors
.iter()
.all(|f| matches!(f.node(), AtomNode::Num(-1)) || *f == x);
let has_x = factors.contains(&x);
return all_neg_one_or_x && has_x;
}
}
}
false
}
fn is_x<'a>(u: Atom<'a>, x: Atom<'a>) -> bool {
u == x
}
fn erf_family<'a>(ctx: &'a AtomArena<'a>, factors: &[Atom<'a>], x: Atom<'a>) -> Option<Atom<'a>> {
if factors.len() != 1 {
return None;
}
let u = as_exp(factors[0])?;
if is_x_squared(u, x) {
let sqrt_pi = ctx.fun("sqrt", &[ctx.var("pi")]);
let erfi = ctx.fun("erfi", &[x]);
return Some(ctx.mul(&[sqrt_pi, ctx.pow(ctx.num(2), ctx.num(-1)), erfi]));
}
if let Some(c) = as_quadratic(u, x) {
let neg_c = ctx.mul(&[ctx.num(-1), c]);
let root = ctx.fun("sqrt", &[neg_c]);
let sqrt_pi = ctx.fun("sqrt", &[ctx.var("pi")]);
let erf = ctx.fun("erf", &[ctx.mul(&[root, x])]);
return Some(ctx.mul(&[
sqrt_pi,
ctx.pow(ctx.mul(&[ctx.num(2), root]), ctx.num(-1)),
erf,
]));
}
None
}
fn ei_family<'a>(ctx: &'a AtomArena<'a>, factors: &[Atom<'a>], x: Atom<'a>) -> Option<Atom<'a>> {
if factors.len() != 2 {
return None;
}
let (exp_f, inv_f) = if as_exp(factors[0]).is_some() {
(factors[0], factors[1])
} else if as_exp(factors[1]).is_some() {
(factors[1], factors[0])
} else {
return None;
};
if !is_x_inv(inv_f, x) {
return None;
}
let u = as_exp(exp_f)?;
if is_x(u, x) {
return Some(ctx.fun("Ei", &[x]));
}
if let AtomNode::Mul(mf) = u.node() {
if mf.len() == 2 && mf[1] == x && matches!(mf[0].node(), AtomNode::Num(_)) {
return Some(ctx.fun("Ei", &[u]));
}
}
None
}
fn trig_integral_family<'a>(
ctx: &'a AtomArena<'a>,
factors: &[Atom<'a>],
x: Atom<'a>,
) -> Option<Atom<'a>> {
if factors.len() != 2 {
return None;
}
let (fun_f, inv_f) = if is_x_inv(factors[1], x) {
(factors[0], factors[1])
} else if is_x_inv(factors[0], x) {
(factors[1], factors[0])
} else {
return None;
};
let _ = inv_f;
let AtomNode::Fun(name, args) = fun_f.node() else {
return None;
};
if args.len() != 1 || !is_x(args[0], x) {
return None;
}
let target = match name.as_str() {
"sin" => "Si",
"cos" => "Ci",
"sinh" => "Shi",
"cosh" => "Chi",
_ => return None,
};
Some(ctx.fun(target, &[x]))
}
fn fresnel_family<'a>(
ctx: &'a AtomArena<'a>,
factors: &[Atom<'a>],
x: Atom<'a>,
) -> Option<Atom<'a>> {
if factors.len() != 1 {
return None;
}
let AtomNode::Fun(name, args) = factors[0].node() else {
return None;
};
if args.len() != 1 || !is_x_squared(args[0], x) {
return None;
}
let target = match name.as_str() {
"sin" => "fresnels",
"cos" => "fresnelc",
_ => return None,
};
let pi = ctx.var("pi");
let two = ctx.num(2);
let prefactor = ctx.fun("sqrt", &[ctx.mul(&[pi, ctx.pow(two, ctx.num(-1))])]);
let inner = ctx.mul(&[
ctx.fun("sqrt", &[ctx.mul(&[two, ctx.pow(pi, ctx.num(-1))])]),
x,
]);
Some(ctx.mul(&[prefactor, ctx.fun(target, &[inner])]))
}
#[cfg(test)]
mod tests {
use ocas_atom::{AtomArena, Symbol};
use ocas_core::arena::Arena;
use super::*;
fn int<'a>(ctx: &'a AtomArena<'a>, expr: Atom<'a>) -> Option<Atom<'a>> {
special_integrate(ctx, expr, ctx.var("x"))
}
#[test]
fn 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 r = int(&ctx, expr).expect("erf form");
assert!(r.to_string().contains("erf"), "got {r}");
}
#[test]
fn 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 r = int(&ctx, expr).expect("Ei form");
assert_eq!(r.to_string(), "Ei(x)");
}
#[test]
fn sin_x_over_x_gives_si() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.fun("sin", &[x]), ctx.pow(x, ctx.num(-1))]);
let r = int(&ctx, expr).expect("Si form");
assert_eq!(r.to_string(), "Si(x)");
}
#[test]
fn cos_x_over_x_gives_ci() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.mul(&[ctx.fun("cos", &[x]), ctx.pow(x, ctx.num(-1))]);
let r = int(&ctx, expr).expect("Ci form");
assert_eq!(r.to_string(), "Ci(x)");
}
#[test]
fn sin_x_squared_gives_fresnel() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.fun("sin", &[ctx.pow(x, ctx.num(2))]);
let r = int(&ctx, expr).expect("Fresnel form");
assert!(r.to_string().contains("fresnels"), "got {r}");
}
#[test]
fn unmatched_returns_none() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
assert!(int(&ctx, ctx.fun("exp", &[x])).is_none());
let _ = Symbol::new("x");
}
}