Skip to main content

ocas_eval/
tree.rs

1//! Owned intermediate representation for expression tree compilation.
2//!
3//! [`EvalTree`] is an arena-free, owned expression tree that can be
4//! constructed from an [`Atom`](ocas_atom::Atom) and then optimized
5//! before instruction generation. It decouples the compilation pipeline
6//! from the arena lifetime.
7
8use ocas_atom::{Atom, AtomNode};
9
10/// An owned intermediate representation of a symbolic expression.
11///
12/// Unlike [`Atom`](ocas_atom::Atom), `EvalTree` owns all its data and
13/// does not depend on an arena. This makes it suitable for multi-pass
14/// compilation and optimization.
15#[derive(Debug, Clone, PartialEq)]
16pub enum EvalTree {
17    /// A numeric constant.
18    Num(f64),
19    /// A variable reference.
20    Var(String),
21    /// A named function applied to arguments.
22    Fun(String, Vec<EvalTree>),
23    /// A sum of terms.
24    Add(Vec<EvalTree>),
25    /// A product of factors.
26    Mul(Vec<EvalTree>),
27    /// A power expression: base^exponent.
28    Pow(Box<EvalTree>, Box<EvalTree>),
29}
30
31impl EvalTree {
32    /// Convert an [`Atom`] into an owned `EvalTree`.
33    ///
34    /// This dissociates the expression from the arena, allowing the
35    /// compilation pipeline to work without lifetime constraints.
36    ///
37    /// # Example
38    ///
39    /// ```ignore
40    /// use ocas_atom::AtomArena;
41    /// use ocas_eval::EvalTree;
42    ///
43    /// let arena = ocas_core::arena::Arena::new();
44    /// let ctx = AtomArena::new(&arena);
45    /// let atom = ctx.add(&[ctx.var("x"), ctx.num(2)]);
46    /// let tree = EvalTree::from_atom(atom);
47    /// ```
48    pub fn from_atom(atom: Atom<'_>) -> Self {
49        match atom.node() {
50            AtomNode::Num(n) => EvalTree::Num(*n as f64),
51            AtomNode::Var(s) => EvalTree::Var(s.as_str().to_string()),
52            AtomNode::Fun(s, args) => {
53                let converted: Vec<EvalTree> =
54                    args.iter().map(|a| EvalTree::from_atom(*a)).collect();
55                EvalTree::Fun(s.as_str().to_string(), converted)
56            }
57            AtomNode::Add(terms) => {
58                let converted: Vec<EvalTree> =
59                    terms.iter().map(|a| EvalTree::from_atom(*a)).collect();
60                EvalTree::Add(converted)
61            }
62            AtomNode::Mul(factors) => {
63                let converted: Vec<EvalTree> =
64                    factors.iter().map(|a| EvalTree::from_atom(*a)).collect();
65                EvalTree::Mul(converted)
66            }
67            AtomNode::Pow(base, exp) => EvalTree::Pow(
68                Box::new(EvalTree::from_atom(*base)),
69                Box::new(EvalTree::from_atom(*exp)),
70            ),
71        }
72    }
73
74    /// Fold constant subtrees and apply algebraic identities.
75    ///
76    /// Rules applied:
77    /// - `Add`: drop `0` terms, sum all-constant terms, collapse single terms
78    /// - `Mul`: absorb `0`, drop `1` factors, multiply all-constant factors,
79    ///   collapse single factors
80    /// - `Pow`: `x^1 → x`, `x^0 → 1`, `Num^Num` evaluated
81    /// - `Fun`: builtin functions of all-constant arguments evaluated
82    ///   (external functions are never folded — they may have side effects)
83    pub fn fold_constants(&self) -> EvalTree {
84        match self {
85            EvalTree::Num(_) | EvalTree::Var(_) => self.clone(),
86            EvalTree::Add(terms) => {
87                let mut folded = Vec::with_capacity(terms.len());
88                let mut const_sum = 0.0f64;
89                let mut has_const = false;
90                for t in terms {
91                    match t.fold_constants() {
92                        EvalTree::Num(n) => {
93                            const_sum += n;
94                            has_const = true;
95                        }
96                        other => folded.push(other),
97                    }
98                }
99                if folded.is_empty() {
100                    return EvalTree::Num(const_sum);
101                }
102                if has_const && const_sum != 0.0 {
103                    folded.push(EvalTree::Num(const_sum));
104                }
105                if folded.len() == 1 {
106                    folded.pop().expect("len checked")
107                } else {
108                    EvalTree::Add(folded)
109                }
110            }
111            EvalTree::Mul(factors) => {
112                let mut folded = Vec::with_capacity(factors.len());
113                let mut const_prod = 1.0f64;
114                let mut has_const = false;
115                for f in factors {
116                    match f.fold_constants() {
117                        EvalTree::Num(n) => {
118                            const_prod *= n;
119                            has_const = true;
120                        }
121                        other => folded.push(other),
122                    }
123                }
124                if has_const && const_prod == 0.0 {
125                    return EvalTree::Num(0.0);
126                }
127                if folded.is_empty() {
128                    return EvalTree::Num(const_prod);
129                }
130                if has_const && const_prod != 1.0 {
131                    folded.push(EvalTree::Num(const_prod));
132                }
133                if folded.len() == 1 {
134                    folded.pop().expect("len checked")
135                } else {
136                    EvalTree::Mul(folded)
137                }
138            }
139            EvalTree::Pow(base, exp) => {
140                let base = base.fold_constants();
141                let exp = exp.fold_constants();
142                match (&base, &exp) {
143                    (EvalTree::Num(b), EvalTree::Num(e)) => EvalTree::Num(b.powf(*e)),
144                    (_, EvalTree::Num(e)) if *e == 0.0 => EvalTree::Num(1.0),
145                    (_, EvalTree::Num(e)) if *e == 1.0 => base,
146                    _ => EvalTree::Pow(Box::new(base), Box::new(exp)),
147                }
148            }
149            EvalTree::Fun(name, args) => {
150                let folded: Vec<EvalTree> = args.iter().map(|a| a.fold_constants()).collect();
151                // Fold builtins with all-constant arguments; external
152                // functions may have side effects and are never folded.
153                if folded.len() == 1
154                    && let (Some(op), EvalTree::Num(x)) =
155                        (crate::instruction::BuiltinOp::from_name(name), &folded[0])
156                {
157                    return EvalTree::Num(apply_builtin_f64(op, *x));
158                }
159                EvalTree::Fun(name.clone(), folded)
160            }
161        }
162    }
163}
164
165/// Apply a builtin operation to an f64 value (used by constant folding).
166fn apply_builtin_f64(op: crate::instruction::BuiltinOp, x: f64) -> f64 {
167    use crate::instruction::BuiltinOp;
168    match op {
169        BuiltinOp::Sin => x.sin(),
170        BuiltinOp::Cos => x.cos(),
171        BuiltinOp::Tan => x.tan(),
172        BuiltinOp::Sec => 1.0 / x.cos(),
173        BuiltinOp::Csc => 1.0 / x.sin(),
174        BuiltinOp::Cot => 1.0 / x.tan(),
175        BuiltinOp::Exp => x.exp(),
176        BuiltinOp::Log => x.ln(),
177        BuiltinOp::Sqrt => x.sqrt(),
178        BuiltinOp::Abs => x.abs(),
179    }
180}
181
182#[cfg(test)]
183mod tests {
184    use super::*;
185    use ocas_atom::AtomArena;
186    use ocas_core::arena::Arena;
187
188    #[test]
189    fn from_atom_num() {
190        let arena = Arena::new();
191        let ctx = AtomArena::new(&arena);
192        let atom = ctx.num(42);
193        let tree = EvalTree::from_atom(atom);
194        assert_eq!(tree, EvalTree::Num(42.0));
195    }
196
197    #[test]
198    fn from_atom_var() {
199        let arena = Arena::new();
200        let ctx = AtomArena::new(&arena);
201        let atom = ctx.var("x");
202        let tree = EvalTree::from_atom(atom);
203        assert_eq!(tree, EvalTree::Var("x".into()));
204    }
205
206    #[test]
207    fn from_atom_add() {
208        let arena = Arena::new();
209        let ctx = AtomArena::new(&arena);
210        let atom = ctx.add(&[ctx.var("x"), ctx.num(2)]);
211        let tree = EvalTree::from_atom(atom);
212        assert_eq!(
213            tree,
214            EvalTree::Add(vec![EvalTree::Var("x".into()), EvalTree::Num(2.0)])
215        );
216    }
217
218    #[test]
219    fn from_atom_pow() {
220        let arena = Arena::new();
221        let ctx = AtomArena::new(&arena);
222        let atom = ctx.pow(ctx.var("x"), ctx.num(2));
223        let tree = EvalTree::from_atom(atom);
224        assert_eq!(
225            tree,
226            EvalTree::Pow(
227                Box::new(EvalTree::Var("x".into())),
228                Box::new(EvalTree::Num(2.0))
229            )
230        );
231    }
232
233    #[test]
234    fn fold_add_all_constants() {
235        let tree = EvalTree::Add(vec![EvalTree::Num(2.0), EvalTree::Num(3.0)]);
236        assert_eq!(tree.fold_constants(), EvalTree::Num(5.0));
237    }
238
239    #[test]
240    fn fold_add_drops_zero() {
241        let tree = EvalTree::Add(vec![EvalTree::Var("x".into()), EvalTree::Num(0.0)]);
242        assert_eq!(tree.fold_constants(), EvalTree::Var("x".into()));
243    }
244
245    #[test]
246    fn fold_add_merges_constants() {
247        let tree = EvalTree::Add(vec![
248            EvalTree::Var("x".into()),
249            EvalTree::Num(2.0),
250            EvalTree::Num(3.0),
251        ]);
252        assert_eq!(
253            tree.fold_constants(),
254            EvalTree::Add(vec![EvalTree::Var("x".into()), EvalTree::Num(5.0)])
255        );
256    }
257
258    #[test]
259    fn fold_mul_absorbs_zero() {
260        let tree = EvalTree::Mul(vec![
261            EvalTree::Var("x".into()),
262            EvalTree::Num(0.0),
263            EvalTree::Var("y".into()),
264        ]);
265        assert_eq!(tree.fold_constants(), EvalTree::Num(0.0));
266    }
267
268    #[test]
269    fn fold_mul_drops_one() {
270        let tree = EvalTree::Mul(vec![EvalTree::Var("x".into()), EvalTree::Num(1.0)]);
271        assert_eq!(tree.fold_constants(), EvalTree::Var("x".into()));
272    }
273
274    #[test]
275    fn fold_pow_one_and_zero() {
276        let x = EvalTree::Var("x".into());
277        assert_eq!(
278            EvalTree::Pow(Box::new(x.clone()), Box::new(EvalTree::Num(1.0))).fold_constants(),
279            x
280        );
281        assert_eq!(
282            EvalTree::Pow(Box::new(x), Box::new(EvalTree::Num(0.0))).fold_constants(),
283            EvalTree::Num(1.0)
284        );
285    }
286
287    #[test]
288    fn fold_pow_constants() {
289        let tree = EvalTree::Pow(Box::new(EvalTree::Num(2.0)), Box::new(EvalTree::Num(10.0)));
290        assert_eq!(tree.fold_constants(), EvalTree::Num(1024.0));
291    }
292
293    #[test]
294    fn fold_builtin_constant() {
295        let tree = EvalTree::Fun("sin".into(), vec![EvalTree::Num(0.0)]);
296        assert_eq!(tree.fold_constants(), EvalTree::Num(0.0));
297    }
298
299    #[test]
300    fn fold_keeps_external_fun() {
301        // Unknown (external) functions are not folded even with constant args
302        let tree = EvalTree::Fun("my_callback".into(), vec![EvalTree::Num(1.0)]);
303        assert_eq!(tree.fold_constants(), tree);
304    }
305
306    #[test]
307    fn fold_nested() {
308        // (2 + 3) * x + 0 → x * 5
309        let tree = EvalTree::Add(vec![
310            EvalTree::Mul(vec![
311                EvalTree::Add(vec![EvalTree::Num(2.0), EvalTree::Num(3.0)]),
312                EvalTree::Var("x".into()),
313            ]),
314            EvalTree::Num(0.0),
315        ]);
316        assert_eq!(
317            tree.fold_constants(),
318            EvalTree::Mul(vec![EvalTree::Var("x".into()), EvalTree::Num(5.0)])
319        );
320    }
321}