use crate::base::arena::Arena;
use crate::base::node::{ExprId, ExprNode};
use crate::base::numeric::Q;
use crate::base::walk;
use rustc_hash::FxHashMap;
pub(crate) fn trig_combine(arena: &mut Arena, expr: ExprId) -> ExprId {
let post_order = walk::post_order_ids(arena, expr);
let mut cache: FxHashMap<ExprId, ExprId> = FxHashMap::default();
for &id in &post_order {
let rebuilt = crate::base::walk::rebuild_with_cache(arena, id, &cache);
let combined = match arena.node(rebuilt).clone() {
ExprNode::Mul(ref children) => try_combine_mul_trig(arena, rebuilt, children),
ExprNode::Add(ref children) => try_combine_add_trig(arena, rebuilt, children),
_ => rebuilt,
};
let combined = crate::transforms::eval::eval(arena, combined); cache.insert(id, combined);
}
cache.get(&expr).copied().unwrap_or(expr)
}
fn try_combine_mul_trig(arena: &mut Arena, original: ExprId, children: &[ExprId]) -> ExprId {
let mut sin_args: Vec<(usize, ExprId)> = Vec::new();
let mut cos_args: Vec<(usize, ExprId)> = Vec::new();
for (idx, &child) in children.iter().enumerate() {
match arena.node(child).clone() {
ExprNode::Sin(arg) => sin_args.push((idx, arg)),
ExprNode::Cos(arg) => cos_args.push((idx, arg)),
_ => {}
}
}
if let (Some(&(si, a)), Some(&(ci, b))) = (sin_args.first(), cos_args.first()) {
let half = arena.rational(1, 2);
let a_plus_b = arena.add(&[a, b]);
let a_minus_b = arena.sub(a, b);
let sin_sum = arena.sin(a_plus_b);
let sin_diff = arena.sin(a_minus_b);
let inner = arena.add(&[sin_sum, sin_diff]);
let result = arena.mul(&[half, inner]);
return mul_with_remaining(arena, result, children, &[si, ci]);
}
if cos_args.len() >= 2 {
let (i1, a) = cos_args[0];
let (i2, b) = cos_args[1];
let half = arena.rational(1, 2);
let a_plus_b = arena.add(&[a, b]);
let a_minus_b = arena.sub(a, b);
let cos_sum = arena.cos(a_plus_b);
let cos_diff = arena.cos(a_minus_b);
let inner = arena.add(&[cos_diff, cos_sum]);
let result = arena.mul(&[half, inner]);
return mul_with_remaining(arena, result, children, &[i1, i2]);
}
if sin_args.len() >= 2 {
let (i1, a) = sin_args[0];
let (i2, b) = sin_args[1];
let half = arena.rational(1, 2);
let a_plus_b = arena.add(&[a, b]);
let a_minus_b = arena.sub(a, b);
let cos_sum = arena.cos(a_plus_b);
let cos_diff = arena.cos(a_minus_b);
let neg_cos_sum = arena.neg(cos_sum);
let inner = arena.add(&[cos_diff, neg_cos_sum]);
let result = arena.mul(&[half, inner]);
return mul_with_remaining(arena, result, children, &[i1, i2]);
}
original
}
fn mul_with_remaining(
arena: &mut Arena,
result: ExprId,
children: &[ExprId],
used: &[usize],
) -> ExprId {
let remaining: Vec<ExprId> = children
.iter()
.enumerate()
.filter(|&(idx, _)| !used.contains(&idx))
.map(|(_, &c)| c)
.collect();
if remaining.is_empty() {
return result;
}
let mut all = remaining;
all.push(result);
arena.mul(&all)
}
struct TrigSquare {
child_idx: usize,
arg: ExprId,
is_sin: bool,
coeff: Q,
}
fn try_combine_add_trig(arena: &mut Arena, original: ExprId, children: &[ExprId]) -> ExprId {
let two_id = arena.int(2);
let mut squares: Vec<TrigSquare> = Vec::new();
for (idx, &child) in children.iter().enumerate() {
let (coeff, term) = arena.as_coeff_term(child);
let node = arena.node(term).clone();
if let ExprNode::Pow(base, exp) = node {
if exp != two_id {
continue;
}
let base_node = arena.node(base).clone();
match base_node {
ExprNode::Sin(arg) => {
squares.push(TrigSquare {
child_idx: idx,
arg,
is_sin: true,
coeff,
});
}
ExprNode::Cos(arg) => {
squares.push(TrigSquare {
child_idx: idx,
arg,
is_sin: false,
coeff,
});
}
_ => {}
}
}
}
for i in 0..squares.len() {
for j in (i + 1)..squares.len() {
if squares[i].arg != squares[j].arg || squares[i].is_sin == squares[j].is_sin {
continue;
}
let (cos_idx, sin_idx) = if squares[i].is_sin { (j, i) } else { (i, j) };
let cos_coeff = &squares[cos_idx].coeff;
let sin_coeff = &squares[sin_idx].coeff;
let arg = squares[cos_idx].arg;
let cos_child_idx = squares[cos_idx].child_idx;
let sin_child_idx = squares[sin_idx].child_idx;
if *cos_coeff == -sin_coeff {
let two_a = arena.mul(&[two_id, arg]);
let cos_2a = arena.cos(two_a);
let replacement = arena.make_coeff_term(cos_coeff.clone(), cos_2a);
let used_indices = [cos_child_idx, sin_child_idx];
let mut new_children: Vec<ExprId> = children
.iter()
.enumerate()
.filter(|&(idx, _)| !used_indices.contains(&idx))
.map(|(_, &c)| c)
.collect();
new_children.push(replacement);
if new_children.len() == 1 {
return new_children[0];
}
return arena.add(&new_children);
}
if *sin_coeff == -cos_coeff {
let two_a = arena.mul(&[two_id, arg]);
let cos_2a = arena.cos(two_a);
let neg_sin_coeff = -sin_coeff;
let replacement = arena.make_coeff_term(neg_sin_coeff, cos_2a);
let used_indices = [cos_child_idx, sin_child_idx];
let mut new_children: Vec<ExprId> = children
.iter()
.enumerate()
.filter(|&(idx, _)| !used_indices.contains(&idx))
.map(|(_, &c)| c)
.collect();
new_children.push(replacement);
if new_children.len() == 1 {
return new_children[0];
}
return arena.add(&new_children);
}
}
}
original
}
#[cfg(test)]
mod tests {
use super::*;
fn sym(a: &mut Arena, name: &str) -> ExprId {
a.symbol(name)
}
fn display(a: &Arena, id: ExprId) -> String {
a.display(id).to_string()
}
#[test]
fn sin_cos_product_to_sum() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let sin_x = a.sin(x);
let cos_y = a.cos(y);
let product = a.mul(&[sin_x, cos_y]);
let result = trig_combine(&mut a, product);
let s = display(&a, result);
assert!(
s.contains("sin"),
"product-to-sum should produce sin terms: {s}"
);
}
#[test]
fn cos_cos_product_to_sum() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let cos_x = a.cos(x);
let cos_y = a.cos(y);
let product = a.mul(&[cos_x, cos_y]);
let result = trig_combine(&mut a, product);
let s = display(&a, result);
assert!(
s.contains("cos"),
"product-to-sum should produce cos terms: {s}"
);
}
#[test]
fn sin_sin_product_to_sum() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let sin_x = a.sin(x);
let sin_y = a.sin(y);
let product = a.mul(&[sin_x, sin_y]);
let result = trig_combine(&mut a, product);
let s = display(&a, result);
assert!(
s.contains("cos"),
"product-to-sum should produce cos terms: {s}"
);
}
#[test]
fn sin_cos_same_arg_is_half_sin_2x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let sin_x = a.sin(x);
let cos_x = a.cos(x);
let product = a.mul(&[sin_x, cos_x]);
let result = trig_combine(&mut a, product);
let s = display(&a, result);
assert!(s.contains("sin"), "sin(x)*cos(x) should → ½sin(2x): {s}");
}
#[test]
fn product_with_coefficient() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let three = a.int(3);
let sin_x = a.sin(x);
let cos_y = a.cos(y);
let product = a.mul(&[three, sin_x, cos_y]);
let result = trig_combine(&mut a, product);
let s = display(&a, result);
assert!(
s.contains("3") || s.contains("sin"),
"should handle coefficients: {s}"
);
}
#[test]
fn no_trig_unchanged() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let sum = a.add(&[x, y]);
let result = trig_combine(&mut a, sum);
assert_eq!(result, sum, "non-trig should be unchanged");
}
#[test]
fn trig_combine_2sincos_clean() {
let mut arena = Arena::new();
let x = arena.symbol("x");
let two = arena.int(2);
let sin_x = arena.sin(x);
let cos_x = arena.cos(x);
let expr = arena.mul(&[two, sin_x, cos_x]); let result = trig_combine(&mut arena, expr);
let result_str = arena.display(result).to_string();
assert!(
!result_str.contains("sin(0)"),
"trig_combine should not leave sin(0) in result, got: {result_str}"
);
assert!(
result_str.contains("sin(2"),
"trig_combine(2sin(x)cos(x)) should produce sin(2x), got: {result_str}"
);
}
}