1use 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#[pyclass(name = "ExpressionEvaluator")]
27pub struct PyExpressionEvaluator {
28 evaluator: ExpressionEvaluator<f64>,
29 param_names: Vec<String>,
30}
31
32#[pymethods]
33impl PyExpressionEvaluator {
34 #[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 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 #[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}