mirror of
https://github.com/gorules/zen.git
synced 2026-10-04 00:02:18 +00:00
feat: implement new loaders across languages (#487)
This commit is contained in:
+128
-11
@@ -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,
|
||||
|
||||
@@ -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
@@ -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]: ...
|
||||
|
||||
|
||||
Reference in New Issue
Block a user