fix: py evalaute options (#358)

This commit is contained in:
stefan-gorules
2025-05-28 10:40:11 +02:00
committed by GitHub
parent 38fcaec7ee
commit 4ba5bf207b
3 changed files with 29 additions and 5 deletions
+3 -3
View File
@@ -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 }
+20 -2
View File
@@ -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<DecisionEngine<PyDecisionLoader, PyCustomNode>>,
}
#[derive(Serialize, Deserialize, FromPyObject)]
#[derive(Serialize, Deserialize)]
pub struct PyZenEvaluateOptions {
pub trace: Option<bool>,
pub max_depth: Option<u8>,
}
impl<'py> FromPyObject<'py> for PyZenEvaluateOptions {
fn extract_bound(ob: &Bound<'py, PyAny>) -> PyResult<Self> {
let dict = ob.downcast::<PyDict>()?;
let trace = dict
.get_item("trace")?
.map(|v| v.extract::<bool>())
.transpose()?;
let max_depth = dict
.get_item("max_depth")?
.map(|v| v.extract::<u8>())
.transpose()?;
Ok(PyZenEvaluateOptions { trace, max_depth })
}
}
impl Default for PyZenEvaluateOptions {
fn default() -> Self {
Self {
+6
View File
@@ -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")