1use std::collections::HashSet;
11
12use ocas_atom::Atom;
13use ocas_core::FastHashMap as HashMap;
14
15use crate::domain::{EvaluationDomain, PowfExtension};
16use crate::error::{EvaluationError, Result};
17use crate::evaluator::ExpressionEvaluator;
18use crate::function_map::FunctionMap;
19use crate::instruction::Instr;
20use crate::optimize;
21use crate::tree::EvalTree;
22
23pub fn compile_atom<T: EvaluationDomain + PowfExtension>(
25 atom: Atom<'_>,
26) -> Result<ExpressionEvaluator<T>> {
27 compile_atom_with(atom, None)
28}
29
30pub fn compile_atom_with<T: EvaluationDomain + PowfExtension>(
32 atom: Atom<'_>,
33 function_map: Option<FunctionMap<T>>,
34) -> Result<ExpressionEvaluator<T>> {
35 let tree = EvalTree::from_atom(atom);
36 compile_tree_with(&tree, function_map)
37}
38
39#[allow(dead_code)]
41pub fn compile_tree<T: EvaluationDomain + PowfExtension>(
42 tree: &EvalTree,
43) -> Result<ExpressionEvaluator<T>> {
44 compile_tree_with(tree, None)
45}
46
47pub fn compile_tree_with<T: EvaluationDomain + PowfExtension>(
49 tree: &EvalTree,
50 function_map: Option<FunctionMap<T>>,
51) -> Result<ExpressionEvaluator<T>> {
52 compile_trees(&[tree], function_map)
53}
54
55pub fn compile_atoms_multi<T: EvaluationDomain + PowfExtension>(
59 atoms: &[Atom<'_>],
60) -> Result<ExpressionEvaluator<T>> {
61 compile_atoms_multi_with(atoms, None)
62}
63
64pub fn compile_atoms_multi_with<T: EvaluationDomain + PowfExtension>(
67 atoms: &[Atom<'_>],
68 function_map: Option<FunctionMap<T>>,
69) -> Result<ExpressionEvaluator<T>> {
70 let trees: Vec<EvalTree> = atoms.iter().map(|a| EvalTree::from_atom(*a)).collect();
71 let refs: Vec<&EvalTree> = trees.iter().collect();
72 compile_trees(&refs, function_map)
73}
74
75pub fn compile_trees_multi<T: EvaluationDomain + PowfExtension>(
78 trees: &[&EvalTree],
79) -> Result<ExpressionEvaluator<T>> {
80 compile_trees(trees, None)
81}
82
83fn compile_trees<T: EvaluationDomain + PowfExtension>(
86 trees: &[&EvalTree],
87 function_map: Option<FunctionMap<T>>,
88) -> Result<ExpressionEvaluator<T>> {
89 let folded: Vec<EvalTree> = trees.iter().map(|t| t.fold_constants()).collect();
91
92 let mut var_names = HashSet::new();
94 let mut const_count = 0usize;
95 for tree in &folded {
96 scan_tree(tree, &mut var_names, &mut const_count);
97 }
98 let param_count = var_names.len();
99
100 let mut sorted_vars: Vec<String> = var_names.into_iter().collect();
102 sorted_vars.sort();
103 let var_to_param: HashMap<String, usize> = sorted_vars
104 .iter()
105 .enumerate()
106 .map(|(i, v)| (v.clone(), i))
107 .collect();
108
109 let temp_base = param_count + const_count;
112 let (instructions, next_temp, constants, result_slots) = {
113 let mut ctx =
114 CompileContext::<T>::new(param_count, temp_base, var_to_param, function_map.as_ref());
115 let mut result_slots = Vec::with_capacity(folded.len());
116 for tree in &folded {
117 result_slots.push(ctx.compile_node(tree)?);
118 }
119 (ctx.instructions, ctx.next_temp, ctx.constants, result_slots)
120 };
121
122 let actual_const_count = constants.len();
123 let (instructions, temp_count, result_indices) =
124 optimize::optimize(instructions, temp_base, next_temp, &result_slots);
125 let stack_size = temp_base + temp_count;
126
127 match function_map {
128 Some(fm) => Ok(ExpressionEvaluator::new_with_functions(
129 instructions,
130 param_count,
131 actual_const_count,
132 stack_size,
133 result_indices,
134 constants,
135 fm,
136 )),
137 None => Ok(ExpressionEvaluator::new(
138 instructions,
139 param_count,
140 actual_const_count,
141 stack_size,
142 result_indices,
143 constants,
144 )),
145 }
146}
147
148fn scan_tree(tree: &EvalTree, vars: &mut HashSet<String>, const_count: &mut usize) {
150 match tree {
151 EvalTree::Num(_) => {
152 *const_count += 1;
153 }
154 EvalTree::Var(name) => {
155 vars.insert(name.clone());
156 }
157 EvalTree::Add(terms) | EvalTree::Mul(terms) => {
158 for t in terms {
159 scan_tree(t, vars, const_count);
160 }
161 }
162 EvalTree::Pow(base, exp) => {
163 scan_tree(base, vars, const_count);
164 scan_tree(exp, vars, const_count);
165 }
166 EvalTree::Fun(_, args) => {
167 for a in args {
168 scan_tree(a, vars, const_count);
169 }
170 }
171 }
172}
173
174struct CompileContext<'a, T: EvaluationDomain> {
175 instructions: Vec<Instr>,
176 next_temp: usize,
178 param_count: usize,
180 temp_base: usize,
182 variables: HashMap<String, usize>,
184 constants: Vec<T>,
186 function_map: Option<&'a FunctionMap<T>>,
188}
189
190impl<'a, T: EvaluationDomain> CompileContext<'a, T> {
191 fn new(
192 param_count: usize,
193 temp_base: usize,
194 variables: HashMap<String, usize>,
195 function_map: Option<&'a FunctionMap<T>>,
196 ) -> Self {
197 Self {
198 instructions: Vec::new(),
199 next_temp: 0,
200 param_count,
201 temp_base,
202 variables,
203 constants: Vec::new(),
204 function_map,
205 }
206 }
207
208 fn alloc_temp(&mut self) -> usize {
209 let slot = self.next_temp + self.temp_base;
210 self.next_temp += 1;
211 slot
212 }
213
214 fn param_slot(&self, name: &str) -> usize {
215 self.variables[name]
216 }
217
218 fn const_slot(&mut self, value: T) -> usize {
219 let idx = self.constants.len();
220 self.constants.push(value);
221 self.param_count + idx
222 }
223
224 fn compile_node(&mut self, node: &EvalTree) -> Result<usize> {
225 match node {
226 EvalTree::Num(n) => {
227 let dst = self.alloc_temp();
228 let const_slot = self.const_slot(T::from_f64(*n));
229 self.instructions.push(Instr::Copy {
230 dst,
231 src: const_slot,
232 });
233 Ok(dst)
234 }
235 EvalTree::Var(name) => {
236 let dst = self.alloc_temp();
237 let param_slot = self.param_slot(name);
238 self.instructions.push(Instr::Copy {
239 dst,
240 src: param_slot,
241 });
242 Ok(dst)
243 }
244 EvalTree::Add(terms) => {
245 let dst = self.alloc_temp();
246 let mut srcs = Vec::with_capacity(terms.len());
247 for term in terms {
248 srcs.push(self.compile_node(term)?);
249 }
250 self.instructions.push(Instr::Add { dst, srcs });
251 Ok(dst)
252 }
253 EvalTree::Mul(factors) => {
254 let dst = self.alloc_temp();
255 let mut srcs = Vec::with_capacity(factors.len());
256 for factor in factors {
257 srcs.push(self.compile_node(factor)?);
258 }
259 self.instructions.push(Instr::Mul { dst, srcs });
260 Ok(dst)
261 }
262 EvalTree::Pow(base, exp) => {
263 let base_slot = self.compile_node(base)?;
264 let dst = self.alloc_temp();
265 if let EvalTree::Num(n) = exp.as_ref()
266 && n.fract() == 0.0
267 && *n >= i64::MIN as f64
268 && *n <= i64::MAX as f64
269 {
270 self.instructions.push(Instr::Pow {
271 dst,
272 base: base_slot,
273 exp: *n as i64,
274 });
275 return Ok(dst);
276 }
277 let exp_slot = self.compile_node(exp)?;
278 self.instructions.push(Instr::Powf {
279 dst,
280 base: base_slot,
281 exp: exp_slot,
282 });
283 Ok(dst)
284 }
285 EvalTree::Fun(name, args) => {
286 if is_builtin(name) && args.len() == 1 {
287 let arg_slot = self.compile_node(&args[0])?;
288 let dst = self.alloc_temp();
289 let op = crate::instruction::BuiltinOp::from_name(name)
290 .expect("is_builtin guarantees known name");
291 self.instructions.push(Instr::BuiltinOp {
292 dst,
293 op,
294 src: arg_slot,
295 });
296 Ok(dst)
297 } else if let Some(fm) = self.function_map {
298 if let Some(_entry) = fm.resolve(name) {
300 let mut srcs = Vec::with_capacity(args.len());
301 for arg in args {
302 srcs.push(self.compile_node(arg)?);
303 }
304 let dst = self.alloc_temp();
305 let fn_idx =
307 fm.index_of(name)
308 .ok_or_else(|| EvaluationError::FunctionNotFound {
309 name: name.clone(),
310 })?;
311 self.instructions
312 .push(Instr::ExternalFun { dst, fn_idx, srcs });
313 Ok(dst)
314 } else {
315 Err(EvaluationError::FunctionNotFound { name: name.clone() })
316 }
317 } else {
318 Err(EvaluationError::FunctionNotFound { name: name.clone() })
319 }
320 }
321 }
322 }
323}
324
325fn is_builtin(name: &str) -> bool {
326 matches!(
327 name.to_lowercase().as_str(),
328 "sin" | "cos" | "tan" | "sec" | "csc" | "cot" | "exp" | "log" | "sqrt" | "abs"
329 )
330}
331
332impl<T: EvaluationDomain + PowfExtension> ExpressionEvaluator<T> {
337 pub fn compile(atom: Atom<'_>) -> Result<Self> {
339 compile_atom(atom)
340 }
341
342 pub fn compile_with(atom: Atom<'_>, map: FunctionMap<T>) -> Result<Self> {
344 compile_atom_with(atom, Some(map))
345 }
346
347 pub fn compile_multi(atoms: &[Atom<'_>]) -> Result<Self> {
350 compile_atoms_multi(atoms)
351 }
352
353 pub fn compile_multi_with(atoms: &[Atom<'_>], map: FunctionMap<T>) -> Result<Self> {
356 compile_atoms_multi_with(atoms, Some(map))
357 }
358}
359
360#[cfg(test)]
361mod tests {
362 use super::*;
363 use ocas_atom::AtomArena;
364 use ocas_core::arena::Arena;
365
366 #[test]
367 fn compile_constant() {
368 let arena = Arena::new();
369 let ctx = AtomArena::new(&arena);
370 let expr = ctx.num(42);
371 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
372 let result = eval.evaluate(&[]).unwrap();
373 assert!((result[0] - 42.0).abs() < 1e-10);
374 }
375
376 #[test]
377 fn compile_single_var() {
378 let arena = Arena::new();
379 let ctx = AtomArena::new(&arena);
380 let expr = ctx.var("x");
381 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
382 assert_eq!(eval.param_count(), 1);
383 let result = eval.evaluate(&[7.0]).unwrap();
384 assert!((result[0] - 7.0).abs() < 1e-10);
385 }
386
387 #[test]
388 fn compile_add_two_vars() {
389 let arena = Arena::new();
390 let ctx = AtomArena::new(&arena);
391 let expr = ctx.add(&[ctx.var("x"), ctx.var("y")]);
392 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
393 let result = eval.evaluate(&[2.0, 3.0]).unwrap();
394 assert!((result[0] - 5.0).abs() < 1e-10);
395 }
396
397 #[test]
398 fn compile_mul_var_const() {
399 let arena = Arena::new();
400 let ctx = AtomArena::new(&arena);
401 let expr = ctx.mul(&[ctx.var("x"), ctx.num(3)]);
402 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
403 let result = eval.evaluate(&[4.0]).unwrap();
404 assert!((result[0] - 12.0).abs() < 1e-10);
405 }
406
407 #[test]
408 fn compile_pow_integer_exp() {
409 let arena = Arena::new();
410 let ctx = AtomArena::new(&arena);
411 let expr = ctx.pow(ctx.var("x"), ctx.num(3));
412 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
413 let result = eval.evaluate(&[2.0]).unwrap();
414 assert!((result[0] - 8.0).abs() < 1e-10);
415 }
416
417 #[test]
418 fn compile_sin() {
419 let arena = Arena::new();
420 let ctx = AtomArena::new(&arena);
421 let expr = ctx.fun("sin", &[ctx.var("x")]);
422 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
423 let result = eval.evaluate(&[std::f64::consts::FRAC_PI_2]).unwrap();
424 assert!((result[0] - 1.0).abs() < 1e-10);
425 }
426
427 #[test]
428 fn compile_cos() {
429 let arena = Arena::new();
430 let ctx = AtomArena::new(&arena);
431 let expr = ctx.fun("cos", &[ctx.var("x")]);
432 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
433 let result = eval.evaluate(&[std::f64::consts::PI]).unwrap();
434 assert!((result[0] + 1.0).abs() < 1e-10);
435 }
436
437 #[test]
438 fn compile_exp_log_roundtrip() {
439 let arena = Arena::new();
440 let ctx = AtomArena::new(&arena);
441 let exp_x = ctx.fun("exp", &[ctx.var("x")]);
442 let expr = ctx.fun("log", &[exp_x]);
443 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
444 let result = eval.evaluate(&[2.0]).unwrap();
445 assert!((result[0] - 2.0).abs() < 1e-10);
446 }
447
448 #[test]
449 fn compile_nested_expression() {
450 let arena = Arena::new();
452 let ctx = AtomArena::new(&arena);
453 let x = ctx.var("x");
454 let x_plus_1 = ctx.add(&[x, ctx.num(1)]);
455 let x_minus_1 = ctx.add(&[x, ctx.num(-1)]);
456 let expr = ctx.mul(&[x_plus_1, x_minus_1]);
457 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
458 let result = eval.evaluate(&[3.0]).unwrap();
459 assert!((result[0] - 8.0).abs() < 1e-10);
460 }
461
462 #[test]
463 fn compile_sqrt() {
464 let arena = Arena::new();
465 let ctx = AtomArena::new(&arena);
466 let expr = ctx.fun("sqrt", &[ctx.num(16)]);
467 let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
468 let result = eval.evaluate(&[]).unwrap();
469 assert!((result[0] - 4.0).abs() < 1e-10);
470 }
471
472 #[test]
473 fn compile_zero_params() {
474 let arena = Arena::new();
475 let ctx = AtomArena::new(&arena);
476 let expr = ctx.fun("sin", &[ctx.num(1)]); let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
478 assert_eq!(eval.param_count(), 0);
479 let result = eval.evaluate(&[]).unwrap();
480 assert!((result[0] - 1.0f64.sin()).abs() < 1e-10);
481 }
482
483 #[test]
484 fn compile_with_external_function() {
485 let arena = Arena::new();
486 let ctx = AtomArena::new(&arena);
487 let expr = ctx.fun("square", &[ctx.var("x")]);
488
489 let mut map = FunctionMap::<f64>::new();
490 map.register("square", 1, Box::new(|args| args[0] * args[0]));
491
492 let eval = ExpressionEvaluator::compile_with(expr, map).unwrap();
493 let result = eval.evaluate(&[3.0]).unwrap();
494 assert!((result[0] - 9.0).abs() < 1e-10);
495 }
496
497 #[test]
498 fn compile_external_function_not_registered() {
499 let arena = Arena::new();
500 let ctx = AtomArena::new(&arena);
501 let expr = ctx.fun("missing_fn", &[ctx.var("x")]);
502
503 let result: Result<ExpressionEvaluator<f64>> = ExpressionEvaluator::compile(expr);
504 assert!(result.is_err());
505 }
506
507 #[test]
508 fn compile_with_case_insensitive_external() {
509 let arena = Arena::new();
510 let ctx = AtomArena::new(&arena);
511 let expr = ctx.fun("Square", &[ctx.num(4)]);
512
513 let mut map = FunctionMap::<f64>::new();
514 map.register("square", 1, Box::new(|args| args[0] * args[0]));
515
516 let eval = ExpressionEvaluator::compile_with(expr, map).unwrap();
517 let result = eval.evaluate(&[]).unwrap();
518 assert!((result[0] - 16.0).abs() < 1e-10);
519 }
520
521 #[test]
522 fn compile_multi_two_outputs() {
523 let arena = Arena::new();
524 let ctx = AtomArena::new(&arena);
525 let sum = ctx.add(&[ctx.var("x"), ctx.var("y")]);
526 let prod = ctx.mul(&[ctx.var("x"), ctx.var("y")]);
527 let eval: ExpressionEvaluator<f64> =
528 ExpressionEvaluator::compile_multi(&[sum, prod]).unwrap();
529 assert_eq!(eval.result_count(), 2);
530 assert_eq!(eval.param_count(), 2);
531 let result = eval.evaluate(&[2.0, 3.0]).unwrap();
532 assert!((result[0] - 5.0).abs() < 1e-10);
533 assert!((result[1] - 6.0).abs() < 1e-10);
534 }
535
536 #[test]
537 fn compile_multi_shared_subexpression() {
538 let arena = Arena::new();
540 let ctx = AtomArena::new(&arena);
541 let sin_x = ctx.fun("sin", &[ctx.var("x")]);
542 let out0 = ctx.add(&[sin_x, ctx.num(1)]);
543 let out1 = ctx.mul(&[sin_x, ctx.num(2)]);
544 let eval: ExpressionEvaluator<f64> =
545 ExpressionEvaluator::compile_multi(&[out0, out1]).unwrap();
546 let result = eval.evaluate(&[std::f64::consts::FRAC_PI_2]).unwrap();
547 assert!((result[0] - 2.0).abs() < 1e-10);
548 assert!((result[1] - 2.0).abs() < 1e-10);
549 }
550
551 #[test]
552 fn compile_multi_constant_folding() {
553 let arena = Arena::new();
555 let ctx = AtomArena::new(&arena);
556 let five_x = ctx.mul(&[ctx.add(&[ctx.num(2), ctx.num(3)]), ctx.var("x")]);
557 let x_pow_1 = ctx.pow(ctx.var("x"), ctx.num(1));
558 let eval: ExpressionEvaluator<f64> =
559 ExpressionEvaluator::compile_multi(&[five_x, x_pow_1]).unwrap();
560 let result = eval.evaluate(&[4.0]).unwrap();
561 assert!((result[0] - 20.0).abs() < 1e-10);
562 assert!((result[1] - 4.0).abs() < 1e-10);
563 }
564
565 #[test]
566 fn compile_multi_with_external_function() {
567 let arena = Arena::new();
568 let ctx = AtomArena::new(&arena);
569 let sq = ctx.fun("square", &[ctx.var("x")]);
570 let cube_arg = ctx.fun("square", &[ctx.var("x")]);
571 let out1 = ctx.mul(&[cube_arg, ctx.var("x")]);
572
573 let mut map = FunctionMap::<f64>::new();
574 map.register("square", 1, Box::new(|args| args[0] * args[0]));
575
576 let eval = ExpressionEvaluator::compile_multi_with(&[sq, out1], map).unwrap();
577 let result = eval.evaluate(&[3.0]).unwrap();
578 assert!((result[0] - 9.0).abs() < 1e-10);
579 assert!((result[1] - 27.0).abs() < 1e-10);
580 }
581}