use crate::base::arena::Arena;
use crate::base::node::{ExprId, ExprNode};
use crate::base::walk;
use crate::simplify::simplify_engine::count_ops;
use crate::transforms::pattern::{Pattern, Rule};
fn walk_has_node_type(arena: &Arena, expr: ExprId, predicate: impl Fn(&ExprNode) -> bool) -> bool {
let post_order = walk::post_order_ids(arena, expr);
for &id in &post_order {
if predicate(arena.node(id)) {
return true;
}
}
false
}
use rustc_hash::FxHashMap;
pub(crate) fn trigsimp(arena: &mut Arena, expr: ExprId) -> ExprId {
let has_trig = walk_has_node_type(arena, expr, |node| {
matches!(
node,
ExprNode::Sin(_)
| ExprNode::Cos(_)
| ExprNode::Tan(_)
| ExprNode::Asin(_)
| ExprNode::Acos(_)
| ExprNode::Atan(_)
)
});
if !has_trig {
tracing::debug!("trigsimp: skipping all strategies (no trig nodes)");
return expr;
}
let s0 = expr;
let s1 = apply_pattern_rules(arena, expr);
let s2 = replace_sin2_with_1_minus_cos2(arena, expr);
let s3 = replace_cos2_with_1_minus_sin2(arena, expr);
let s4 = strategy_trig_combine(arena, expr);
let s5 = strategy_expand_trig_then_simplify(arena, expr);
let s6 = strategy_trig_identity_rules(arena, expr);
let candidates = [s0, s1, s2, s3, s4, s5, s6];
let _strategy_names = [
"original",
"pattern_rules",
"sin2_to_1_minus_cos2",
"cos2_to_1_minus_sin2",
"trig_combine",
"expand_trig_then_simplify",
"trig_identity_rules",
];
let best_idx = candidates
.iter()
.enumerate()
.min_by_key(|&(_, &e)| count_ops(arena, e))
.map(|(i, _)| i)
.unwrap_or(0);
candidates[best_idx]
}
fn apply_pattern_rules(arena: &mut Arena, expr: ExprId) -> ExprId {
let evaled = crate::transforms::eval::eval(arena, expr);
let rules = crate::transforms::pattern::basic_rules(arena);
let (result, _) = crate::transforms::pattern::apply_rules(arena, evaled, &rules);
result
}
fn replace_sin2_with_1_minus_cos2(arena: &mut Arena, expr: ExprId) -> ExprId {
let replaced = walk_replace_trig_square(arena, expr, TrigKind::Sin);
let evaled = crate::transforms::eval::eval(arena, replaced);
let expanded = crate::transforms::expand::expand(arena, evaled);
let evaled2 = crate::transforms::eval::eval(arena, expanded);
let rules = crate::transforms::pattern::basic_rules(arena);
let (result, _) = crate::transforms::pattern::apply_rules(arena, evaled2, &rules);
result
}
fn replace_cos2_with_1_minus_sin2(arena: &mut Arena, expr: ExprId) -> ExprId {
let replaced = walk_replace_trig_square(arena, expr, TrigKind::Cos);
let evaled = crate::transforms::eval::eval(arena, replaced);
let expanded = crate::transforms::expand::expand(arena, evaled);
let evaled2 = crate::transforms::eval::eval(arena, expanded);
let rules = crate::transforms::pattern::basic_rules(arena);
let (result, _) = crate::transforms::pattern::apply_rules(arena, evaled2, &rules);
result
}
fn strategy_trig_combine(arena: &mut Arena, expr: ExprId) -> ExprId {
let evaled = crate::transforms::eval::eval(arena, expr);
let combined = crate::simplify::trig_combine::trig_combine(arena, evaled);
let evaled2 = crate::transforms::eval::eval(arena, combined);
let rules = crate::transforms::pattern::basic_rules(arena);
let (result, _) = crate::transforms::pattern::apply_rules(arena, evaled2, &rules);
result
}
fn strategy_expand_trig_then_simplify(arena: &mut Arena, expr: ExprId) -> ExprId {
let evaled = crate::transforms::eval::eval(arena, expr);
let expanded = crate::simplify::trig_expand::expand_trig(arena, evaled);
let evaled2 = crate::transforms::eval::eval(arena, expanded);
let rules = crate::transforms::pattern::basic_rules(arena);
let (result, _) = crate::transforms::pattern::apply_rules(arena, evaled2, &rules);
result
}
fn strategy_trig_identity_rules(arena: &mut Arena, expr: ExprId) -> ExprId {
let rules = trig_identity_rules(arena);
let mut current = crate::transforms::eval::eval(arena, expr);
for _ in 0..8 {
let (next, steps) = crate::transforms::pattern::apply_rules(arena, current, &rules);
let next = crate::transforms::eval::eval(arena, next);
if steps.is_empty() || next == current {
break;
}
current = next;
}
current
}
fn rule2(
arena: &mut Arena,
name: &'static str,
lhs: impl Fn(&mut Arena, ExprId, ExprId) -> ExprId,
rhs: impl Fn(&mut Arena, ExprId, ExprId) -> ExprId,
) -> Rule {
let (a, wa) = arena.wild();
let (b, wb) = arena.wild();
let root = lhs(arena, a, b);
let template = rhs(arena, a, b);
let mut wilds = rustc_hash::FxHashMap::default();
wilds.insert(a, wa);
wilds.insert(b, wb);
Rule::new(name, Pattern { root, wilds }, template)
}
fn rule1(
arena: &mut Arena,
name: &'static str,
lhs: impl Fn(&mut Arena, ExprId) -> ExprId,
rhs: impl Fn(&mut Arena, ExprId) -> ExprId,
) -> Rule {
let (a, wa) = arena.wild();
let root = lhs(arena, a);
let template = rhs(arena, a);
let mut wilds = rustc_hash::FxHashMap::default();
wilds.insert(a, wa);
Rule::new(name, Pattern { root, wilds }, template)
}
pub(crate) fn trig_identity_rules(arena: &mut Arena) -> Vec<Rule> {
let two = arena.int(2);
let neg_two = arena.int(-2);
let neg_one = arena.neg_one;
let one = arena.one;
vec![
rule2(
arena,
"sin_add",
|ar, a, b| {
let (sa, cb, ca, sb) = (ar.sin(a), ar.cos(b), ar.cos(a), ar.sin(b));
let t1 = ar.mul(&[sa, cb]);
let t2 = ar.mul(&[ca, sb]);
ar.add(&[t1, t2])
},
|ar, a, b| {
let s = ar.add(&[a, b]);
ar.sin(s)
},
),
rule2(
arena,
"sin_sub",
|ar, a, b| {
let (sa, cb, ca, sb) = (ar.sin(a), ar.cos(b), ar.cos(a), ar.sin(b));
let t1 = ar.mul(&[sa, cb]);
let t2 = ar.mul(&[neg_one, ca, sb]);
ar.add(&[t1, t2])
},
|ar, a, b| {
let d = ar.sub(a, b);
ar.sin(d)
},
),
rule2(
arena,
"cos_add",
|ar, a, b| {
let (ca, cb, sa, sb) = (ar.cos(a), ar.cos(b), ar.sin(a), ar.sin(b));
let t1 = ar.mul(&[ca, cb]);
let t2 = ar.mul(&[neg_one, sa, sb]);
ar.add(&[t1, t2])
},
|ar, a, b| {
let s = ar.add(&[a, b]);
ar.cos(s)
},
),
rule2(
arena,
"cos_sub",
|ar, a, b| {
let (ca, cb, sa, sb) = (ar.cos(a), ar.cos(b), ar.sin(a), ar.sin(b));
let t1 = ar.mul(&[ca, cb]);
let t2 = ar.mul(&[sa, sb]);
ar.add(&[t1, t2])
},
|ar, a, b| {
let d = ar.sub(a, b);
ar.cos(d)
},
),
rule1(
arena,
"cos_double_sq",
|ar, a| {
let (ca, sa) = (ar.cos(a), ar.sin(a));
let c2 = ar.pow(ca, two);
let s2 = ar.pow(sa, two);
let ns2 = ar.mul(&[neg_one, s2]);
ar.add(&[c2, ns2])
},
|ar, a| {
let d = ar.mul(&[two, a]);
ar.cos(d)
},
),
rule1(
arena,
"cos_double_sin",
|ar, a| {
let sa = ar.sin(a);
let s2 = ar.pow(sa, two);
let t = ar.mul(&[neg_two, s2]);
ar.add(&[one, t])
},
|ar, a| {
let d = ar.mul(&[two, a]);
ar.cos(d)
},
),
rule1(
arena,
"cos_double_cos",
|ar, a| {
let ca = ar.cos(a);
let c2 = ar.pow(ca, two);
let t = ar.mul(&[two, c2]);
ar.add(&[neg_one, t])
},
|ar, a| {
let d = ar.mul(&[two, a]);
ar.cos(d)
},
),
rule1(
arena,
"sin_double",
|ar, a| {
let (sa, ca) = (ar.sin(a), ar.cos(a));
ar.mul(&[two, sa, ca])
},
|ar, a| {
let d = ar.mul(&[two, a]);
ar.sin(d)
},
),
rule1(
arena,
"one_minus_cos_sq",
|ar, a| {
let ca = ar.cos(a);
let c2 = ar.pow(ca, two);
let t = ar.mul(&[neg_one, c2]);
ar.add(&[one, t])
},
|ar, a| {
let sa = ar.sin(a);
ar.pow(sa, two)
},
),
rule1(
arena,
"one_minus_sin_sq",
|ar, a| {
let sa = ar.sin(a);
let s2 = ar.pow(sa, two);
let t = ar.mul(&[neg_one, s2]);
ar.add(&[one, t])
},
|ar, a| {
let ca = ar.cos(a);
ar.pow(ca, two)
},
),
rule1(
arena,
"sinh_double",
|ar, a| {
let (sa, ca) = (ar.sinh(a), ar.cosh(a));
ar.mul(&[two, sa, ca])
},
|ar, a| {
let d = ar.mul(&[two, a]);
ar.sinh(d)
},
),
rule1(
arena,
"cosh_double",
|ar, a| {
let (ca, sa) = (ar.cosh(a), ar.sinh(a));
let c2 = ar.pow(ca, two);
let s2 = ar.pow(sa, two);
ar.add(&[c2, s2])
},
|ar, a| {
let d = ar.mul(&[two, a]);
ar.cosh(d)
},
),
rule1(
arena,
"cosh_sinh_sq",
|ar, a| {
let (ca, sa) = (ar.cosh(a), ar.sinh(a));
let c2 = ar.pow(ca, two);
let s2 = ar.pow(sa, two);
let ns2 = ar.mul(&[neg_one, s2]);
ar.add(&[c2, ns2])
},
|ar, _a| ar.one,
),
rule1(
arena,
"sin_div_cos",
|ar, a| {
let (sa, ca) = (ar.sin(a), ar.cos(a));
let inv = ar.pow(ca, neg_one);
ar.mul(&[sa, inv])
},
|ar, a| ar.tan(a),
),
rule1(
arena,
"sinh_div_cosh",
|ar, a| {
let (sa, ca) = (ar.sinh(a), ar.cosh(a));
let inv = ar.pow(ca, neg_one);
ar.mul(&[sa, inv])
},
|ar, a| ar.tanh(a),
),
]
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum TrigKind {
Sin,
Cos,
}
fn walk_replace_trig_square(arena: &mut Arena, expr: ExprId, kind: TrigKind) -> ExprId {
let post_order = walk::post_order_ids(arena, expr);
let mut cache: FxHashMap<ExprId, ExprId> = FxHashMap::default();
for &id in &post_order {
if let Some(replaced) = try_replace_trig_square(arena, id, &cache, kind) {
cache.insert(id, replaced);
continue;
}
let rebuilt = crate::base::walk::rebuild_with_cache(arena, id, &cache);
cache.insert(id, rebuilt);
}
cache.get(&expr).copied().unwrap_or(expr)
}
fn try_replace_trig_square(
arena: &mut Arena,
id: ExprId,
cache: &FxHashMap<ExprId, ExprId>,
kind: TrigKind,
) -> Option<ExprId> {
let node = arena.node(id).clone();
if let ExprNode::Pow(base, exp) = node {
if !is_integer_two(arena, exp) {
return None;
}
let base_node = arena.node(base).clone();
let inner = match (&base_node, kind) {
(ExprNode::Sin(inner), TrigKind::Sin) => Some(*inner),
(ExprNode::Cos(inner), TrigKind::Cos) => Some(*inner),
_ => None,
}?;
let new_inner = cache.get(&inner).copied().unwrap_or(inner);
let other = match kind {
TrigKind::Sin => arena.cos(new_inner),
TrigKind::Cos => arena.sin(new_inner),
};
let two = arena.int(2);
let other_sq = arena.pow(other, two);
let one = arena.one;
let result = arena.sub(one, other_sq);
return Some(result);
}
None
}
fn is_integer_two(arena: &Arena, id: ExprId) -> bool {
if let Some(val) = arena.as_num(id) {
*val == num_rational::Ratio::from_integer(num_bigint::BigInt::from(2))
} else {
false
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sym(arena: &mut Arena, name: &str) -> ExprId {
arena.symbol(name)
}
fn display(arena: &Arena, id: ExprId) -> String {
arena.display(id).to_string()
}
#[test]
fn trigsimp_pythagorean_identity() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let sin_x = arena.sin(x);
let cos_x = arena.cos(x);
let two = arena.int(2);
let sin2 = arena.pow(sin_x, two);
let cos2 = arena.pow(cos_x, two);
let expr = arena.add(&[sin2, cos2]);
let result = trigsimp(&mut arena, expr);
assert_eq!(display(&arena, result), "1");
}
#[test]
fn trigsimp_leaves_simple_trig_alone() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let sin_x = arena.sin(x);
let result = trigsimp(&mut arena, sin_x);
assert_eq!(display(&arena, result), "sin(x)");
}
#[test]
fn trigsimp_pythagorean_plus_constant() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let sin_x = arena.sin(x);
let cos_x = arena.cos(x);
let two = arena.int(2);
let sin2 = arena.pow(sin_x, two);
let cos2 = arena.pow(cos_x, two);
let five = arena.int(5);
let expr = arena.add(&[sin2, cos2, five]);
let result = trigsimp(&mut arena, expr);
assert_eq!(display(&arena, result), "6");
}
#[test]
fn trigsimp_trig_combine_strategy() {
let mut arena = Arena::new();
let x = sym(&mut arena, "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 = trigsimp(&mut arena, expr);
let result_ops = count_ops(&arena, result);
let expr_ops = count_ops(&arena, expr);
assert!(
result_ops <= expr_ops,
"trigsimp should not increase complexity: got {} ops vs original {} ops, result={}",
result_ops,
expr_ops,
display(&arena, result)
);
}
#[test]
fn trigsimp_picks_best_strategy() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let sin_x = arena.sin(x);
let cos_x = arena.cos(x);
let two = arena.int(2);
let sin2 = arena.pow(sin_x, two);
let cos2 = arena.pow(cos_x, two);
let expr = arena.sub(cos2, sin2);
let result = trigsimp(&mut arena, expr);
let result_ops = count_ops(&arena, result);
let expr_ops = count_ops(&arena, expr);
assert!(
result_ops <= expr_ops,
"trigsimp should simplify cos²-sin²: got {} ops, original {} ops, result={}",
result_ops,
expr_ops,
display(&arena, result)
);
}
#[test]
fn identity_rules_sum_difference_double_angle() {
let mut arena = Arena::new();
let (x, y) = (sym(&mut arena, "x"), sym(&mut arena, "y"));
let rules = trig_identity_rules(&mut arena);
let (sx, cx, sy, cy) = (arena.sin(x), arena.cos(x), arena.sin(y), arena.cos(y));
let t1 = arena.mul(&[sx, cy]);
let t2 = arena.mul(&[cx, sy]);
let e = arena.add(&[t1, t2]);
let (r, steps) = crate::transforms::pattern::apply_rules(&mut arena, e, &rules);
assert_eq!(display(&arena, r), "sin(x + y)");
assert_eq!(steps[0].rule_name, "sin_add");
let two = arena.int(2);
let d = arena.mul(&[two, sx, cx]);
let (r, _) = crate::transforms::pattern::apply_rules(&mut arena, d, &rules);
assert_eq!(display(&arena, r), "sin(2*x)");
let sh = arena.sinh(x);
let ch = arena.cosh(x);
let dh = arena.mul(&[two, sh, ch]);
let (r, _) = crate::transforms::pattern::apply_rules(&mut arena, dh, &rules);
assert_eq!(display(&arena, r), "sinh(2*x)");
}
#[test]
fn identity_rules_do_not_fire_on_mismatch() {
let mut arena = Arena::new();
let (x, y, z) = (
sym(&mut arena, "x"),
sym(&mut arena, "y"),
sym(&mut arena, "z"),
);
let rules = trig_identity_rules(&mut arena);
let (sx, cx, sz, cy) = (arena.sin(x), arena.cos(x), arena.sin(z), arena.cos(y));
let t1 = arena.mul(&[sx, cy]);
let t2 = arena.mul(&[cx, sz]);
let e = arena.add(&[t1, t2]);
let (r, steps) = crate::transforms::pattern::apply_rules(&mut arena, e, &rules);
assert_eq!(r, e);
assert!(steps.is_empty());
}
#[test]
fn trigsimp_uses_identity_rules() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let (cx, sx) = (arena.cos(x), arena.sin(x));
let two = arena.int(2);
let c2 = arena.pow(cx, two);
let s2 = arena.pow(sx, two);
let neg_two = arena.int(-2);
let t = arena.mul(&[neg_two, s2]);
let one = arena.one;
let e = arena.add(&[one, t]); let r = trigsimp(&mut arena, e);
assert_eq!(display(&arena, r), "cos(2*x)");
let two_c2 = arena.mul(&[two, c2]);
let neg_one = arena.neg_one;
let f = arena.add(&[two_c2, neg_one]);
let r = trigsimp(&mut arena, f);
assert_eq!(display(&arena, r), "cos(2*x)");
}
}