From 4ba5bf207b552dee2e3cca53b2f9318766550e55 Mon Sep 17 00:00:00 2001 From: stefan-gorules <127550877+stefan-gorules@users.noreply.github.com> Date: Wed, 28 May 2025 10:40:11 +0200 Subject: [PATCH] fix: py evalaute options (#358) --- bindings/python/Cargo.toml | 6 +++--- bindings/python/src/engine.rs | 22 ++++++++++++++++++++-- bindings/python/test_sync.py | 6 ++++++ 3 files changed, 29 insertions(+), 5 deletions(-) diff --git a/bindings/python/Cargo.toml b/bindings/python/Cargo.toml index 5c6cb4f1..809c83b2 100644 --- a/bindings/python/Cargo.toml +++ b/bindings/python/Cargo.toml @@ -12,9 +12,9 @@ crate-type = ["cdylib"] [dependencies] anyhow = { workspace = true } either = "1" -pyo3 = { version = "0.24", features = ["anyhow", "serde", "either"] } -pyo3-async-runtimes = { version = "0.24", features = ["tokio-runtime", "attributes"] } -pythonize = "0.24" +pyo3 = { version = "0.25", features = ["anyhow", "serde", "either"] } +pyo3-async-runtimes = { version = "0.25", features = ["tokio-runtime", "attributes"] } +pythonize = "0.25" json_dotpath = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } diff --git a/bindings/python/src/engine.rs b/bindings/python/src/engine.rs index 1ec9e7b6..455229ec 100644 --- a/bindings/python/src/engine.rs +++ b/bindings/python/src/engine.rs @@ -8,7 +8,7 @@ use crate::mt::{block_on, worker_pool}; use crate::value::PyValue; use crate::variable::PyVariable; use anyhow::{anyhow, Context}; -use pyo3::prelude::PyDictMethods; +use pyo3::prelude::{PyAnyMethods, PyDictMethods}; use pyo3::types::PyDict; use pyo3::{pyclass, pymethods, Bound, FromPyObject, IntoPyObjectExt, Py, PyAny, PyResult, Python}; use pyo3_async_runtimes::tokio::get_current_locals; @@ -23,12 +23,30 @@ pub struct PyZenEngine { engine: Arc>, } -#[derive(Serialize, Deserialize, FromPyObject)] +#[derive(Serialize, Deserialize)] pub struct PyZenEvaluateOptions { pub trace: Option, pub max_depth: Option, } +impl<'py> FromPyObject<'py> for PyZenEvaluateOptions { + fn extract_bound(ob: &Bound<'py, PyAny>) -> PyResult { + let dict = ob.downcast::()?; + + let trace = dict + .get_item("trace")? + .map(|v| v.extract::()) + .transpose()?; + + let max_depth = dict + .get_item("max_depth")? + .map(|v| v.extract::()) + .transpose()?; + + Ok(PyZenEvaluateOptions { trace, max_depth }) + } +} + impl Default for PyZenEvaluateOptions { fn default() -> Self { Self { diff --git a/bindings/python/test_sync.py b/bindings/python/test_sync.py index 629f9a33..b6f444c0 100644 --- a/bindings/python/test_sync.py +++ b/bindings/python/test_sync.py @@ -90,6 +90,12 @@ class ZenEngine(unittest.TestCase): engine.evaluate("http-function.json", {}) self.assertTrue(True) + def test_additional_options(self): + engine = zen.ZenEngine({"loader": loader, "customHandler": custom_handler}) + + engine.evaluate("sleep-function.json", {}, {"trace": True}) + self.assertTrue(True) + def test_evaluate_graphs(self): engine = zen.ZenEngine({"loader": graph_loader}) json_files = glob.glob("../../test-data/graphs/*.json")