feat: add batch evaluation for python and tests in workflows

This commit is contained in:
Stefan
2026-06-29 14:32:25 +02:00
parent 7805da79f7
commit 35f6519c71
6 changed files with 198 additions and 1 deletions
+27
View File
@@ -293,3 +293,30 @@ pub fn response_to_py(py: Python<'_>, response: DecisionGraphResponse) -> PyResu
Ok(dict.into_any().unbind())
}
pub fn batch_results_to_py(
py: Python<'_>,
outcomes: Vec<Result<PortableResponse, Value>>,
) -> PyResult<Py<PyAny>> {
let list = PyList::empty(py);
for outcome in outcomes {
let item = PyDict::new(py);
match outcome {
Ok(response) => {
item.set_item("success", true)?;
item.set_item("data", response.into_py(py)?)?;
item.set_item("error", py.None())?;
}
Err(error) => {
item.set_item("success", false)?;
item.set_item("data", py.None())?;
item.set_item("error", value_to_object(py, &error)?)?;
}
}
list.append(item)?;
}
Ok(list.into_any().unbind())
}
+82
View File
@@ -165,6 +165,88 @@ impl PyZenEngine {
Ok(result.unbind())
}
#[pyo3(signature = (requests, opts=None))]
pub fn evaluate_batch(
&self,
py: Python,
requests: Vec<(String, PyValue)>,
opts: Option<PyZenEvaluateOptions>,
) -> PyResult<Py<PyAny>> {
let options: EvaluationOptions = opts.unwrap_or_default().into();
let trace = options.trace;
let max_depth = options.max_depth;
let engine = self.engine.clone();
let outcomes = py.allow_threads(|| {
block_on(async move {
let mut handles = Vec::with_capacity(requests.len());
for (key, ctx) in requests {
let engine = engine.clone();
handles.push(worker_pool().spawn_pinned(move || async move {
let options = EvaluationOptions { trace, max_depth };
engine
.evaluate_with_opts(key, ctx.0.into(), options)
.await
.map(crate::convert::PortableResponse::build)
.map_err(|e| serde_json::to_value(e.as_ref()).unwrap_or_default())
}));
}
let mut outcomes = Vec::with_capacity(handles.len());
for handle in handles {
outcomes.push(match handle.await {
Ok(outcome) => outcome,
Err(_) => Err(Value::String("evaluation worker panicked".to_string())),
});
}
outcomes
})
});
crate::convert::batch_results_to_py(py, outcomes)
}
#[pyo3(signature = (requests, opts=None))]
pub fn async_evaluate_batch<'py>(
&'py self,
py: Python<'py>,
requests: Vec<(String, PyValue)>,
opts: Option<PyZenEvaluateOptions>,
) -> PyResult<Py<PyAny>> {
let options: EvaluationOptions = opts.unwrap_or_default().into();
let trace = options.trace;
let max_depth = options.max_depth;
let engine = self.engine.clone();
let result = tokio::future_into_py_with_locals(py, get_current_locals(py)?, async move {
let mut handles = Vec::with_capacity(requests.len());
for (key, ctx) in requests {
let engine = engine.clone();
handles.push(worker_pool().spawn_pinned(move || async move {
let options = EvaluationOptions { trace, max_depth };
engine
.evaluate_with_opts(key, ctx.0.into(), options)
.await
.map(crate::convert::PortableResponse::build)
.map_err(|e| serde_json::to_value(e.as_ref()).unwrap_or_default())
}));
}
let mut outcomes = Vec::with_capacity(handles.len());
for handle in handles {
outcomes.push(match handle.await {
Ok(outcome) => outcome,
Err(_) => Err(Value::String("evaluation worker panicked".to_string())),
});
}
Python::with_gil(|py| crate::convert::batch_results_to_py(py, outcomes))
})?;
Ok(result.unbind())
}
pub fn create_decision(&self, content: PyZenDecisionContentJson) -> PyResult<PyZenDecision> {
let decision = self
.engine
+16
View File
@@ -79,6 +79,22 @@ class AsyncZenEngine(unittest.IsolatedAsyncioTestCase):
r = await functionDecision.async_evaluate({"input": 15})
self.assertEqual(r["result"]["output"], 30)
async def test_async_evaluate_batch(self):
engine = zen.ZenEngine({"loader": loader})
results = await engine.async_evaluate_batch([
("table.json", {"input": 12}),
("table.json", {"input": 2}),
("does-not-exist.json", {}),
])
self.assertEqual(len(results), 3)
self.assertTrue(results[0]["success"])
self.assertEqual(results[0]["data"]["result"]["output"], 10)
self.assertTrue(results[1]["success"])
self.assertEqual(results[1]["data"]["result"]["output"], 0)
self.assertFalse(results[2]["success"])
self.assertIsNotNone(results[2]["error"])
async def test_evaluate_graphs(self):
engine = zen.ZenEngine({"loader": graph_loader})
json_files = glob.glob("../../test-data/graphs/*.json")
+16
View File
@@ -58,6 +58,22 @@ class ZenEngine(unittest.TestCase):
r = functionDecision.evaluate({"input": 15})
self.assertEqual(r["result"]["output"], 30)
def test_evaluate_batch(self):
engine = zen.ZenEngine({"loader": loader})
results = engine.evaluate_batch([
("table.json", {"input": 12}),
("table.json", {"input": 2}),
("does-not-exist.json", {}),
])
self.assertEqual(len(results), 3)
self.assertTrue(results[0]["success"])
self.assertEqual(results[0]["data"]["result"]["output"], 10)
self.assertTrue(results[1]["success"])
self.assertEqual(results[1]["data"]["result"]["output"], 0)
self.assertFalse(results[2]["success"])
self.assertIsNotNone(results[2]["error"])
def test_engine_custom_handler(self):
engine = zen.ZenEngine({"loader": loader, "customHandler": custom_handler})
r1 = engine.evaluate("custom.json", {"a": 10})
+14
View File
@@ -13,8 +13,15 @@ class EvaluateResponse(TypedDict):
trace: dict
class BatchEvaluateResult(TypedDict):
success: bool
data: Optional[EvaluateResponse]
error: Optional[Any]
ZenContext: TypeAlias = Union[str, bytes, dict]
ZenDecisionContentInput: TypeAlias = Union[str, ZenDecisionContent]
ZenBatchRequest: TypeAlias = "tuple[str, ZenContext]"
class ZenEngine:
@@ -26,6 +33,13 @@ class ZenEngine:
def async_evaluate(self, key: str, context: ZenContext, options: Optional[DecisionEvaluateOptions] = None) -> \
Awaitable[EvaluateResponse]: ...
def evaluate_batch(self, requests: "list[ZenBatchRequest]",
options: Optional[DecisionEvaluateOptions] = None) -> "list[BatchEvaluateResult]": ...
def async_evaluate_batch(self, requests: "list[ZenBatchRequest]",
options: Optional[DecisionEvaluateOptions] = None) -> \
Awaitable["list[BatchEvaluateResult]"]: ...
def create_decision(self, content: ZenDecisionContentInput) -> ZenDecision: ...
def get_decision(self, key: str) -> ZenDecision: ...