Skip to main content

ocas_py/
eval.rs

1//! Python evaluation module — numeric evaluation of expressions.
2//!
3//! Provides [`PyExpressionEvaluator`] wrapping
4//! [`ocas_eval::ExpressionEvaluator<f64>`]. The evaluator compiles an
5//! expression to an instruction sequence once, then evaluates it many
6//! times with different parameter values.
7
8use ocas_atom::AtomArena;
9use ocas_core::arena::Arena;
10use ocas_eval::{ExpressionEvaluator, compile_atom};
11use ocas_parse::parse;
12use pyo3::exceptions::PyValueError;
13use pyo3::prelude::*;
14
15/// A compiled expression evaluator for fast numeric evaluation.
16///
17/// Compile once, evaluate many times:
18///
19/// ```python
20/// from ocas import ExpressionEvaluator
21///
22/// evaluator = ExpressionEvaluator("x^2 + y", ["x", "y"])
23/// print(evaluator.evaluate([3.0, 1.0]))  # [10.0]
24/// print(evaluator.evaluate([2.0, 0.0]))  # [4.0]
25/// ```
26#[pyclass(name = "ExpressionEvaluator")]
27pub struct PyExpressionEvaluator {
28    evaluator: ExpressionEvaluator<f64>,
29    param_names: Vec<String>,
30}
31
32#[pymethods]
33impl PyExpressionEvaluator {
34    /// Compile `input` with the given parameter names (in order).
35    #[new]
36    fn new(input: &str, param_names: Vec<String>) -> PyResult<Self> {
37        let arena = Arena::new();
38        let ctx = AtomArena::new(&arena);
39        let atom =
40            parse(&ctx, input).map_err(|e| PyValueError::new_err(format!("parse error: {e}")))?;
41        let evaluator = compile_atom::<f64>(atom)
42            .map_err(|e| PyValueError::new_err(format!("compile error: {e}")))?;
43        Ok(PyExpressionEvaluator {
44            evaluator,
45            param_names,
46        })
47    }
48
49    /// Evaluate the compiled expression with the given parameter values.
50    ///
51    /// `values` must be a list of floats with the same length as the
52    /// parameter names passed to the constructor. Returns a list of result
53    /// floats.
54    fn evaluate(&self, values: Vec<f64>) -> PyResult<Vec<f64>> {
55        if values.len() != self.param_names.len() {
56            return Err(PyValueError::new_err(format!(
57                "expected {} values, got {}",
58                self.param_names.len(),
59                values.len()
60            )));
61        }
62        self.evaluator
63            .evaluate(&values)
64            .map_err(|e| PyValueError::new_err(format!("evaluation error: {e}")))
65    }
66
67    /// Number of parameters.
68    #[getter]
69    fn n_params(&self) -> usize {
70        self.param_names.len()
71    }
72
73    fn __repr__(&self) -> String {
74        format!("ExpressionEvaluator(params={:?})", self.param_names)
75    }
76}