use num_traits::Signed;
use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use ocas_domain::{Domain, Integer, IntegerDomain, Rational, RationalDomain};
use ocas_poly::{Lex, SparseMultivariatePolynomial};
use crate::tower::convert::{GeneratorField, atom_to_rational, rational_to_atom};
type Sparse = SparseMultivariatePolynomial<ocas_domain::RationalDomain, Lex>;
const MAX_WORK_UNITS: u64 = 4_000_000;
const MAX_COEFF_COST: usize = 20_000;
const MAX_FIELD_DIVISORS: usize = 256;
const MAX_SPLIT_CANDIDATES: usize = 2048;
const MAX_FACTOR_TERMS: usize = 64;
const MAX_SPARSE_TERMS: usize = 1024;
const MAX_STEP_PRODUCT: usize = 100_000;
const MAX_GCD_PRODUCT: usize = 100_000;
thread_local! {
static WORK_UNITS: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
#[cfg(test)]
static PEAK_WORK_UNITS: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
}
fn reset_work_budget() {
WORK_UNITS.with(|c| c.set(0));
}
fn charge_work(units: u64) -> bool {
WORK_UNITS.with(|c| {
let v = c.get().saturating_add(units);
c.set(v);
#[cfg(test)]
PEAK_WORK_UNITS.with(|p| p.set(p.get().max(v)));
v > MAX_WORK_UNITS
})
}
fn budget_exhausted() -> bool {
WORK_UNITS.with(|c| c.get() > MAX_WORK_UNITS)
}
fn fmul(a: &GeneratorField, b: &GeneratorField) -> GeneratorField {
charge_work(
(coeff_cost(a) as u64)
.saturating_mul(coeff_cost(b) as u64)
.max(1),
);
a.mul(b)
}
fn fadd(a: &GeneratorField, b: &GeneratorField) -> GeneratorField {
charge_work((coeff_cost(a) as u64).saturating_add(coeff_cost(b) as u64));
a.add(b)
}
fn fdiv(a: &GeneratorField, b: &GeneratorField) -> Option<GeneratorField> {
charge_work(
(coeff_cost(a) as u64)
.saturating_mul(coeff_cost(b) as u64)
.max(1),
);
a.div(b)
}
#[cfg(test)]
fn peak_work_units() -> u64 {
PEAK_WORK_UNITS.with(|p| p.get())
}
#[cfg(test)]
fn reset_peak_work() {
PEAK_WORK_UNITS.with(|p| p.set(0));
}
#[derive(Clone, Debug)]
struct FPoly {
terms: Vec<(usize, GeneratorField)>,
}
fn rat_const(n: i64, n_vars: usize) -> GeneratorField {
GeneratorField::from_polynomial(Sparse::from_terms(
RationalDomain,
n_vars,
vec![(vec![0; n_vars], Rational::new(n, 1))],
))
}
fn fpoly_one(n_vars: usize) -> FPoly {
FPoly {
terms: vec![(0, rat_const(1, n_vars))],
}
}
fn fpoly_zero(n_vars: usize) -> FPoly {
FPoly {
terms: vec![(0, GeneratorField::zero(&RationalDomain, n_vars))],
}
}
fn mono_reduce(mut g: GeneratorField) -> GeneratorField {
if g.numerator.n_terms() != 1 || g.denominator.n_terms() != 1 {
return g;
}
let (en, cn) = match g.numerator.terms_ref().iter().next() {
Some((e, c)) => (e.clone(), c.clone()),
None => return g,
};
let (ed, cd) = match g.denominator.terms_ref().iter().next() {
Some((e, c)) => (e.clone(), c.clone()),
None => return g,
};
if en.len() != ed.len() {
return g;
}
let common: Vec<usize> = en.iter().zip(&ed).map(|(a, b)| (*a).min(*b)).collect();
if common.iter().all(|&v| v == 0) {
return g;
}
let new_en: Vec<usize> = en.iter().zip(&common).map(|(a, b)| a - b).collect();
let new_ed: Vec<usize> = ed.iter().zip(&common).map(|(a, b)| a - b).collect();
g.numerator = Sparse::from_terms(RationalDomain, en.len(), vec![(new_en, cn)]);
g.denominator = Sparse::from_terms(RationalDomain, ed.len(), vec![(new_ed, cd)]);
g
}
impl FPoly {
fn n_vars(&self) -> usize {
self.terms[0].1.n_vars()
}
fn is_zero(&self) -> bool {
self.terms.iter().all(|(_, c)| c.is_zero())
}
fn degree(&self) -> Option<usize> {
self.terms
.iter()
.rev()
.find(|(_, c)| !c.is_zero())
.map(|(p, _)| *p)
}
fn leading_coeff(&self) -> Option<GeneratorField> {
self.terms
.iter()
.rev()
.find(|(_, c)| !c.is_zero())
.map(|(_, c)| c.clone())
}
fn trim(&mut self) {
let n_vars = self.n_vars();
self.terms.retain(|(_, c)| !c.is_zero());
self.terms.sort_by_key(|(p, _)| *p);
if self.terms.is_empty() {
self.terms
.push((0, GeneratorField::zero(&RationalDomain, n_vars)));
}
}
fn add(&self, other: &Self) -> Self {
let mut out = self.clone();
for (p, c) in &other.terms {
match out.terms.iter_mut().find(|(q, _)| q == p) {
Some((_, acc)) => *acc = fadd(acc, c),
None => out.terms.push((*p, c.clone())),
}
}
out.trim();
out
}
fn neg(&self) -> Self {
Self {
terms: self.terms.iter().map(|(p, c)| (*p, c.neg())).collect(),
}
}
fn sub(&self, other: &Self) -> Self {
self.add(&other.neg())
}
fn mul(&self, other: &Self) -> Self {
let mut out = fpoly_zero(self.n_vars());
for (p, c) in &self.terms {
for (q, d) in &other.terms {
let pow = p + q;
let prod = mono_reduce(fmul(c, d));
match out.terms.iter_mut().find(|(r, _)| *r == pow) {
Some((_, acc)) => *acc = mono_reduce(fadd(acc, &prod)),
None => out.terms.push((pow, prod)),
}
}
}
out.trim();
out
}
fn derivative(&self) -> Self {
let mut out: Vec<(usize, GeneratorField)> = self
.terms
.iter()
.filter(|(p, _)| *p > 0)
.map(|(p, c)| {
(
*p - 1,
mono_reduce(fmul(c, &rat_const(*p as i64, self.n_vars()))),
)
})
.collect();
if out.is_empty() {
out.push((0, GeneratorField::zero(&RationalDomain, self.n_vars())));
}
Self { terms: out }
}
fn to_sparse(&self) -> Sparse {
let n_vars = self.n_vars() + 1;
let mut terms: Vec<(Vec<usize>, _)> = Vec::new();
for (pow, c) in &self.terms {
for (e, coeff) in c.numerator.terms_ref() {
let mut e2 = vec![*pow];
e2.extend_from_slice(e);
terms.push((e2, coeff.clone()));
}
for (e, coeff) in c.denominator.terms_ref() {
let mut e2 = vec![*pow];
e2.extend_from_slice(e);
terms.push((e2, RationalDomain.neg(coeff)));
}
}
Sparse::from_terms(RationalDomain, n_vars, terms)
}
fn from_sparse(p: &Sparse) -> Self {
let nsym = p.n_vars() - 1;
let mut out = fpoly_zero(nsym);
for (e, c) in p.terms_ref() {
let sym = GeneratorField::from_polynomial(Sparse::from_terms(
RationalDomain,
nsym,
vec![(e[1..].to_vec(), c.clone())],
));
out = out.add(&FPoly {
terms: vec![(e[0], sym)],
});
}
out
}
fn div_rem(&self, den: &Self) -> Option<(Self, Self)> {
let mut q = fpoly_zero(self.n_vars());
let mut r = self.clone();
let dd = den.degree()?;
let lc_d = den.leading_coeff()?;
let lc_d_cost = coeff_cost(&lc_d);
let den_cost = fpoly_cost(den);
let mut steps_left = self.degree().map_or(0, |d| d.saturating_sub(dd) + 1);
while let Some(dr) = r.degree() {
if dr < dd {
break;
}
let lc_r = r.leading_coeff()?;
let r_cost = fpoly_cost(&r);
let t_cost = coeff_cost(&lc_r).saturating_add(lc_d_cost);
if steps_left == 0
|| budget_exhausted()
|| r_cost.saturating_mul(t_cost) > MAX_STEP_PRODUCT
|| charge_work(1 + r_cost as u64 + den_cost as u64)
{
return None;
}
steps_left -= 1;
let c = mono_reduce(fdiv(&lc_r, &lc_d)?);
let t = FPoly {
terms: vec![(dr - dd, c)],
};
q = q.add(&t);
r = r.sub(&den.mul(&t));
if fpoly_cost(&r) + fpoly_cost(&q) + den_cost > MAX_COEFF_COST {
return None;
}
}
Some((q, r))
}
fn rem(&self, den: &Self) -> Option<Self> {
self.div_rem(den).map(|(_, r)| r)
}
fn eval(&self, v: &GeneratorField) -> Option<GeneratorField> {
let mut acc = GeneratorField::zero(&RationalDomain, self.n_vars());
let mut prev: Option<usize> = None;
for (exp, c) in self.terms.iter().rev() {
if budget_exhausted() {
return None;
}
let gap = prev.map(|p| p - *exp).unwrap_or(0);
for _ in 0..=gap {
acc = fmul(&acc, v);
}
acc = fadd(&acc, c);
prev = Some(*exp);
}
if let Some(p) = prev {
for _ in 0..p {
acc = fmul(&acc, v);
}
}
Some(acc)
}
fn to_atom<'a>(&self, ctx: &'a AtomArena<'a>, gens: &[Atom<'a>]) -> Option<Atom<'a>> {
let mut terms: Vec<Atom<'a>> = Vec::new();
for (pow, c) in &self.terms {
let c_atom = rational_to_atom(ctx, c, &gens[1..])?;
let x = gens[0];
let xp = if *pow == 0 {
ctx.num(1)
} else {
ctx.pow(x, ctx.num(*pow as i64))
};
terms.push(ctx.mul(&[c_atom, xp]));
}
if terms.is_empty() {
return Some(ctx.num(0));
}
Some(ocas_atom::normalize::normalize(ctx, ctx.add(&terms)))
}
}
fn fpoly_gcd(a: &FPoly, b: &FPoly) -> Option<FPoly> {
if a.is_zero() {
return Some(b.clone());
}
if b.is_zero() {
return Some(a.clone());
}
if fpoly_cost(a).saturating_mul(fpoly_cost(b)) > MAX_GCD_PRODUCT
|| charge_work(1 + fpoly_cost(a) as u64 + fpoly_cost(b) as u64)
{
return None;
}
let mut old_r = a.clone();
let mut r = b.clone();
let mut steps = 0usize;
while !r.is_zero() {
steps += 1;
if steps > 512 || budget_exhausted() || fpoly_cost(&old_r) + fpoly_cost(&r) > MAX_COEFF_COST
{
return None;
}
let (_, rem) = old_r.div_rem(&r)?;
old_r = r;
r = rem;
}
let mut g = old_r;
if let Some(lc) = g.leading_coeff() {
if let Some(inv) = lc.inv() {
g = g.scale(&inv);
}
}
g.trim();
Some(g)
}
fn fpoly_cost(p: &FPoly) -> usize {
p.terms.iter().map(|(_, c)| coeff_cost(c)).sum()
}
fn coeff_cost(c: &GeneratorField) -> usize {
c.numerator.n_terms() + c.denominator.n_terms()
}
impl FPoly {
fn scale(&self, c: &GeneratorField) -> Self {
Self {
terms: self
.terms
.iter()
.map(|(p, q)| (*p, mono_reduce(fmul(q, c))))
.collect(),
}
}
}
fn square_free_factors(p: &FPoly) -> Option<Vec<(FPoly, usize)>> {
if p.degree() == Some(0) {
return Some(Vec::new());
}
let p_prime = p.derivative();
let a0 = fpoly_gcd(p, &p_prime)?;
if a0.degree() == Some(0) {
return Some(vec![(p.clone(), 1)]);
}
let (b1, _) = p.div_rem(&a0)?;
let (c1, _) = p_prime.div_rem(&a0)?;
let mut b = b1;
let mut d = c1.sub(&b.derivative());
let mut result: Vec<(FPoly, usize)> = Vec::new();
let mut i = 1usize;
let mut steps_left = 2 * (p.degree().unwrap_or(0) + 2);
while b.degree() != Some(0) {
if steps_left == 0
|| budget_exhausted()
|| charge_work(1 + fpoly_cost(&b) as u64 + fpoly_cost(&d) as u64)
{
return None;
}
steps_left -= 1;
let ai = fpoly_gcd(&b, &d)?;
let (b_next, _) = b.div_rem(&ai)?;
let (c_next, _) = d.div_rem(&ai)?;
let d_next = c_next.sub(&b_next.derivative());
if ai.degree() != Some(0) {
result.push((ai, i));
}
b = b_next;
d = d_next;
i += 1;
}
if result.is_empty() {
result.push((p.clone(), 1));
}
Some(result)
}
type HermiteParts = (Vec<(FPoly, FPoly)>, FPoly, FPoly);
fn hermite_reduce(num: &FPoly, den: &FPoly) -> Option<HermiteParts> {
let mut b_parts: Vec<(FPoly, FPoly)> = Vec::new();
let mut a = num.clone();
let mut d = den.clone();
let mut steps_left = d.degree().unwrap_or(0) + 1;
loop {
if fpoly_cost(&a) + fpoly_cost(&d) > MAX_COEFF_COST || budget_exhausted() {
return None;
}
if steps_left == 0 || charge_work(1 + fpoly_cost(&a) as u64 + fpoly_cost(&d) as u64) {
return None;
}
steps_left -= 1;
let factors = square_free_factors(&d)?;
let Some((f, m)) = factors.iter().find(|(_, m)| *m >= 2).cloned() else {
break;
};
let mut f_pow = fpoly_one(d.n_vars());
for _ in 0..m {
f_pow = f_pow.mul(&f);
}
let (d1, r) = d.div_rem(&f_pow)?;
if !r.is_zero() {
return None;
}
let f_prime = f.derivative();
let w = f_prime
.scale(&rat_const((m - 1) as i64, d.n_vars()))
.mul(&d1);
let (_, t) = extended_gcd(&f, &w)?;
let b = a.mul(&t).neg().rem(&f)?;
let b_prime = b.derivative();
let inner = b_prime
.mul(&f)
.sub(
&b.mul(&f_prime)
.scale(&rat_const((m - 1) as i64, d.n_vars())),
)
.mul(&d1);
let (c, r2) = a.sub(&inner).div_rem(&f)?;
if !r2.is_zero() {
return None;
}
let mut f_pow_m1 = fpoly_one(d.n_vars());
for _ in 0..(m - 1) {
f_pow_m1 = f_pow_m1.mul(&f);
}
b_parts.push((b, f_pow_m1.clone()));
a = c;
d = d1.mul(&f_pow_m1);
}
Some((b_parts, a, d))
}
fn extended_gcd(a: &FPoly, b: &FPoly) -> Option<(FPoly, FPoly)> {
let mut old_r = a.clone();
let mut r = b.clone();
let mut old_s = fpoly_one(a.n_vars());
let mut s = fpoly_zero(a.n_vars());
let mut old_t = fpoly_zero(a.n_vars());
let mut t = fpoly_one(a.n_vars());
let mut steps_left = a.degree().unwrap_or(0) + b.degree().unwrap_or(0) + 2;
while !r.is_zero() {
if fpoly_cost(&old_r) + fpoly_cost(&r) + fpoly_cost(&s) + fpoly_cost(&t) > MAX_COEFF_COST {
return None;
}
if steps_left == 0
|| budget_exhausted()
|| charge_work(1 + fpoly_cost(&old_r) as u64 + fpoly_cost(&r) as u64)
{
return None;
}
steps_left -= 1;
let (q, rem) = old_r.div_rem(&r)?;
old_r = r;
r = rem;
let new_s = old_s.sub(&q.mul(&s));
old_s = s;
s = new_s;
let new_t = old_t.sub(&q.mul(&t));
old_t = t;
t = new_t;
}
if let Some(lc) = old_r.leading_coeff() {
if let Some(inv) = lc.inv() {
old_s = old_s.scale(&inv);
old_t = old_t.scale(&inv);
}
}
Some((old_s, old_t))
}
fn linear_coeffs(f: &FPoly) -> Option<(GeneratorField, GeneratorField)> {
let a1 = f
.terms
.iter()
.find(|(p, _)| *p == 1)
.map(|(_, c)| c.clone())
.unwrap_or_else(|| GeneratorField::zero(&RationalDomain, f.n_vars()));
let a0 = f
.terms
.iter()
.find(|(p, _)| *p == 0)
.map(|(_, c)| c.clone())
.unwrap_or_else(|| GeneratorField::zero(&RationalDomain, f.n_vars()));
Some((a1, a0))
}
fn collect_symbols(expr: Atom<'_>, var: Symbol, out: &mut Vec<Symbol>) {
match expr.node() {
AtomNode::Var(v) => {
if *v != var && !out.contains(v) {
out.push(*v);
}
}
AtomNode::Num(_) => {}
AtomNode::Fun(_, args) => {
for a in *args {
collect_symbols(*a, var, out);
}
}
AtomNode::Add(args) | AtomNode::Mul(args) => {
for a in *args {
collect_symbols(*a, var, out);
}
}
AtomNode::Pow(base, exp) => {
collect_symbols(*base, var, out);
collect_symbols(*exp, var, out);
}
}
}
fn discriminant(a: &GeneratorField, b: &GeneratorField, c: &GeneratorField) -> GeneratorField {
let four = rat_const(4, a.n_vars());
fmul(&fmul(&four, a), c).sub(&fmul(b, b))
}
fn rational_square_root(delta: &GeneratorField) -> Option<GeneratorField> {
if budget_exhausted() {
return None;
}
let sqrt_sparse = |p: &Sparse| -> Option<Sparse> {
let mut terms: Vec<(Vec<usize>, Rational)> = Vec::new();
for (e, c) in p.terms_ref() {
if e.iter().any(|&v| v % 2 != 0) {
return None;
}
if c.denom().to_i64()? != 1 {
return None;
}
let rp = isqrt_i64(c.numer().to_i64()?)?;
terms.push((e.iter().map(|v| v / 2).collect(), Rational::new(rp, 1)));
}
Some(Sparse::from_terms(RationalDomain, p.n_vars(), terms))
};
let n = sqrt_sparse(&delta.numerator)?;
let d = sqrt_sparse(&delta.denominator)?;
let candidate = GeneratorField::from_num_den(n, d);
if fmul(&candidate, &candidate).sub(delta).is_zero() {
Some(candidate)
} else {
None
}
}
fn factor_via_integer(f: &FPoly) -> Option<Vec<(FPoly, usize)>> {
if charge_work(1 + fpoly_cost(f) as u64) {
return None;
}
let sparse = f.to_sparse();
if sparse.terms_ref().len() > MAX_FACTOR_TERMS {
return None;
}
for e in sparse.terms_ref().keys() {
if e[1..].iter().any(|&v| v != 0) {
return None;
}
}
let mut d: i64 = 1;
for c in sparse.terms_ref().values() {
d = d / gcd_i64(d, c.denom().to_i64()?) * c.denom().to_i64()?;
}
let mut int_terms: Vec<(Vec<usize>, Integer)> = Vec::new();
for (e, c) in sparse.terms_ref() {
let den = c.denom().to_i64()?;
let num = c.numer().to_i64()? * (d / den);
int_terms.push((e.to_vec(), Integer::from(num)));
}
let int_poly: SparseMultivariatePolynomial<IntegerDomain, Lex> =
SparseMultivariatePolynomial::from_terms(IntegerDomain, sparse.n_vars(), int_terms);
let fac = int_poly.factor();
let mut out: Vec<(FPoly, usize)> = Vec::new();
for (g, k) in fac {
let mut terms: Vec<(Vec<usize>, Rational)> = Vec::new();
for (e, c) in g.terms_ref() {
let num = c.to_i64()?;
terms.push((e.to_vec(), Rational::new(num, 1)));
}
let gq = Sparse::from_terms(RationalDomain, sparse.n_vars(), terms);
out.push((FPoly::from_sparse(&gq), k));
}
let _ = d;
if out.is_empty() {
return None;
}
Some(out)
}
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 field_divisors(el: &GeneratorField) -> Option<Vec<GeneratorField>> {
let mut out: Vec<GeneratorField> = Vec::new();
for e in el.numerator.terms_ref().keys() {
let mut exps: Vec<Vec<usize>> = vec![vec![0; e.len()]];
for (i, &v) in e.iter().enumerate() {
if exps.len().saturating_mul(v + 1) > MAX_FIELD_DIVISORS {
return None;
}
let mut next = Vec::new();
for cur in &exps {
for k in 0..=v {
let mut c = cur.clone();
c[i] = k;
next.push(c);
}
}
exps = next;
}
for exp in exps {
if out.len() >= MAX_FIELD_DIVISORS || charge_work(1) {
return None;
}
let p = Sparse::from_terms(
RationalDomain,
el.numerator.n_vars(),
vec![(exp, Rational::new(1, 1))],
);
out.push(GeneratorField::from_polynomial(p));
}
}
for e in el.denominator.terms_ref().keys() {
let mut exps: Vec<Vec<usize>> = vec![vec![0; e.len()]];
for (i, &v) in e.iter().enumerate() {
if exps.len().saturating_mul(v + 1) > MAX_FIELD_DIVISORS {
return None;
}
let mut next = Vec::new();
for cur in &exps {
for k in 0..=v {
let mut c = cur.clone();
c[i] = k;
next.push(c);
}
}
exps = next;
}
for exp in exps {
if out.len() >= MAX_FIELD_DIVISORS || charge_work(1) {
return None;
}
let p = Sparse::from_terms(
RationalDomain,
el.denominator.n_vars(),
vec![(exp, Rational::new(1, 1))],
);
out.push(GeneratorField::from_polynomial(p).inv()?);
}
}
Some(out)
}
fn split_linear_candidates(f: &FPoly) -> Option<Vec<(FPoly, usize)>> {
let lc = f.leading_coeff()?;
let a0 = f
.terms
.iter()
.find(|(p, _)| *p == 0)
.map(|(_, c)| c.clone())
.unwrap_or_else(|| GeneratorField::zero(&RationalDomain, f.n_vars()));
let mut candidates: Vec<GeneratorField> = Vec::new();
if a0.is_zero() {
candidates.push(GeneratorField::zero(&RationalDomain, f.n_vars()));
} else {
let divisors_a0 = field_divisors(&a0)?;
let divisors_lc = field_divisors(&lc)?;
if divisors_a0
.len()
.saturating_mul(divisors_lc.len())
.saturating_mul(2)
> MAX_SPLIT_CANDIDATES
{
return None;
}
for da in divisors_a0 {
for dl in &divisors_lc {
if !dl.is_zero() {
let r = fdiv(&da, dl)?;
candidates.push(r.neg());
candidates.push(r);
}
}
}
}
for r in candidates {
if budget_exhausted() || charge_work(1 + f.degree().unwrap_or(0) as u64) {
return None;
}
if f.eval(&r)?.is_zero() {
let one = GeneratorField::one(&RationalDomain, f.n_vars());
let mut g = FPoly {
terms: vec![(1, one), (0, r.neg())],
};
g.trim();
let (q, rem) = f.div_rem(&g)?;
if rem.is_zero() && q.degree()? >= 1 {
return Some(vec![(g, 1), (q, 1)]);
}
}
}
None
}
fn split_squarefree_factors(factors: Vec<(FPoly, usize)>) -> Option<Vec<(FPoly, usize)>> {
let mut out: Vec<(FPoly, usize)> = Vec::new();
for (f, m) in factors {
if budget_exhausted() || charge_work(1 + fpoly_cost(&f) as u64) {
return None;
}
let deg = f.degree()?;
if deg > 2 && deg <= 3 {
if let Some(sub) = split_linear_candidates(&f) {
let mut sub = split_squarefree_factors(sub)?;
out.append(&mut sub);
continue;
}
if let Some(sub) = factor_via_integer(&f) {
let mut sub = split_squarefree_factors(sub)?;
out.append(&mut sub);
continue;
}
out.push((f, m));
continue;
}
if deg == 2 {
let (a, b, c) = quadratic_coeffs_fpoly(&f)?;
let delta = fmul(&b, &b).sub(&fmul(&fmul(&rat_const(4, f.n_vars()), &a), &c));
if let Some(s) = rational_square_root(&delta) {
let two_a = fmul(&a, &rat_const(2, f.n_vars()));
let r1 = fadd(&b.neg(), &s).div(&two_a)?;
let r2 = b.neg().sub(&s).div(&two_a)?;
let one = GeneratorField::one(&RationalDomain, f.n_vars());
let mut g1 = FPoly {
terms: vec![(1, a.clone()), (0, fmul(&a, &r1).neg())],
};
g1.trim();
let mut g2 = FPoly {
terms: vec![(1, one), (0, r2.neg())],
};
g2.trim();
out.push((g1, m));
out.push((g2, m));
continue;
}
}
out.push((f, m));
}
Some(out)
}
fn isqrt_i64(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)
}
fn constant_rational(delta: &GeneratorField) -> Option<ocas_domain::Rational> {
if delta.numerator.n_terms() == 1 && delta.denominator.n_terms() == 1 {
let (e, c) = delta.numerator.terms_ref().iter().next()?;
if e.iter().all(|&v| v == 0) {
return Some(c.clone());
}
}
None
}
pub(crate) fn rational_complexity_ok<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> bool {
reset_work_budget();
let mut symbols: Vec<Symbol> = Vec::new();
collect_symbols(expr, var, &mut symbols);
if symbols.len() > 5 {
return false;
}
let x = ctx.var(var.as_str());
let mut gens: Vec<Atom<'a>> = vec![x];
for s in &symbols {
gens.push(ctx.var(s.as_str()));
}
let Some(rf) = atom_to_rational(expr, &gens) else {
return false;
};
let sparse_terms = rf.numerator.n_terms() + rf.denominator.n_terms();
if sparse_terms > MAX_SPARSE_TERMS
|| charge_work((sparse_terms as u64).saturating_mul(sparse_terms as u64))
{
return false;
}
let num = FPoly::from_sparse(&rf.numerator);
let mut den = FPoly::from_sparse(&rf.denominator);
if let Some(g) = fpoly_gcd(&num, &den) {
if g.degree() != Some(0) {
if let (Some((_, r1)), Some((dq, r2))) = (num.div_rem(&g), den.div_rem(&g)) {
if r1.is_zero() && r2.is_zero() {
den = dq;
}
}
}
}
den.degree().is_some_and(|d| d <= 6)
}
pub(crate) fn integrate_rational_symbolic<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
reset_work_budget();
integrate_rational_symbolic_inner(ctx, expr, var)
}
fn integrate_rational_symbolic_inner<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var: Symbol,
) -> Option<Atom<'a>> {
let x = ctx.var(var.as_str());
let mut symbols: Vec<Symbol> = Vec::new();
collect_symbols(expr, var, &mut symbols);
symbols.sort_by_key(|s| s.as_str().to_string());
let mut gens: Vec<Atom<'a>> = vec![x];
for s in &symbols {
gens.push(ctx.var(s.as_str()));
}
let rf = atom_to_rational(expr, &gens)?;
let sparse_terms = rf.numerator.n_terms() + rf.denominator.n_terms();
if sparse_terms > MAX_SPARSE_TERMS
|| charge_work((sparse_terms as u64).saturating_mul(sparse_terms as u64))
{
return None;
}
let mut num = FPoly::from_sparse(&rf.numerator);
let mut den = FPoly::from_sparse(&rf.denominator);
let den_deg = den.degree()?;
if den_deg > 6 {
return None;
}
if den_deg > 0 && symbols.len() >= 5 && num.degree().unwrap_or(0) + den_deg > 8 {
return None;
}
if let Some(g) = fpoly_gcd(&num, &den) {
if g.degree() != Some(0) {
if let (Some((nq, r1)), Some((dq, r2))) = (num.div_rem(&g), den.div_rem(&g)) {
if r1.is_zero() && r2.is_zero() {
num = nq;
den = dq;
}
}
}
}
let n_vars = den.n_vars();
let mut parts: Vec<Atom<'a>> = Vec::new();
let (quotient, remainder) = num.div_rem(&den)?;
for (p, c) in quotient.terms {
if p == 0 {
let c_atom = rational_to_atom(ctx, &c, &gens[1..])?;
parts.push(ctx.mul(&[c_atom, x]));
} else {
let p1 = p as i64 + 1;
let c_atom = rational_to_atom(ctx, &c, &gens[1..])?;
parts.push(ctx.mul(&[
c_atom,
ctx.pow(x, ctx.num(p1)),
ctx.pow(ctx.num(p1), ctx.num(-1)),
]));
}
}
let (b_parts, c_num, squarefree_den) = match hermite_reduce(&remainder, &den) {
Some(v) => v,
None => {
return None;
}
};
let mut c_num = c_num;
let mut squarefree_den = squarefree_den;
if let Some(lc) = squarefree_den.leading_coeff() {
if let Some(inv) = lc.inv() {
c_num = c_num.scale(&inv);
squarefree_den = squarefree_den.scale(&inv);
}
}
for (b, d2) in b_parts {
let b_atom = b.to_atom(ctx, &gens)?;
let d2_atom = d2.to_atom(ctx, &gens)?;
parts.push(ctx.mul(&[b_atom, ctx.pow(d2_atom, ctx.num(-1))]));
}
let factors = split_squarefree_factors(square_free_factors(&squarefree_den)?)?;
let has_high_degree = factors.iter().any(|(f, _)| f.degree().unwrap_or(0) > 2);
if has_high_degree {
let integrand = ctx.mul(&[
c_num.to_atom(ctx, &gens)?,
ctx.pow(squarefree_den.to_atom(ctx, &gens)?, ctx.num(-1)),
]);
parts.push(ctx.fun("Integral", &[integrand, x]));
return assemble(ctx, parts);
}
for (f, m) in &factors {
debug_assert_eq!(*m, 1);
if budget_exhausted() {
return None;
}
let f_deg = f.degree()?;
if f_deg == 1 {
let (alpha, beta) = linear_coeffs(f)?;
let r = fdiv(&beta.neg(), &alpha)?;
let mut denom = alpha.clone();
for (g, _) in &factors {
if g.terms == f.terms {
continue;
}
denom = fmul(&denom, &g.eval(&r)?);
}
let coeff = fdiv(&c_num.eval(&r)?, &denom)?;
let coeff_atom = rational_to_atom(ctx, &coeff, &gens[1..])?;
let f_atom = f.to_atom(ctx, &gens)?;
parts.push(ctx.mul(&[coeff_atom, ctx.fun("log", &[f_atom])]));
} else {
let (m, n) = quadratic_coeffs(&c_num, f, &factors)?;
let (a, b, c) = quadratic_coeffs_fpoly(f)?;
let two_a = fmul(&a, &rat_const(2, n_vars));
let m_over = fdiv(&m, &two_a)?;
let f_atom = f.to_atom(ctx, &gens)?;
let m_atom = rational_to_atom(ctx, &m_over, &gens[1..])?;
if !m_over.is_zero() {
parts.push(ctx.mul(&[m_atom, ctx.fun("log", &[f_atom])]));
}
let delta = discriminant(&a, &b, &c);
let mb_over = fdiv(&fmul(&m, &b), &two_a)?;
let n_shift = n.sub(&mb_over);
let fun = match constant_rational(&delta) {
Some(d) if d.inner().is_positive() => "atan",
Some(_) => "atanh",
None => {
let lead = c
.numerator
.terms_ref()
.iter()
.next()
.map(|(_, cc)| cc.clone());
match lead {
Some(cc) if cc.numer().is_negative() => "atanh",
_ => "atan",
}
}
};
let two = if fun == "atanh" { -2 } else { 2 };
let coeff = n_shift.mul(&rat_const(two, n_vars));
let coeff_atom = rational_to_atom(ctx, &coeff, &gens[1..])?;
let mag = if fun == "atanh" {
delta.neg()
} else {
delta.clone()
};
let sqrt_atom: Atom<'a> = if let Some(r) = rational_square_root(&mag) {
rational_to_atom(ctx, &r, &gens[1..])?
} else {
let d_atom = rational_to_atom(ctx, &mag, &gens[1..])?;
ctx.pow(d_atom, ctx.pow(ctx.num(2), ctx.num(-1)))
};
let lin = FPoly {
terms: vec![(1, two_a), (0, b)],
};
let arg_atom = ctx.mul(&[lin.to_atom(ctx, &gens)?, ctx.pow(sqrt_atom, ctx.num(-1))]);
parts.push(ctx.mul(&[
coeff_atom,
ctx.pow(sqrt_atom, ctx.num(-1)),
ctx.fun(fun, &[arg_atom]),
]));
}
}
if let Some(answer) = assemble(ctx, parts) {
if !emission_is_verified(ctx, expr, answer, var, &symbols) {
return None;
}
return Some(answer);
}
None
}
const GUARD_SAMPLES: [f64; 3] = [0.37, 0.83, 1.27];
fn guard_param(i: usize) -> f64 {
const TABLE: [f64; 8] = [2.0, 3.0, 5.0, 7.0, 11.0, 13.0, 0.5, 1.5];
TABLE[i % TABLE.len()]
}
fn eval_num(expr: Atom<'_>, env: &[(Symbol, f64)]) -> Option<f64> {
match expr.node() {
AtomNode::Num(n) => Some(*n as f64),
AtomNode::Var(v) => {
if let Some((_, val)) = env.iter().find(|(s, _)| s == v) {
return Some(*val);
}
match v.as_str() {
"pi" => Some(std::f64::consts::PI),
"e" | "E" => Some(std::f64::consts::E),
_ => None,
}
}
AtomNode::Add(args) => {
let mut acc = 0.0;
for a in args.iter() {
acc += eval_num(*a, env)?;
}
Some(acc)
}
AtomNode::Mul(args) => {
let mut acc = 1.0;
for a in args.iter() {
acc *= eval_num(*a, env)?;
}
Some(acc)
}
AtomNode::Pow(b, e) => {
let (b, e) = (eval_num(*b, env)?, eval_num(*e, env)?);
if e.fract() == 0.0 || b >= 0.0 {
Some(b.powf(e))
} else {
None
}
}
AtomNode::Fun(name, args) => {
let v = eval_num(*args.first()?, env)?;
Some(match name.as_str() {
"log" => v.abs().ln(),
"sqrt" => {
if v < 0.0 {
return None;
}
v.sqrt()
}
"atan" => v.atan(),
"atanh" => {
if v.abs() >= 1.0 {
return None;
}
v.atanh()
}
"asin" => {
if !(-1.0..=1.0).contains(&v) {
return None;
}
v.asin()
}
"abs" => v.abs(),
_ => return None,
})
}
}
}
fn emission_is_verified<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
answer: Atom<'a>,
var: Symbol,
symbols: &[Symbol],
) -> bool {
let base: Vec<(Symbol, f64)> = symbols
.iter()
.enumerate()
.map(|(i, s)| (*s, guard_param(i)))
.collect();
let derivative = crate::diff(ctx, answer, var);
for &x in GUARD_SAMPLES.iter() {
let mut env = base.clone();
env.push((var, x));
let (Some(lhs), Some(rhs)) = (eval_num(derivative, &env), eval_num(expr, &env)) else {
continue;
};
if !lhs.is_finite() || !rhs.is_finite() {
continue;
}
if (lhs - rhs).abs() > 1e-4 * rhs.abs().max(1.0) {
return false;
}
}
true
}
fn assemble<'a>(ctx: &'a AtomArena<'a>, parts: Vec<Atom<'a>>) -> Option<Atom<'a>> {
if parts.is_empty() {
return Some(ctx.num(0));
}
Some(ocas_atom::normalize::normalize(ctx, ctx.add(&parts)))
}
fn quadratic_coeffs_fpoly(f: &FPoly) -> Option<(GeneratorField, GeneratorField, GeneratorField)> {
let get = |p: usize| -> GeneratorField {
f.terms
.iter()
.find(|(q, _)| *q == p)
.map(|(_, c)| c.clone())
.unwrap_or_else(|| GeneratorField::zero(&RationalDomain, f.n_vars()))
};
Some((get(2), get(1), get(0)))
}
fn quadratic_coeffs(
num: &FPoly,
f: &FPoly,
factors: &[(FPoly, usize)],
) -> Option<(GeneratorField, GeneratorField)> {
let mut q = fpoly_one(f.n_vars());
for (g, _) in factors {
if budget_exhausted() {
return None;
}
if g.terms == f.terms {
continue;
}
q = q.mul(g);
}
let (a, b, c) = quadratic_coeffs_fpoly(f)?;
let num_mod = num.rem(f)?;
let q_mod = q.rem(f)?;
let (q1, q0) = linear_coeffs(&q_mod)?;
let (p1, p0) = linear_coeffs(&num_mod)?;
let a_inv = a.inv()?;
let m11 = fadd(&q0, &fmul(&fmul(&q1, &b), &a_inv).neg());
let m21 = fmul(&fmul(&q1, &c), &a_inv).neg();
let det = fadd(&fmul(&m11, &q0), &fmul(&q1, &m21).neg());
let det_inv = det.inv()?;
let m = fmul(&fadd(&fmul(&p1, &q0), &fmul(&q1, &p0).neg()), &det_inv);
let n = fmul(&fadd(&fmul(&m11, &p0), &fmul(&p1, &m21).neg()), &det_inv);
Some((m, n))
}
#[cfg(test)]
mod tests {
use super::*;
use ocas_core::arena::Arena;
fn consts() -> Vec<(Symbol, f64)> {
[
("a", 1.3),
("b", 0.7),
("c", 0.4),
("d", 0.9),
("e", 0.5),
("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(),
_ => return None,
})
}
}
}
fn int_str(input: &str, var: &str) -> String {
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("), "{input} left a residue: {r}");
}
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 = crate::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})"
);
}
}
#[test]
fn symbolic_rationals() {
assert_solved("1/(a+b*x^2)");
assert_solved("1/(a-b*x^2)");
assert_solved("1/(x*(a+b*x)^2)");
assert_solved("1/(x^2*(a-b*x^2))");
assert_solved("(d+e*x)/(x^3*(a+c*x^2))");
assert_solved("(A+B*x)/(a+b*x+c*x^2)");
assert_solved("x^2/(a-b*x^2)^3");
}
#[test]
fn numeric_quadratic_irreducible() {
assert_solved("1/(x^2+2*x+3)");
}
const HANG_SHAPES: &[&str] = &[
"sinh(x)^3/(a + b*sinh(x))", "1/((a - b*x)*(a + b*x)*(c + d*x)^3)", ];
#[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 in [
"1/(a+b*x^2)",
"(d+e*x)/(x^3*(a+c*x^2))",
"(A+B*x)/(a+b*x+c*x^2)",
"1/(x^2*(a-b*x^2))",
] {
assert_solved(input);
}
}
#[test]
fn budget_does_not_leak_across_calls() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let t = Symbol::new("_t");
let expr = ocas_parse::parse(
&ctx,
"(2*((1 + _t^2)^-1))*((a + b*(2*_t*(1 + _t^2)^-1))^-3)",
)
.unwrap();
let first = integrate_rational_symbolic(&ctx, expr, t).map(|a| a.to_string());
assert!(first.is_none(), "expected a decline, got {first:?}");
for i in 0..64 {
let r = integrate_rational_symbolic(&ctx, expr, t).map(|a| a.to_string());
assert_eq!(r, first, "call {i} diverged");
}
}
#[test]
fn symbolic_backend_declines_heavy_t_form() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let t = Symbol::new("_t");
let expr = ocas_parse::parse(
&ctx,
"(2*((1 + _t^2)^-1))*((a + b*(2*_t*(1 + _t^2)^-1))^-3)",
)
.unwrap();
reset_peak_work();
let first = integrate_rational_symbolic(&ctx, expr, t);
let first_peak = peak_work_units();
reset_peak_work();
let second = integrate_rational_symbolic(&ctx, expr, t);
let second_peak = peak_work_units();
if let Some(r) = &first {
assert!(
!r.to_string().contains("Integral("),
"a decline was expected, got {r}"
);
}
assert_eq!(
first.map(|a| a.to_string()),
second.map(|a| a.to_string()),
"two invocations disagreed"
);
assert!(
first_peak.abs_diff(second_peak) <= 8,
"budget leaked between invocations ({first_peak} then {second_peak})"
);
assert!(
first_peak < 4 * MAX_WORK_UNITS && second_peak < 4 * MAX_WORK_UNITS,
"budget did not contain the t-form: {first_peak} then {second_peak} units"
);
}
fn field_of(terms: &[(&[usize], i64)]) -> GeneratorField {
GeneratorField::from_polynomial(Sparse::from_terms(
RationalDomain,
2,
terms
.iter()
.map(|(e, c)| (e.to_vec(), Rational::new(*c, 1)))
.collect(),
))
}
#[test]
fn sum_discriminant_is_not_a_square() {
let delta = field_of(&[(&[2, 0], 4), (&[0, 2], 4)]);
assert!(
rational_square_root(&delta).is_none(),
"4a² + 4b² accepted as a square"
);
let square = field_of(&[(&[2, 0], 4), (&[1, 1], 8), (&[0, 2], 4)]);
assert!(rational_square_root(&square).is_none());
}
#[test]
fn monomial_discriminant_still_splits() {
for (delta, want) in [
(field_of(&[(&[0, 2], 4)]), field_of(&[(&[0, 1], 2)])),
(field_of(&[(&[2, 0], 4)]), field_of(&[(&[1, 0], 2)])),
(field_of(&[(&[2, 0], 1)]), field_of(&[(&[1, 0], 1)])),
] {
let got = rational_square_root(&delta).expect("genuine square declined");
assert!(got.mul(&got).sub(&delta).is_zero(), "√Δ is not a root of Δ");
assert!(
got.sub(&want).is_zero(),
"unexpected root: {} vs {}",
got,
want
);
}
assert!(rational_square_root(&field_of(&[(&[2, 0], 2)])).is_none());
}
#[test]
fn sum_square_discriminant_integrands_are_not_wrong() {
for input in ["1/(b*x^2+2*a*x-b)", "2/(b*x^2+2*a*x-b)", "1/(a+b*sinh(x))"] {
assert_returns_not_wrong(input);
}
}
#[test]
fn monomial_discriminant_integrand_unchanged() {
let r = int_str("1/(b*x^2-b)", "x");
assert!(r.contains("Integral("), "expected the fallback, got {r}");
}
}