From 4c9c26b4296fb9b6d75a46215976fbe51637ba59 Mon Sep 17 00:00:00 2001 From: Scott Thompson Date: Thu, 4 Jul 2024 05:10:20 -0500 Subject: [PATCH] feat: add async support to python binding (#185) * feat: Add async support to pythong binding * fix: return NodeResult instead of panic when future missing --- bindings/python/Cargo.toml | 3 ++- bindings/python/index.py | 43 +++++++++++++++++++++++++++++- bindings/python/pyproject.toml | 2 +- bindings/python/src/custom_node.rs | 31 +++++++++++++++++---- bindings/python/src/decision.rs | 34 ++++++++++++++++++++++- bindings/python/src/engine.rs | 36 ++++++++++++++++++++++++- 6 files changed, 139 insertions(+), 10 deletions(-) diff --git a/bindings/python/Cargo.toml b/bindings/python/Cargo.toml index 213fb6b4..77d482a8 100644 --- a/bindings/python/Cargo.toml +++ b/bindings/python/Cargo.toml @@ -20,4 +20,5 @@ serde_json = { workspace = true } futures = { workspace = true } zen-engine = { path = "../../core/engine" } zen-expression = { path = "../../core/expression" } -zen-tmpl = { path = "../../core/template" } \ No newline at end of file +zen-tmpl = { path = "../../core/template" } +pyo3-asyncio = { version = "0.20.0", features = ["tokio-runtime"] } diff --git a/bindings/python/index.py b/bindings/python/index.py index bf8bad8a..9d205108 100644 --- a/bindings/python/index.py +++ b/bindings/python/index.py @@ -1,5 +1,5 @@ import zen -import time +import asyncio import unittest def loader(key): @@ -12,6 +12,13 @@ def custom_handler(request): "output": { "sum": p1 } } +async def custom_async_handler(request): + p1 = request.get_field("prop1") + await asyncio.sleep(1) + return { + "output": { "sum": p1 } + } + # The test based on unittest module class ZenEngine(unittest.TestCase): def test_decision_using_loader(self): @@ -69,5 +76,39 @@ class ZenEngine(unittest.TestCase): result = zen.render_template("{{ a + b }}", { "a": 10, "b": 20 }) self.assertEqual(result, 30) + +class AsyncZenEngine(unittest.IsolatedAsyncioTestCase): + async def test_async_evaluate(self): + engine = zen.ZenEngine({"loader": loader}) + r1 = engine.async_evaluate("function.json", {"input": 5}) + r2 = engine.async_evaluate("table.json", {"input": 2}) + r3 = engine.async_evaluate("table.json", {"input": 12}) + + results = await asyncio.gather(r1, r2, r3) + self.assertEqual(results[0]["result"]["output"], 10) + self.assertEqual(results[1]["result"]["output"], 0) + self.assertEqual(results[2]["result"]["output"], 10) + + async def test_async_evaluate_custom_handler(self): + engine = zen.ZenEngine({"loader": loader, "customHandler": custom_async_handler}) + r1 = engine.async_evaluate("custom.json", {"a": 10}) + r2 = engine.async_evaluate("custom.json", {"a": 20}) + r3 = engine.async_evaluate("custom.json", {"a": 30}) + + results = await asyncio.gather(r1, r2, r3) + self.assertEqual(results[0]["result"]["sum"], 20) + self.assertEqual(results[1]["result"]["sum"], 30) + self.assertEqual(results[2]["result"]["sum"], 40) + + async def test_create_decisions_from_content(self): + engine = zen.ZenEngine() + with open("../../test-data/function.json", "r") as f: + functionContent = f.read() + functionDecision = engine.create_decision(functionContent) + + r = await functionDecision.async_evaluate({"input": 15}) + self.assertEqual(r["result"]["output"], 30) + + # run the test unittest.main() diff --git a/bindings/python/pyproject.toml b/bindings/python/pyproject.toml index 52478a8b..9ba1ff59 100644 --- a/bindings/python/pyproject.toml +++ b/bindings/python/pyproject.toml @@ -32,7 +32,7 @@ keywords = ["gorules", ] [project.optional-dependencies] -dev = ["black", "bumpver", "isort", "pip-tools", "pytest"] +dev = ["black", "bumpver", "isort", "pip-tools", "pytest", "asyncio"] [project.urls] Homepage = "https://github.com/gorules/zen" diff --git a/bindings/python/src/custom_node.rs b/bindings/python/src/custom_node.rs index 64bc41ee..ac945bc4 100644 --- a/bindings/python/src/custom_node.rs +++ b/bindings/python/src/custom_node.rs @@ -1,6 +1,7 @@ use anyhow::anyhow; use pyo3::types::PyDict; -use pyo3::{PyObject, Python}; +use pyo3::{PyObject, PyResult, Python}; +use pyo3_asyncio::tokio::into_future; use pythonize::depythonize; use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; @@ -23,20 +24,40 @@ impl From> for PyCustomNode { } } +fn extract_custom_node_response(result: PyObject, py: Python<'_>) -> NodeResult { + let dict = result.extract::<&PyDict>(py)?; + let response: NodeResponse = depythonize(dict)?; + Ok(response) +} + impl CustomNodeAdapter for PyCustomNode { async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { let Some(callable) = &self.0 else { return Err(anyhow!("Custom node handler not provided")); }; - let content: NodeResponse = Python::with_gil(|py| { + let (future, result) = Python::with_gil(|py| -> PyResult<_> { let req = PyNodeRequest::from_request(py, request)?; let result = callable.call1(py, (req,))?; - - let dict = result.extract::<&PyDict>(py)?; - depythonize(dict) + let is_coroutine = result.getattr(py, "__await__").is_ok(); + if is_coroutine { + return Ok((Some(into_future(result.as_ref(py))), None)); + } + Ok((None, Some(extract_custom_node_response(result, py)))) })?; + if let Some(result) = result { + return result; + } + + let result = future + .ok_or_else(|| anyhow!("Future or result must be present"))?? + .await?; + + let content = Python::with_gil(|py| -> PyResult<_> { + Ok(extract_custom_node_response(result, py)) + })??; + Ok(content) } } diff --git a/bindings/python/src/decision.rs b/bindings/python/src/decision.rs index 900c736f..3328e346 100644 --- a/bindings/python/src/decision.rs +++ b/bindings/python/src/decision.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use anyhow::{anyhow, Context}; use pyo3::types::PyDict; -use pyo3::{pyclass, pymethods, PyObject, PyResult, Python, ToPyObject}; +use pyo3::{pyclass, pymethods, PyAny, PyObject, PyResult, Python, ToPyObject}; use pythonize::depythonize; use zen_engine::{Decision, EvaluationOptions}; @@ -48,6 +48,38 @@ impl PyZenDecision { Ok(PyValue(value).to_object(py)) } + pub fn async_evaluate<'py>( + &'py self, + py: Python<'py>, + ctx: &PyDict, + opts: Option<&PyDict>, + ) -> PyResult<&PyAny> { + let context = depythonize(ctx).context("Failed to convert dict")?; + let options: PyZenEvaluateOptions = if let Some(op) = opts { + depythonize(op).context("Failed to convert dict")? + } else { + Default::default() + }; + + let decision = self.0.clone(); + pyo3_asyncio::tokio::future_into_py(py, async move { + let result = futures::executor::block_on(decision.evaluate_with_opts( + &context, + EvaluationOptions { + max_depth: options.max_depth, + trace: options.trace, + }, + )) + .map_err(|e| { + anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) + })?; + + let value = serde_json::to_value(result).context("Failed to serialize result")?; + + Python::with_gil(|py| Ok(PyValue(value).to_object(py))) + }) + } + pub fn validate(&self) -> PyResult<()> { let decision = self.0.clone(); decision diff --git a/bindings/python/src/engine.rs b/bindings/python/src/engine.rs index 338b7a2e..2f94657f 100644 --- a/bindings/python/src/engine.rs +++ b/bindings/python/src/engine.rs @@ -4,7 +4,7 @@ use crate::loader::PyDecisionLoader; use crate::value::PyValue; use anyhow::{anyhow, Context}; use pyo3::types::PyDict; -use pyo3::{pyclass, pymethods, PyObject, PyResult, Python, ToPyObject}; +use pyo3::{pyclass, pymethods, PyAny, PyObject, PyResult, Python, ToPyObject}; use pythonize::depythonize; use serde::{Deserialize, Serialize}; use std::sync::Arc; @@ -102,6 +102,40 @@ impl PyZenEngine { Ok(PyValue(value).to_object(py)) } + pub fn async_evaluate<'py>( + &'py self, + py: Python<'py>, + key: String, + ctx: &PyDict, + opts: Option<&PyDict>, + ) -> PyResult<&PyAny> { + let context = depythonize(ctx).context("Failed to convert dict")?; + let options: PyZenEvaluateOptions = if let Some(op) = opts { + depythonize(op).context("Failed to convert dict")? + } else { + Default::default() + }; + + let graph = self.graph.clone(); + pyo3_asyncio::tokio::future_into_py(py, async move { + let result = futures::executor::block_on(graph.evaluate_with_opts( + key, + &context, + EvaluationOptions { + max_depth: options.max_depth, + trace: options.trace, + }, + )) + .map_err(|e| { + anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) + })?; + + let value = serde_json::to_value(result).context("Failed to serialize result")?; + + Python::with_gil(|py| Ok(PyValue(value).to_object(py))) + }) + } + pub fn create_decision(&self, content: String) -> PyResult { let decision_content: DecisionContent = serde_json::from_str(&content).context("Failed to serialize decision content")?;