use ocas_atom::{Atom, AtomArena, AtomNode};
use ocas_domain::{Domain, Rational, RationalDomain};
use ocas_poly::{Lex, RationalPolynomial, SparseMultivariatePolynomial};
pub type GeneratorField = RationalPolynomial<RationalDomain, Lex>;
type Poly = SparseMultivariatePolynomial<RationalDomain, Lex>;
pub fn atom_to_rational<'a>(atom: Atom<'a>, gens: &[Atom<'a>]) -> Option<GeneratorField> {
atom_to_rational_extended(atom, gens, gens.len())
}
pub(crate) fn atom_to_rational_extended<'a>(
atom: Atom<'a>,
gens: &[Atom<'a>],
n_vars: usize,
) -> Option<GeneratorField> {
debug_assert!(n_vars >= gens.len());
match atom.node() {
AtomNode::Num(n) => Some(constant(*n, n_vars)),
AtomNode::Var(_) | AtomNode::Fun(..) => {
let idx = gens.iter().position(|g| *g == atom)?;
Some(variable(idx, n_vars))
}
AtomNode::Add(args) => {
let mut acc = GeneratorField::zero(&RationalDomain, n_vars);
for a in args.iter() {
acc = acc.add(&atom_to_rational_extended(*a, gens, n_vars)?);
}
Some(acc)
}
AtomNode::Mul(args) => {
let mut acc = GeneratorField::one(&RationalDomain, n_vars);
for a in args.iter() {
acc = acc.mul(&atom_to_rational_extended(*a, gens, n_vars)?);
}
Some(acc)
}
AtomNode::Pow(base, exp) => {
let AtomNode::Num(n) = exp.node() else {
return None;
};
let b = atom_to_rational_extended(*base, gens, n_vars)?;
if *n >= 0 {
Some(pow_u64(&b, *n as u64))
} else {
GeneratorField::one(&RationalDomain, n_vars).div(&pow_u64(&b, n.unsigned_abs()))
}
}
}
}
pub fn rational_to_atom<'a>(
ctx: &'a AtomArena<'a>,
rf: &GeneratorField,
gens: &[Atom<'a>],
) -> Option<Atom<'a>> {
let num = poly_to_atom(ctx, &rf.numerator, gens)?;
if is_const_one(&rf.denominator) {
return Some(num);
}
if is_const_one(&rf.numerator) {
let den = poly_to_atom(ctx, &rf.denominator, gens)?;
return Some(ctx.pow(den, ctx.num(-1)));
}
let den = poly_to_atom(ctx, &rf.denominator, gens)?;
Some(ctx.mul(&[num, ctx.pow(den, ctx.num(-1))]))
}
fn constant(n: i64, n_vars: usize) -> GeneratorField {
if n == 0 {
return GeneratorField::zero(&RationalDomain, n_vars);
}
let poly = Poly::from_terms(
RationalDomain,
n_vars,
vec![(vec![0; n_vars], Rational::new(n, 1))],
);
GeneratorField::from_polynomial(poly)
}
fn variable(idx: usize, n_vars: usize) -> GeneratorField {
let mut exp = vec![0usize; n_vars];
exp[idx] = 1;
let poly = Poly::from_terms(RationalDomain, n_vars, vec![(exp, RationalDomain.one())]);
GeneratorField::from_polynomial(poly)
}
fn pow_u64(base: &GeneratorField, mut exp: u64) -> GeneratorField {
let mut result = GeneratorField::one(&RationalDomain, base.n_vars());
let mut b = base.clone();
while exp > 0 {
if exp & 1 == 1 {
result = result.mul(&b);
}
b = b.mul(&b);
exp >>= 1;
}
result
}
fn is_const_one(p: &Poly) -> bool {
p.n_terms() == 1 && p.domain().is_one(&p.coeff(&vec![0; p.n_vars()]))
}
pub fn rational_const_to_atom<'a>(ctx: &'a AtomArena<'a>, coeff: &Rational) -> Option<Atom<'a>> {
term_to_atom(ctx, &[], coeff, &[])
}
fn poly_to_atom<'a>(ctx: &'a AtomArena<'a>, poly: &Poly, gens: &[Atom<'a>]) -> Option<Atom<'a>> {
let mut terms: Vec<_> = poly.terms_ref().iter().collect();
terms.sort_by(|(e1, _), (e2, _)| e2.iter().rev().cmp(e1.iter().rev()));
let mut sum = Vec::with_capacity(terms.len());
for (exp, coeff) in terms {
sum.push(term_to_atom(ctx, exp, coeff, gens)?);
}
Some(match sum.len() {
0 => ctx.num(0),
1 => sum[0],
_ => ctx.add(&sum),
})
}
fn term_to_atom<'a>(
ctx: &'a AtomArena<'a>,
exp: &[usize],
coeff: &Rational,
gens: &[Atom<'a>],
) -> Option<Atom<'a>> {
let p = coeff.numer().to_i64()?;
let q = coeff.denom().to_i64()?;
let mut factors: Vec<Atom> = Vec::new();
let has_monomial = exp.iter().any(|&e| e > 0);
if q != 1 {
if p != 1 {
factors.push(ctx.num(p));
}
factors.push(ctx.pow(ctx.num(q), ctx.num(-1)));
} else if p != 1 || !has_monomial {
factors.push(ctx.num(p));
}
for (i, &e) in exp.iter().enumerate() {
if e == 0 {
continue;
}
let g = gens[i];
factors.push(if e == 1 {
g
} else {
ctx.pow(g, ctx.num(e as i64))
});
}
Some(match factors.len() {
0 => ctx.num(1),
1 => factors[0],
_ => ctx.mul(&factors),
})
}
#[cfg(test)]
mod tests {
use ocas_atom::AtomArena;
use ocas_core::arena::Arena;
use super::*;
#[test]
fn polynomial_roundtrip() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.add(&[
ctx.pow(x, ctx.num(2)),
ctx.mul(&[ctx.num(2), x]),
ctx.num(1),
]);
let rf = atom_to_rational(expr, &[x]).unwrap();
assert!(is_const_one(&rf.denominator));
let back = rational_to_atom(&ctx, &rf, &[x]).unwrap();
assert_eq!(back.to_string(), "(x^2) + (2*x) + 1");
}
#[test]
fn rational_coefficient_roundtrip() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let half_x = ctx.mul(&[ctx.pow(ctx.num(2), ctx.num(-1)), x]);
let three_over_x = ctx.mul(&[ctx.num(3), ctx.pow(x, ctx.num(-1))]);
let rf = atom_to_rational(ctx.add(&[half_x, three_over_x]), &[x]).unwrap();
let back = rational_to_atom(&ctx, &rf, &[x]).unwrap();
assert_eq!(back.to_string(), "(((2^-1)*(x^2)) + 3)*(x^-1)");
let rf2 = atom_to_rational(back, &[x]).unwrap();
assert_eq!(rf, rf2);
}
#[test]
fn function_generators() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let log_x = ctx.fun("log", &[x]);
let expr = ctx.add(&[ctx.mul(&[x, log_x]), ctx.num(1)]);
let rf = atom_to_rational(expr, &[x, log_x]).unwrap();
let back = rational_to_atom(&ctx, &rf, &[x, log_x]).unwrap();
assert_eq!(back.to_string(), "(x*(log(x))) + 1");
}
#[test]
fn negative_exponent_denominator() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let expr = ctx.pow(ctx.add(&[x, ctx.num(1)]), ctx.num(-2));
let rf = atom_to_rational(expr, &[x]).unwrap();
assert!(is_const_one(&rf.numerator));
let back = rational_to_atom(&ctx, &rf, &[x]).unwrap();
assert_eq!(back.to_string(), "((x^2) + (2*x) + 1)^-1");
}
#[test]
fn rejects_non_rational_input() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
assert!(atom_to_rational(ctx.var("y"), &[x]).is_none());
assert!(atom_to_rational(ctx.fun("sin", &[x]), &[x]).is_none());
assert!(atom_to_rational(ctx.pow(x, ctx.var("y")), &[x]).is_none());
assert!(atom_to_rational(ctx.pow(ctx.num(0), ctx.num(-1)), &[x]).is_none());
}
#[test]
fn zero_and_constants() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let x = ctx.var("x");
let zero = atom_to_rational(ctx.num(0), &[x]).unwrap();
assert!(zero.is_zero());
assert_eq!(
rational_to_atom(&ctx, &zero, &[x]).unwrap().to_string(),
"0"
);
let five = atom_to_rational(ctx.num(5), &[x]).unwrap();
assert_eq!(
rational_to_atom(&ctx, &five, &[x]).unwrap().to_string(),
"5"
);
}
}