use ocas_atom::AtomArena;
use ocas_core::arena::Arena;
use ocas_eval::{ExpressionEvaluator, compile_atom};
use ocas_parse::parse;
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
#[pyclass(name = "ExpressionEvaluator")]
pub struct PyExpressionEvaluator {
evaluator: ExpressionEvaluator<f64>,
param_names: Vec<String>,
}
#[pymethods]
impl PyExpressionEvaluator {
#[new]
fn new(input: &str, param_names: Vec<String>) -> PyResult<Self> {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let atom =
parse(&ctx, input).map_err(|e| PyValueError::new_err(format!("parse error: {e}")))?;
let evaluator = compile_atom::<f64>(atom)
.map_err(|e| PyValueError::new_err(format!("compile error: {e}")))?;
Ok(PyExpressionEvaluator {
evaluator,
param_names,
})
}
fn evaluate(&self, values: Vec<f64>) -> PyResult<Vec<f64>> {
if values.len() != self.param_names.len() {
return Err(PyValueError::new_err(format!(
"expected {} values, got {}",
self.param_names.len(),
values.len()
)));
}
self.evaluator
.evaluate(&values)
.map_err(|e| PyValueError::new_err(format!("evaluation error: {e}")))
}
#[getter]
fn n_params(&self) -> usize {
self.param_names.len()
}
fn __repr__(&self) -> String {
format!("ExpressionEvaluator(params={:?})", self.param_names)
}
}