Skip to main content

ocas_eval/
evaluator.rs

1//! Stack-based expression evaluator.
2//!
3//! The [`ExpressionEvaluator`] executes a pre-compiled instruction sequence
4//! on a flat stack, producing numeric results from input parameters.
5
6use crate::domain::{EvaluationDomain, PowfExtension};
7use crate::error::{EvaluationError, Result};
8use crate::function_map::FunctionMap;
9use crate::instruction::Instr;
10
11/// A compiled expression ready for numeric evaluation.
12///
13/// The evaluator holds a sequence of [`Instr`]s and a pre-allocated stack.
14/// Call [`evaluate`](ExpressionEvaluator::evaluate) with parameter values
15/// to compute the result.
16pub struct ExpressionEvaluator<T: EvaluationDomain> {
17    /// The instruction sequence to execute.
18    instructions: Vec<Instr>,
19    /// Number of parameter slots at the start of the stack.
20    param_count: usize,
21    /// Number of constant slots (after params).
22    #[allow(dead_code)]
23    const_count: usize,
24    /// Total stack size (params + constants + temporaries + outputs).
25    stack_size: usize,
26    /// Indices of result slots in the stack.
27    result_indices: Vec<usize>,
28    /// Pre-computed constant values.
29    constants: Vec<T>,
30    /// Optional user-defined function registry.
31    function_map: Option<FunctionMap<T>>,
32}
33
34impl<T: EvaluationDomain + PowfExtension> ExpressionEvaluator<T> {
35    /// Create an evaluator from compiled instruction data.
36    ///
37    /// This is a low-level constructor used by the compiler. Most users
38    /// should use [`ExpressionEvaluator::compile`] instead.
39    #[allow(dead_code)]
40    pub(crate) fn new(
41        instructions: Vec<Instr>,
42        param_count: usize,
43        const_count: usize,
44        stack_size: usize,
45        result_indices: Vec<usize>,
46        constants: Vec<T>,
47    ) -> Self {
48        Self {
49            instructions,
50            param_count,
51            const_count,
52            stack_size,
53            result_indices,
54            constants,
55            function_map: None,
56        }
57    }
58
59    /// Create an evaluator with a function map for user-defined functions.
60    #[allow(dead_code)]
61    pub(crate) fn new_with_functions(
62        instructions: Vec<Instr>,
63        param_count: usize,
64        const_count: usize,
65        stack_size: usize,
66        result_indices: Vec<usize>,
67        constants: Vec<T>,
68        function_map: FunctionMap<T>,
69    ) -> Self {
70        Self {
71            instructions,
72            param_count,
73            const_count,
74            stack_size,
75            result_indices,
76            constants,
77            function_map: Some(function_map),
78        }
79    }
80
81    /// Return the number of parameters expected by this evaluator.
82    pub fn param_count(&self) -> usize {
83        self.param_count
84    }
85
86    /// Return the number of results produced by this evaluator.
87    pub fn result_count(&self) -> usize {
88        self.result_indices.len()
89    }
90
91    /// Return the stack size required to evaluate this expression.
92    ///
93    /// Callers reusing a stack buffer across evaluations should
94    /// pre-allocate a `Vec` with this capacity.
95    pub fn stack_size(&self) -> usize {
96        self.stack_size
97    }
98
99    /// Evaluate the expression with the given parameter values.
100    ///
101    /// Returns a vector of result values. The number of results equals
102    /// `result_indices.len()`.
103    ///
104    /// # Errors
105    ///
106    /// Returns [`EvaluationError`] if the number of parameters does not
107    /// match, or if an arithmetic error occurs (e.g. division by zero).
108    ///
109    /// # Example
110    ///
111    /// ```ignore
112    /// let result = evaluator.evaluate(&[2.0_f64])?;
113    /// assert_eq!(result.len(), 1);
114    /// ```
115    pub fn evaluate(&self, params: &[T]) -> Result<Vec<T>> {
116        let mut stack: Vec<T> = Vec::with_capacity(self.stack_size);
117        let mut results: Vec<T> = Vec::with_capacity(self.result_indices.len());
118        self.evaluate_with_stack(params, &mut stack, &mut results)?;
119        Ok(results)
120    }
121
122    /// Evaluate the expression, reusing caller-provided buffers.
123    ///
124    /// `stack` must have capacity for at least
125    /// [`stack_size`](ExpressionEvaluator::stack_size) elements (it is
126    /// resized as needed); `results` is filled with the result values.
127    /// Reusing buffers across calls avoids per-evaluation heap
128    /// allocation, which matters for streaming and batch workloads.
129    ///
130    /// # Errors
131    ///
132    /// Returns [`EvaluationError`] if the number of parameters does not
133    /// match, or if an arithmetic error occurs.
134    pub fn evaluate_with_stack(
135        &self,
136        params: &[T],
137        stack: &mut Vec<T>,
138        results: &mut Vec<T>,
139    ) -> Result<()> {
140        if params.len() != self.param_count {
141            return Err(EvaluationError::WrongArity {
142                name: "<expr>".into(),
143                expected: self.param_count,
144                got: params.len(),
145            });
146        }
147
148        stack.clear();
149        stack.resize(self.stack_size, T::zero());
150
151        // Fill parameters
152        for (i, p) in params.iter().enumerate() {
153            stack[i] = p.clone();
154        }
155
156        // Fill constants
157        for (i, c) in self.constants.iter().enumerate() {
158            stack[self.param_count + i] = c.clone();
159        }
160
161        // Execute instructions
162        for instr in &self.instructions {
163            match instr {
164                Instr::Add { dst, srcs } => {
165                    let mut sum = stack[srcs[0]].clone();
166                    for idx in &srcs[1..] {
167                        sum = sum.add_ref(&stack[*idx]);
168                    }
169                    stack[*dst] = sum;
170                }
171                Instr::Mul { dst, srcs } => {
172                    let mut prod = stack[srcs[0]].clone();
173                    for idx in &srcs[1..] {
174                        prod = prod.mul_ref(&stack[*idx]);
175                    }
176                    stack[*dst] = prod;
177                }
178                Instr::Pow { dst, base, exp } => {
179                    stack[*dst] = stack[*base].powi_ref(*exp);
180                }
181                Instr::Powf { dst, base, exp } => {
182                    let result = stack[*base].powf_ref(&stack[*exp])?;
183                    stack[*dst] = result;
184                }
185                Instr::BuiltinOp { dst, op, src } => {
186                    let name = match op {
187                        crate::instruction::BuiltinOp::Sin => "sin",
188                        crate::instruction::BuiltinOp::Cos => "cos",
189                        crate::instruction::BuiltinOp::Tan => "tan",
190                        crate::instruction::BuiltinOp::Sec => "sec",
191                        crate::instruction::BuiltinOp::Csc => "csc",
192                        crate::instruction::BuiltinOp::Cot => "cot",
193                        crate::instruction::BuiltinOp::Exp => "exp",
194                        crate::instruction::BuiltinOp::Log => "log",
195                        crate::instruction::BuiltinOp::Sqrt => "sqrt",
196                        crate::instruction::BuiltinOp::Abs => "abs",
197                    };
198                    let result = T::resolve_builtin(name, &stack[*src])?;
199                    stack[*dst] = result;
200                }
201                Instr::ExternalFun { dst, fn_idx, srcs } => {
202                    let args: Vec<T> = srcs.iter().map(|&i| stack[i].clone()).collect();
203                    let result = self
204                        .function_map
205                        .as_ref()
206                        .and_then(|fm| fm.call_by_index(*fn_idx, &args))
207                        .ok_or_else(|| EvaluationError::FunctionNotFound {
208                            name: format!("external function at index {fn_idx}"),
209                        })?;
210                    stack[*dst] = result;
211                }
212                Instr::Copy { dst, src } => {
213                    stack[*dst] = stack[*src].clone();
214                }
215            }
216        }
217
218        // Collect results
219        results.clear();
220        results.extend(self.result_indices.iter().map(|&i| stack[i].clone()));
221
222        Ok(())
223    }
224}
225
226/// Cranelift JIT compilation support.
227#[cfg(feature = "jit")]
228impl ExpressionEvaluator<f64> {
229    /// Compile this evaluator's instruction sequence to native machine
230    /// code via the Cranelift JIT backend.
231    ///
232    /// Constants are embedded as immediates and all result slots are
233    /// written in order, so multi-output evaluators are fully supported.
234    ///
235    /// # Errors
236    ///
237    /// Returns [`EvaluationError::UnsupportedOperation`] if the expression
238    /// contains instructions the JIT cannot lower (e.g. external
239    /// functions or `sec`/`csc`/`cot`), or
240    /// [`EvaluationError::JitCompilationError`] on backend failure.
241    pub fn compile_jit(&self) -> Result<crate::jit::JitCompiledFunction> {
242        let constants: Vec<f64> = self.constants.clone();
243        crate::jit::JitEngine::compile(
244            &self.instructions,
245            self.param_count,
246            &constants,
247            &self.result_indices,
248        )
249    }
250
251    /// Compile this evaluator's instruction sequence to single-precision
252    /// native code. Constants are narrowed from f64 to f32; results have
253    /// f32 precision.
254    ///
255    /// # Errors
256    ///
257    /// Same conditions as [`compile_jit`](ExpressionEvaluator::compile_jit).
258    pub fn compile_jit_f32(&self) -> Result<crate::jit::JitCompiledF32> {
259        let constants: Vec<f32> = self.constants.iter().map(|&c| c as f32).collect();
260        crate::jit::JitEngine::compile_f32(
261            &self.instructions,
262            self.param_count,
263            &constants,
264            &self.result_indices,
265        )
266    }
267}
268
269/// SIMD batch evaluation support.
270#[cfg(feature = "simd")]
271impl ExpressionEvaluator<f64> {
272    /// Compile this evaluator into a [`VectorEvaluator`] for batch SIMD evaluation.
273    ///
274    /// The resulting evaluator processes multiple input values simultaneously
275    /// using the best available SIMD width (SSE2/AVX2/AVX-512).
276    ///
277    /// # Errors
278    ///
279    /// Returns [`EvaluationError::UnsupportedOperation`] if the expression
280    /// contains external functions, which are not supported in SIMD mode.
281    pub fn compile_vector_evaluator(&self) -> Result<crate::simd::VectorEvaluator> {
282        // Check for unsupported instructions
283        for instr in &self.instructions {
284            if let Instr::ExternalFun { .. } = instr {
285                return Err(EvaluationError::UnsupportedOperation {
286                    message: "external functions not supported in SIMD mode".into(),
287                });
288            }
289        }
290
291        Ok(crate::simd::VectorEvaluator::new(
292            self.instructions.clone(),
293            self.param_count,
294            self.const_count,
295            self.stack_size,
296            self.result_indices.clone(),
297            self.constants.clone(),
298        ))
299    }
300
301    /// Compile this evaluator into a single-precision
302    /// [`VectorEvaluatorF32`](crate::simd::VectorEvaluatorF32) for batch
303    /// SIMD evaluation. Constants are narrowed from f64 to f32; on the
304    /// same hardware this doubles the SIMD lane count.
305    ///
306    /// # Errors
307    ///
308    /// Returns [`EvaluationError::UnsupportedOperation`] if the expression
309    /// contains external functions, which are not supported in SIMD mode.
310    pub fn compile_vector_evaluator_f32(&self) -> Result<crate::simd::VectorEvaluatorF32> {
311        // Check for unsupported instructions
312        for instr in &self.instructions {
313            if let Instr::ExternalFun { .. } = instr {
314                return Err(EvaluationError::UnsupportedOperation {
315                    message: "external functions not supported in SIMD mode".into(),
316                });
317            }
318        }
319
320        let constants: Vec<f32> = self.constants.iter().map(|&c| c as f32).collect();
321        Ok(crate::simd::VectorEvaluatorF32::new(
322            self.instructions.clone(),
323            self.param_count,
324            self.const_count,
325            self.stack_size,
326            self.result_indices.clone(),
327            constants,
328        ))
329    }
330}
331
332#[cfg(test)]
333mod tests {
334    use super::*;
335
336    fn make_simple_evaluator() -> ExpressionEvaluator<f64> {
337        // Evaluate: x + 1
338        // Stack layout: [param(0)=x] [const(0)=1.0] [temp(0)=result]
339        // Instructions: Add(temp(0), param(0), const(0))
340        let instructions = vec![Instr::Add {
341            dst: 2,
342            srcs: vec![0, 1],
343        }];
344        let constants = vec![1.0f64];
345        ExpressionEvaluator::new(instructions, 1, 1, 3, vec![2], constants)
346    }
347
348    #[test]
349    fn simple_add() {
350        let eval = make_simple_evaluator();
351        assert_eq!(eval.param_count(), 1);
352        let result = eval.evaluate(&[2.0]).unwrap();
353        assert!((result[0] - 3.0).abs() < 1e-10);
354    }
355
356    #[test]
357    fn wrong_param_count() {
358        let eval = make_simple_evaluator();
359        assert!(eval.evaluate(&[1.0, 2.0]).is_err());
360        assert!(eval.evaluate(&[]).is_err());
361    }
362
363    #[test]
364    fn mul_expression() {
365        // Evaluate: x * 2
366        let instructions = vec![Instr::Mul {
367            dst: 2,
368            srcs: vec![0, 1],
369        }];
370        let constants = vec![2.0f64];
371        let eval = ExpressionEvaluator::new(instructions, 1, 1, 3, vec![2], constants);
372        assert!((eval.evaluate(&[3.0]).unwrap()[0] - 6.0).abs() < 1e-10);
373    }
374
375    #[test]
376    fn pow_expression() {
377        // Evaluate: x^3
378        let instructions = vec![Instr::Pow {
379            dst: 1,
380            base: 0,
381            exp: 3,
382        }];
383        let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
384        assert!((eval.evaluate(&[2.0]).unwrap()[0] - 8.0).abs() < 1e-10);
385    }
386
387    #[test]
388    fn builtin_sin() {
389        // Evaluate: sin(x)
390        let instructions = vec![Instr::BuiltinOp {
391            dst: 1,
392            op: crate::instruction::BuiltinOp::Sin,
393            src: 0,
394        }];
395        let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
396        let result = eval.evaluate(&[std::f64::consts::FRAC_PI_2]).unwrap();
397        assert!((result[0] - 1.0).abs() < 1e-10);
398    }
399
400    #[test]
401    fn copy_instruction() {
402        // Evaluate: x (identity)
403        let instructions = vec![Instr::Copy { dst: 1, src: 0 }];
404        let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
405        assert!((eval.evaluate(&[42.0]).unwrap()[0] - 42.0).abs() < 1e-10);
406    }
407
408    #[test]
409    fn evaluate_with_stack_reuses_buffers() {
410        let eval = make_simple_evaluator();
411        let mut stack: Vec<f64> = Vec::with_capacity(eval.stack_size());
412        let mut results: Vec<f64> = Vec::with_capacity(eval.result_count());
413
414        for x in 0..100 {
415            eval.evaluate_with_stack(&[x as f64], &mut stack, &mut results)
416                .unwrap();
417            assert!((results[0] - (x as f64 + 1.0)).abs() < 1e-10);
418        }
419        // Buffers retain capacity across calls (no reallocation needed)
420        assert!(stack.capacity() >= eval.stack_size());
421    }
422
423    #[test]
424    fn evaluate_with_stack_wrong_arity() {
425        let eval = make_simple_evaluator();
426        let mut stack = Vec::new();
427        let mut results = Vec::new();
428        assert!(
429            eval.evaluate_with_stack(&[1.0, 2.0], &mut stack, &mut results)
430                .is_err()
431        );
432    }
433
434    #[test]
435    fn result_and_stack_getters() {
436        let eval = make_simple_evaluator();
437        assert_eq!(eval.result_count(), 1);
438        assert_eq!(eval.stack_size(), 3);
439    }
440}