feat: implement new loaders across languages (#487)

This commit is contained in:
stefan-gorules
2026-07-20 12:20:48 +02:00
committed by GitHub
parent fe4790bee4
commit 1818371c90
19 changed files with 938 additions and 43 deletions
+128 -11
View File
@@ -8,13 +8,15 @@ use crate::mt::{block_on, worker_pool};
use crate::value::PyValue;
use crate::variable::PyVariable;
use anyhow::{anyhow, Context};
use pyo3::prelude::{PyAnyMethods, PyDictMethods};
use pyo3::types::PyDict;
use pyo3::prelude::{PyAnyMethods, PyDictMethods, PyListMethods};
use pyo3::types::{PyDict, PyList};
use pyo3::{pyclass, pymethods, Bound, FromPyObject, IntoPyObjectExt, Py, PyAny, PyResult, Python};
use pyo3_async_runtimes::tokio::get_current_locals;
use pyo3_async_runtimes::{tokio, TaskLocals};
use pythonize::{depythonize, pythonize};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use zen_engine::loader::{DynamicLoader, LoaderConfig};
use zen_engine::{DecisionEngine, EvaluationOptions};
#[pyclass]
@@ -76,6 +78,53 @@ impl Default for PyZenEngine {
}
}
pub struct PyBatchRequest {
key: String,
context: Value,
}
impl<'py> FromPyObject<'py> for PyBatchRequest {
fn extract_bound(ob: &Bound<'py, PyAny>) -> PyResult<Self> {
let dict = ob.downcast::<PyDict>()?;
let key: String = dict
.get_item("key")?
.ok_or_else(|| anyhow!("batch request requires a 'key'"))?
.extract()?;
let context = dict
.get_item("context")?
.ok_or_else(|| anyhow!("batch request requires a 'context'"))?
.extract::<PyValue>()?;
Ok(Self {
key,
context: context.0,
})
}
}
impl PyZenEngine {
fn config_loader(config: &Bound<'_, PyDict>) -> PyResult<DynamicLoader> {
let loader_type: Option<String> =
config.get_item("type")?.map(|v| v.extract()).transpose()?;
let loader_config = match loader_type.as_deref() {
Some("zip") => {
let bytes = config
.get_item("bytes")?
.ok_or_else(|| anyhow!("zip loader requires a 'bytes' value"))?;
LoaderConfig::Zip {
bytes: bytes.extract()?,
}
}
_ => depythonize(config.as_any())?,
};
Ok(loader_config.into_loader()?)
}
}
#[pymethods]
impl PyZenEngine {
#[new]
@@ -85,11 +134,6 @@ impl PyZenEngine {
return Ok(Default::default());
};
let loader = match options.get_item("loader")? {
Some(loader) => Some(loader.into_py_any(py)?),
None => None,
};
let custom_node = match options.get_item("customHandler")? {
Some(custom_node) => Some(custom_node.into_py_any(py)?),
None => None,
@@ -102,11 +146,25 @@ impl PyZenEngine {
.flatten()
};
let loader: DynamicLoader = match options.get_item("loader")? {
Some(loader) => match loader.downcast::<PyDict>() {
Ok(config) => Self::config_loader(config)?,
Err(_) => Arc::new(PyDecisionLoader::new(
Some(loader.into_py_any(py)?),
make_locals(),
)),
},
None => Arc::new(PyDecisionLoader::default()),
};
let engine = DecisionEngine::new(
loader,
Arc::new(PyCustomNode::new(custom_node, make_locals())),
);
engine.compile();
Ok(Self {
engine: Arc::new(DecisionEngine::new(
Arc::new(PyDecisionLoader::new(loader, make_locals())),
Arc::new(PyCustomNode::new(custom_node, make_locals())),
)),
engine: Arc::new(engine),
})
}
@@ -131,6 +189,65 @@ impl PyZenEngine {
crate::convert::response_to_py(py, result)
}
#[pyo3(signature = (requests, opts=None))]
pub fn evaluate_batch(
&self,
py: Python,
requests: Vec<PyBatchRequest>,
opts: Option<PyZenEvaluateOptions>,
) -> PyResult<Py<PyAny>> {
let options: EvaluationOptions = opts.unwrap_or_default().into();
let handles: Vec<_> = requests
.into_iter()
.map(|request| {
let engine = self.engine.clone();
worker_pool().spawn_pinned(move || async move {
engine
.evaluate_with_opts(request.key, request.context.into(), options)
.await
.map(crate::convert::PortableResponse::build)
.map_err(|e| {
serde_json::to_value(e.as_ref())
.unwrap_or_else(|_| Value::String(e.to_string()))
})
})
})
.collect();
let results = py.allow_threads(|| {
block_on(async move {
let mut out = Vec::with_capacity(handles.len());
for handle in handles {
out.push(handle.await);
}
out
})
});
let list = PyList::empty(py);
for result in results {
let item = PyDict::new(py);
match result {
Ok(Ok(response)) => {
item.set_item("success", true)?;
item.set_item("data", response.into_py(py)?)?;
}
Ok(Err(error)) => {
item.set_item("success", false)?;
item.set_item("error", pythonize(py, &error)?)?;
}
Err(_) => {
item.set_item("success", false)?;
item.set_item("error", "evaluation worker panicked")?;
}
}
list.append(item)?;
}
Ok(list.into_py_any(py)?)
}
#[pyo3(signature = (key, ctx, opts=None))]
pub fn async_evaluate<'py>(
&'py self,
+56
View File
@@ -68,6 +68,62 @@ class ZenEngine(unittest.TestCase):
self.assertEqual(r2["result"]["sum"], 30)
self.assertEqual(r3["result"]["sum"], 40)
def test_static_loader_config(self):
with open("../../test-data/table.json", "r") as f:
table_content = json.loads(f.read())
engine = zen.ZenEngine({"loader": {"type": "static", "content": {"table.json": table_content}}})
r1 = engine.evaluate("table.json", {"input": 2})
r2 = engine.evaluate("table.json", {"input": 12})
self.assertEqual(r1["result"]["output"], 0)
self.assertEqual(r2["result"]["output"], 10)
self.assertRaises(RuntimeError, engine.evaluate, "missing.json", {})
def test_fs_loader_config(self):
engine = zen.ZenEngine({"loader": {"type": "fs", "path": "../../test-data"}})
r1 = engine.evaluate("table.json", {"input": 2})
r2 = engine.evaluate("table.json", {"input": 12})
self.assertEqual(r1["result"]["output"], 0)
self.assertEqual(r2["result"]["output"], 10)
def test_zip_loader_config(self):
import io
import zipfile
buffer = io.BytesIO()
with zipfile.ZipFile(buffer, "w", zipfile.ZIP_DEFLATED) as archive:
with open("../../test-data/table.json", "rb") as f:
archive.writestr("table.json", f.read())
engine = zen.ZenEngine({"loader": {"type": "zip", "bytes": buffer.getvalue()}})
r1 = engine.evaluate("table.json", {"input": 2})
r2 = engine.evaluate("table.json", {"input": 12})
self.assertEqual(r1["result"]["output"], 0)
self.assertEqual(r2["result"]["output"], 10)
def test_evaluate_batch(self):
engine = zen.ZenEngine({"loader": {"type": "fs", "path": "../../test-data"}})
results = engine.evaluate_batch([
{"key": "table.json", "context": {"input": 12}},
{"key": "missing.json", "context": {}},
{"key": "table.json", "context": {"input": 5}},
])
self.assertEqual(len(results), 3)
self.assertTrue(results[0]["success"])
self.assertEqual(results[0]["data"]["result"]["output"], 10)
self.assertFalse(results[1]["success"])
self.assertIn("error", results[1])
self.assertTrue(results[2]["success"])
self.assertEqual(results[2]["data"]["result"]["output"], 0)
def test_evaluate_batch_empty(self):
engine = zen.ZenEngine({"loader": loader})
self.assertEqual(engine.evaluate_batch([]), [])
def test_evaluate_expression(self):
result = zen.evaluate_expression("sum(a)", {"a": [1, 2, 3, 4]})
self.assertEqual(result, 10)
+40 -2
View File
@@ -1,4 +1,4 @@
from collections.abc import Awaitable
from collections.abc import Awaitable, Callable
from typing import Any, Optional, TypedDict, Literal, TypeAlias, Union
@@ -17,12 +17,50 @@ ZenContext: TypeAlias = Union[str, bytes, dict]
ZenDecisionContentInput: TypeAlias = Union[str, ZenDecisionContent]
class StaticLoaderConfig(TypedDict):
type: Literal["static"]
content: dict[str, dict]
class FilesystemLoaderConfig(TypedDict):
type: Literal["fs"]
path: str
class ZipLoaderConfig(TypedDict):
type: Literal["zip"]
bytes: bytes
ZenLoaderConfig: TypeAlias = Union[StaticLoaderConfig, FilesystemLoaderConfig, ZipLoaderConfig]
ZenLoaderCallback: TypeAlias = Callable[[str], Union[str, dict, ZenDecisionContent, Awaitable[Union[str, dict, ZenDecisionContent]]]]
class ZenEngineOptions(TypedDict, total=False):
loader: Union[ZenLoaderCallback, ZenLoaderConfig]
customHandler: Callable
class EvaluateBatchRequest(TypedDict):
key: str
context: Any
class EvaluateBatchResult(TypedDict, total=False):
success: bool
data: EvaluateResponse
error: Any
class ZenEngine:
def __init__(self, options: Optional[dict] = None) -> None: ...
def __init__(self, options: Optional[ZenEngineOptions] = None) -> None: ...
def evaluate(self, key: str, context: ZenContext,
options: Optional[DecisionEvaluateOptions] = None) -> EvaluateResponse: ...
def evaluate_batch(self, requests: list[EvaluateBatchRequest],
options: Optional[DecisionEvaluateOptions] = None) -> list[EvaluateBatchResult]: ...
def async_evaluate(self, key: str, context: ZenContext, options: Optional[DecisionEvaluateOptions] = None) -> \
Awaitable[EvaluateResponse]: ...