1use ocas_atom::{Atom, AtomNode};
9
10#[derive(Debug, Clone, PartialEq)]
16pub enum EvalTree {
17 Num(f64),
19 Var(String),
21 Fun(String, Vec<EvalTree>),
23 Add(Vec<EvalTree>),
25 Mul(Vec<EvalTree>),
27 Pow(Box<EvalTree>, Box<EvalTree>),
29}
30
31impl EvalTree {
32 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 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 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
165fn 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 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 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}