Skip to main content

qualia_core_db/specialized_libs/
symbolic_trig.rs

1//! **Trigonometric simplification** (Gap analysis §3.3) — identity-driven rewrites over the
2//! CAS's `Sin`/`Cos`/`Tan` nodes that plain [`simplify`](super::symbolic_algebra::simplify)
3//! cannot do (it knows only constant folding and algebraic identities).
4//!
5//! Implemented identities, each value-preserving for all real inputs:
6//! - **Pythagorean:** `k·sin²(u) + k·cos²(u) → k` (any common scalar `k`, either order),
7//!   and `1 − sin²(u) → cos²(u)`, `1 − cos²(u) → sin²(u)`.
8//! - **Parity:** `sin(−u) → −sin(u)`, `cos(−u) → cos(u)`, `tan(−u) → −tan(u)`.
9//! - **Quotient:** `sin(u)/cos(u) → tan(u)`.
10//!
11//! Applied bottom-up to a bounded fixpoint and interleaved with the base `simplify`, so the
12//! Pythagorean collapse also fires on nested sums after the algebra normalises them.
13
14use super::symbolic_algebra::{add, c, cos, div, mul, neg, pow, simplify, sin, sub, tan, Expr};
15
16/// Simplify trigonometric structure in `expr`, layered on the always-sound base `simplify`.
17pub fn simplify_trig(expr: &Expr) -> Expr {
18    let mut cur = simplify(expr);
19    for _ in 0..16 {
20        let next = simplify(&rewrite(&cur));
21        if next == cur {
22            break;
23        }
24        cur = next;
25    }
26    cur
27}
28
29/// If `e` is `k·trig(u)²` (or `trig(u)²` with `k = 1`), return `(k, is_sin, u)`.
30fn as_scaled_square(e: &Expr) -> Option<(Expr, bool, Expr)> {
31    let classify = |b: &Expr| -> Option<(bool, Expr)> {
32        match b {
33            Expr::Sin(u) => Some((true, (**u).clone())),
34            Expr::Cos(u) => Some((false, (**u).clone())),
35            _ => None,
36        }
37    };
38    match e {
39        Expr::Pow(base, 2) => classify(base).map(|(is_sin, u)| (c(1.0), is_sin, u)),
40        Expr::Mul(a, b) => {
41            if let (Expr::Const(_), Expr::Pow(base, 2)) = (&**a, &**b) {
42                classify(base).map(|(is_sin, u)| ((**a).clone(), is_sin, u))
43            } else if let (Expr::Pow(base, 2), Expr::Const(_)) = (&**a, &**b) {
44                classify(base).map(|(is_sin, u)| ((**b).clone(), is_sin, u))
45            } else {
46                None
47            }
48        }
49        _ => None,
50    }
51}
52
53/// `k·sin²(u) + k·cos²(u) → k` when the two summands share a coefficient and argument.
54fn pythagorean(a: &Expr, b: &Expr) -> Option<Expr> {
55    let (ka, sin_a, ua) = as_scaled_square(a)?;
56    let (kb, sin_b, ub) = as_scaled_square(b)?;
57    if ka == kb && ua == ub && sin_a != sin_b {
58        Some(ka)
59    } else {
60        None
61    }
62}
63
64fn rewrite(e: &Expr) -> Expr {
65    // Bottom-up.
66    let e = match e {
67        Expr::Add(a, b) => add(rewrite(a), rewrite(b)),
68        Expr::Sub(a, b) => sub(rewrite(a), rewrite(b)),
69        Expr::Mul(a, b) => mul(rewrite(a), rewrite(b)),
70        Expr::Div(a, b) => div(rewrite(a), rewrite(b)),
71        Expr::Pow(a, n) => pow(rewrite(a), *n),
72        Expr::Neg(a) => neg(rewrite(a)),
73        Expr::Sqrt(a) => Expr::Sqrt(Box::new(rewrite(a))),
74        Expr::Exp(a) => Expr::Exp(Box::new(rewrite(a))),
75        Expr::Ln(a) => Expr::Ln(Box::new(rewrite(a))),
76        Expr::Sin(a) => sin(rewrite(a)),
77        Expr::Cos(a) => cos(rewrite(a)),
78        Expr::Tan(a) => tan(rewrite(a)),
79        Expr::Const(_) | Expr::Var(_) => e.clone(),
80    };
81
82    match &e {
83        // Parity.
84        Expr::Sin(u) => {
85            if let Expr::Neg(inner) = &**u {
86                return neg(sin((**inner).clone()));
87            }
88            e
89        }
90        Expr::Cos(u) => {
91            if let Expr::Neg(inner) = &**u {
92                return cos((**inner).clone());
93            }
94            e
95        }
96        Expr::Tan(u) => {
97            if let Expr::Neg(inner) = &**u {
98                return neg(tan((**inner).clone()));
99            }
100            e
101        }
102        // sin(u)/cos(u) → tan(u).
103        Expr::Div(n, d) => {
104            if let (Expr::Sin(un), Expr::Cos(ud)) = (&**n, &**d) {
105                if un == ud {
106                    return tan((**un).clone());
107                }
108            }
109            e
110        }
111        // Pythagorean collapse on a sum (either order).
112        Expr::Add(a, b) => {
113            if let Some(r) = pythagorean(a, b).or_else(|| pythagorean(b, a)) {
114                return r;
115            }
116            e
117        }
118        // 1 − sin²(u) → cos²(u) ; 1 − cos²(u) → sin²(u).
119        Expr::Sub(a, b) => {
120            if let Expr::Const(one) = &**a {
121                if *one == 1.0 {
122                    if let Expr::Pow(base, 2) = &**b {
123                        match &**base {
124                            Expr::Sin(u) => return pow(cos((**u).clone()), 2),
125                            Expr::Cos(u) => return pow(sin((**u).clone()), 2),
126                            _ => {}
127                        }
128                    }
129                }
130            }
131            e
132        }
133        _ => e,
134    }
135}
136
137#[cfg(test)]
138mod tests {
139    use super::super::symbolic_algebra::{var, Expr};
140    use super::*;
141    use std::collections::HashMap;
142
143    fn val(e: &Expr, x: f64) -> f64 {
144        let mut env = HashMap::new();
145        env.insert("x".to_string(), x);
146        e.eval(&env).unwrap()
147    }
148
149    #[test]
150    fn pythagorean_identity_collapses() {
151        // sin²(x) + cos²(x) → 1.
152        let e = add(pow(sin(var("x")), 2), pow(cos(var("x")), 2));
153        assert_eq!(simplify_trig(&e), c(1.0));
154        // cos²(x) + sin²(x) → 1 (other order).
155        let e2 = add(pow(cos(var("x")), 2), pow(sin(var("x")), 2));
156        assert_eq!(simplify_trig(&e2), c(1.0));
157        // 3·sin²(x) + 3·cos²(x) → 3.
158        let e3 = add(
159            mul(c(3.0), pow(sin(var("x")), 2)),
160            mul(c(3.0), pow(cos(var("x")), 2)),
161        );
162        assert_eq!(simplify_trig(&e3), c(3.0));
163    }
164
165    #[test]
166    fn pythagorean_in_a_larger_sum() {
167        // sin²(x) + cos²(x) + x  →  1 + x (value-checked; structure may be (1 + x)).
168        let e = add(add(pow(sin(var("x")), 2), pow(cos(var("x")), 2)), var("x"));
169        let s = simplify_trig(&e);
170        for &x in &[0.3, 1.2, 2.7] {
171            assert!((val(&s, x) - (1.0 + x)).abs() < 1e-9);
172        }
173    }
174
175    #[test]
176    fn one_minus_square() {
177        // 1 − sin²(x) → cos²(x) (value-checked).
178        let s = simplify_trig(&sub(c(1.0), pow(sin(var("x")), 2)));
179        for &x in &[0.3, 1.2, 2.7] {
180            assert!((val(&s, x) - x.cos().powi(2)).abs() < 1e-9);
181        }
182    }
183
184    #[test]
185    fn parity_and_quotient() {
186        // cos(−x) → cos(x) ; sin(−x) → −sin(x).
187        assert_eq!(simplify_trig(&cos(neg(var("x")))), cos(var("x")));
188        assert_eq!(simplify_trig(&sin(neg(var("x")))), neg(sin(var("x"))));
189        // sin(x)/cos(x) → tan(x).
190        assert_eq!(
191            simplify_trig(&div(sin(var("x")), cos(var("x")))),
192            tan(var("x"))
193        );
194    }
195}