use crate::node::Node;
use crate::{CasError, CasErrorKind, Inner};
use cas_domain::Rational;
use std::collections::HashMap;
const MAX_POW_EXPAND: u64 = 10_000;
const MAX_EXPAND_TERMS: u64 = 1_000_000;
enum Info {
Add(Vec<u32>),
Mul(Vec<u32>),
Pow { base: u32, exp: u32 },
Fn { head: u32, args: Vec<u32> },
Atom(u32),
}
fn info_of(inner: &Inner, id: u32) -> Info {
match &inner.nodes[id as usize] {
Node::Add { args: sp } => Info::Add(inner.node_args(*sp).to_vec()),
Node::Mul { args: sp } => Info::Mul(inner.node_args(*sp).to_vec()),
Node::Pow { base, exp } => Info::Pow {
base: *base,
exp: *exp,
},
Node::Fn { head, args: sp } => Info::Fn {
head: *head,
args: inner.node_args(*sp).to_vec(),
},
_ => Info::Atom(id),
}
}
impl Inner {
pub(crate) fn expand_at(&mut self, id: u32, depth: u32) -> u32 {
self.expand_impl(id, depth, true)
}
pub(crate) fn expand_impl(&mut self, id: u32, depth: u32, fast: bool) -> u32 {
assert!(depth <= 10_000, "表达式嵌套过深");
if fast {
if let Info::Mul(_) | Info::Pow { .. } = info_of(self, id) {
if let Some((ring, var_ids, p)) = crate::poly_bridge::to_poly(self, id) {
return crate::poly_bridge::from_poly(self, &ring, &var_ids, &p);
}
}
}
match info_of(self, id) {
Info::Atom(_) => id,
Info::Add(args) => {
let v: Vec<u32> = args
.iter()
.map(|&a| self.expand_impl(a, depth + 1, fast))
.collect();
self.make_add(&v)
}
Info::Mul(args) => {
let v: Vec<u32> = args
.iter()
.map(|&a| self.expand_impl(a, depth + 1, fast))
.collect();
let mut acc = self.sum_terms(v[0]);
for &f in &v[1..] {
let ft = self.sum_terms(f);
let mut next = Vec::with_capacity(acc.len() * ft.len());
for &a in &acc {
for &b in &ft {
next.push(self.make_mul(&[a, b]));
}
}
let s = self.make_add(&next);
acc = self.sum_terms(s);
}
self.make_add(&acc)
}
Info::Pow { base, exp } => {
let k = match &self.nodes[exp as usize] {
Node::Int(v) => v.to_i64(),
_ => None,
};
if let Some(k) = k {
if (2..=MAX_POW_EXPAND as i64).contains(&k) {
let b = self.expand_impl(base, depth + 1, fast);
let bt = self.sum_terms(b);
if bt.len() as u64 * k as u64 <= MAX_EXPAND_TERMS {
let one = self.lit_int(1);
let mut acc = vec![one];
for _ in 0..k {
let mut next = Vec::with_capacity(acc.len() * bt.len());
for &a in &acc {
for &b in &bt {
next.push(self.make_mul(&[a, b]));
}
}
let s = self.make_add(&next);
acc = self.sum_terms(s);
}
return self.make_add(&acc);
}
}
}
let b = self.expand_impl(base, depth + 1, fast);
let e = self.expand_impl(exp, depth + 1, fast);
self.make_pow(b, e)
}
Info::Fn { head, args } => {
let v: Vec<u32> = args
.iter()
.map(|&a| self.expand_impl(a, depth + 1, fast))
.collect();
self.fn_node_by_id(head, &v)
}
}
}
fn sum_terms(&self, id: u32) -> Vec<u32> {
match &self.nodes[id as usize] {
Node::Add { args: sp } => self.node_args(*sp).to_vec(),
_ => vec![id],
}
}
pub(crate) fn cancel_at(&mut self, id: u32) -> u32 {
let args: Vec<u32> = match &self.nodes[id as usize] {
Node::Mul { args: sp } => self.node_args(*sp).to_vec(),
_ => vec![id],
};
let mut coeff = Rational::one();
let mut extras: Vec<u32> = Vec::new(); let mut num: Vec<u32> = Vec::new();
let mut den: Vec<u32> = Vec::new();
enum Piece {
Coeff(Rational),
Extra,
Num,
Den(u32, i64),
}
for &a in &args {
let piece = match &self.nodes[a as usize] {
Node::Int(v) => Piece::Coeff(Rational::from_integer(v)),
Node::Rat(r) => Piece::Coeff(r.clone()),
Node::Float { .. } => Piece::Extra,
Node::Pow { base, exp } => {
let k = match &self.nodes[*exp as usize] {
Node::Int(v) => v.to_i64(),
_ => None,
};
match k {
Some(k) if k > 0 => Piece::Num,
Some(k) => Piece::Den(*base, -k),
None => return id, }
}
_ => Piece::Num,
};
match piece {
Piece::Coeff(v) => coeff = coeff.mul(&v),
Piece::Extra => extras.push(a),
Piece::Num => num.push(a),
Piece::Den(base, m) => {
let lit = self.lit_int(m);
den.push(self.make_pow(base, lit));
}
}
}
if den.is_empty() {
return id; }
let n_id = if num.is_empty() {
self.lit_int(1)
} else {
self.make_mul(&num)
};
let d_id = self.make_mul(&den);
let mut syms: Vec<u32> = Vec::new();
crate::poly_bridge::collect_syms(self, n_id, &mut syms, 0);
crate::poly_bridge::collect_syms(self, d_id, &mut syms, 0);
let Some((ring, var_ids, vi)) = crate::poly_bridge::ring_for(self, &syms) else {
return id;
};
let np = match crate::poly_bridge::to_poly_with(self, n_id, &ring, &vi) {
Some(p) => p,
None => return id,
};
let dp = match crate::poly_bridge::to_poly_with(self, d_id, &ring, &vi) {
Some(p) => p,
None => return id,
};
if dp.is_constant() {
return id; }
if np.is_zero() {
return self.lit_int(0);
}
let g = np.gcd(&dp);
if g.is_constant() {
return id; }
let n2 = np.exact_div(&g).expect("gcd 整除分子");
let d2 = dp.exact_div(&g).expect("gcd 整除分母");
let mut n2 = n2;
if d2.is_constant() {
let cd = d2
.terms()
.next()
.map(|(_, c)| c.clone())
.unwrap_or_else(Rational::one);
coeff = coeff.mul(&cd.inv_reduced().expect("gcd 非零"));
if coeff.is_negative() && !n2.is_constant() {
coeff = coeff.neg();
n2 = n2.neg();
}
}
let mut factors: Vec<u32> = Vec::with_capacity(4);
factors.push(self.lit_rational(&coeff));
factors.extend(extras.iter().copied());
factors.push(crate::poly_bridge::from_poly(self, &ring, &var_ids, &n2));
if !d2.is_constant() {
let d_expr = crate::poly_bridge::from_poly(self, &ring, &var_ids, &d2);
let neg1 = self.lit_int(-1);
factors.push(self.make_pow(d_expr, neg1));
}
self.make_mul(&factors)
}
pub(crate) fn subst_at(&mut self, id: u32, map: &HashMap<u32, u32>, depth: u32) -> u32 {
assert!(depth <= 10_000, "表达式嵌套过深");
match info_of(self, id) {
Info::Atom(aid) => match &self.nodes[aid as usize] {
Node::Sym(s) => map.get(s).copied().unwrap_or(aid),
_ => aid,
},
Info::Add(args) => {
let v: Vec<u32> = args
.iter()
.map(|&a| self.subst_at(a, map, depth + 1))
.collect();
self.make_add(&v)
}
Info::Mul(args) => {
let v: Vec<u32> = args
.iter()
.map(|&a| self.subst_at(a, map, depth + 1))
.collect();
self.make_mul(&v)
}
Info::Pow { base, exp } => {
let b = self.subst_at(base, map, depth + 1);
let e = self.subst_at(exp, map, depth + 1);
self.make_pow(b, e)
}
Info::Fn { head, args } => {
let v: Vec<u32> = args
.iter()
.map(|&a| self.subst_at(a, map, depth + 1))
.collect();
self.fn_node_by_id(head, &v)
}
}
}
}
pub(crate) struct Mono {
pub(crate) exps: Vec<u32>,
pub(crate) coef: Vec<u32>,
}
const MAX_MONO_TERMS: usize = 1_000_000;
const MAX_MONO_DEPTH: u32 = 10_000;
fn one_mono(n: usize) -> Mono {
Mono {
exps: vec![0; n],
coef: Vec::new(),
}
}
fn mul_lists(a: &[Mono], b: &[Mono], n: usize) -> Vec<Mono> {
let mut out = Vec::with_capacity(a.len() * b.len());
for x in a {
for y in b {
let mut exps = x.exps.clone();
for (e, add) in exps.iter_mut().zip(&y.exps) {
*e += *add;
}
let mut coef = x.coef.clone();
coef.extend_from_slice(&y.coef);
let _ = n;
out.push(Mono { exps, coef });
}
}
out
}
fn contains_var(inner: &Inner, id: u32, vars: &[u32], depth: u32) -> Result<bool, CasError> {
if depth > MAX_MONO_DEPTH {
return Err(CasError::new(CasErrorKind::NonPolynomial, "表达式嵌套过深"));
}
Ok(match &inner.nodes[id as usize] {
Node::Sym(s) => vars.contains(s),
Node::Fn { args, .. } => {
let ids: Vec<u32> = inner.node_args(*args).to_vec();
let mut found = false;
for a in ids {
if contains_var(inner, a, vars, depth + 1)? {
found = true;
break;
}
}
found
}
Node::Pow { base, exp } => {
contains_var(inner, *base, vars, depth + 1)?
|| contains_var(inner, *exp, vars, depth + 1)?
}
Node::Mul { args } | Node::Add { args } => {
let ids: Vec<u32> = inner.node_args(*args).to_vec();
let mut found = false;
for a in ids {
if contains_var(inner, a, vars, depth + 1)? {
found = true;
break;
}
}
found
}
Node::Int(_) | Node::Rat(_) | Node::Float { .. } => false,
})
}
pub(crate) fn monomials_at(
inner: &Inner,
id: u32,
vars: &[u32],
depth: u32,
) -> Result<Vec<Mono>, CasError> {
let n = vars.len();
if depth > MAX_MONO_DEPTH {
return Err(CasError::new(CasErrorKind::NonPolynomial, "表达式嵌套过深"));
}
Ok(match &inner.nodes[id as usize] {
Node::Int(_) | Node::Rat(_) | Node::Float { .. } => vec![Mono {
exps: vec![0; n],
coef: vec![id],
}],
Node::Sym(s) => match vars.iter().position(|v| v == s) {
Some(k) => {
let mut exps = vec![0; n];
exps[k] = 1;
vec![Mono { exps, coef: vec![] }]
}
None => vec![Mono {
exps: vec![0; n],
coef: vec![id],
}],
},
Node::Fn { args, .. } => {
let ids: Vec<u32> = inner.node_args(*args).to_vec();
for a in ids {
if contains_var(inner, a, vars, depth + 1)? {
return Err(CasError::new(
CasErrorKind::NonPolynomial,
"函数节点内部含变量",
));
}
}
vec![Mono {
exps: vec![0; n],
coef: vec![id],
}]
}
Node::Pow { base, exp } => {
let k = match &inner.nodes[*exp as usize] {
Node::Int(v) => v.to_i64().ok_or_else(|| {
CasError::new(CasErrorKind::NonPolynomial, "幂指数不是定值整数")
})?,
_ => {
return Err(CasError::new(
CasErrorKind::NonPolynomial,
"幂指数不是整数常量",
));
}
};
if k < 0 {
let is_var =
matches!(&inner.nodes[*base as usize], Node::Sym(s) if vars.contains(s));
if is_var {
return Err(CasError::new(
CasErrorKind::NegativeExponent,
"变量出现负指数",
));
}
return Ok(vec![Mono {
exps: vec![0; n],
coef: vec![id],
}]);
}
if k == 0 {
vec![one_mono(n)]
} else if let Node::Sym(s) = &inner.nodes[*base as usize] {
match vars.iter().position(|v| v == s) {
Some(pos) => {
let mut exps = vec![0; n];
exps[pos] = k as u32;
vec![Mono { exps, coef: vec![] }]
}
None => vec![Mono {
exps: vec![0; n],
coef: vec![id],
}],
}
} else {
let base_monos = monomials_at(inner, *base, vars, depth + 1)?;
let mut acc = vec![one_mono(n)];
for _ in 0..k {
acc = mul_lists(&acc, &base_monos, n);
if acc.len() > MAX_MONO_TERMS {
return Err(CasError::new(
CasErrorKind::NonPolynomial,
"幂展开项数超过守卫上限",
));
}
}
acc
}
}
Node::Mul { args } => {
let ids: Vec<u32> = inner.node_args(*args).to_vec();
let mut acc = vec![one_mono(n)];
for a in ids {
let b = monomials_at(inner, a, vars, depth + 1)?;
acc = mul_lists(&acc, &b, n);
if acc.len() > MAX_MONO_TERMS {
return Err(CasError::new(
CasErrorKind::NonPolynomial,
"乘法展开项数超过守卫上限",
));
}
}
acc
}
Node::Add { args } => {
let ids: Vec<u32> = inner.node_args(*args).to_vec();
let mut out = Vec::new();
for a in ids {
out.extend(monomials_at(inner, a, vars, depth + 1)?);
if out.len() > MAX_MONO_TERMS {
return Err(CasError::new(
CasErrorKind::NonPolynomial,
"加法展开项数超过守卫上限",
));
}
}
out
}
})
}