From 55a632058075a911deceb6130010ed26e90c2210 Mon Sep 17 00:00:00 2001 From: stefan-gorules <127550877+stefan-gorules@users.noreply.github.com> Date: Sat, 8 Feb 2025 17:28:58 +0100 Subject: [PATCH] fix: python async (#309) * fix: python async * fix python async --- bindings/python/Cargo.toml | 4 +- bindings/python/src/custom_node.rs | 34 +++++++++------ bindings/python/src/decision.rs | 59 +++++++++++++++---------- bindings/python/src/engine.rs | 69 +++++++++++++++++------------- bindings/python/src/lib.rs | 1 + bindings/python/src/mt.rs | 25 +++++++++++ bindings/python/test_async.py | 15 ++++++- bindings/python/test_sync.py | 12 ++++++ test-data/http-function.json | 34 +++++++++++++++ test-data/sleep-function.json | 34 +++++++++++++++ 10 files changed, 219 insertions(+), 68 deletions(-) create mode 100644 bindings/python/src/mt.rs create mode 100644 test-data/http-function.json create mode 100644 test-data/sleep-function.json diff --git a/bindings/python/Cargo.toml b/bindings/python/Cargo.toml index aa56f6dd..a98e061e 100644 --- a/bindings/python/Cargo.toml +++ b/bindings/python/Cargo.toml @@ -13,12 +13,12 @@ crate-type = ["cdylib"] anyhow = { workspace = true } either = "1.13" pyo3 = { version = "0.23", features = ["anyhow", "serde"] } -pyo3-async-runtimes = { version = "0.23", features = ["tokio-runtime"] } +pyo3-async-runtimes = { version = "0.23", features = ["tokio-runtime", "attributes"] } pythonize = "0.23" json_dotpath = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } -futures = "0.3" +tokio-util = { version = "0.7", features = ["rt"] } zen-engine = { path = "../../core/engine" } zen-expression = { path = "../../core/expression" } zen-tmpl = { path = "../../core/template" } diff --git a/bindings/python/src/custom_node.rs b/bindings/python/src/custom_node.rs index 86ce9192..75093c7c 100644 --- a/bindings/python/src/custom_node.rs +++ b/bindings/python/src/custom_node.rs @@ -2,7 +2,7 @@ use anyhow::anyhow; use either::Either; use pyo3::types::PyDict; use pyo3::{Bound, IntoPyObjectExt, Py, PyAny, PyObject, PyResult, Python}; -use pyo3_async_runtimes::tokio; +use pyo3_async_runtimes::TaskLocals; use pythonize::depythonize; use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; @@ -11,17 +11,17 @@ use zen_engine::handler::node::{NodeResponse, NodeResult}; use crate::types::PyNodeRequest; #[derive(Default)] -pub(crate) struct PyCustomNode(Option>); - -impl From> for PyCustomNode { - fn from(value: Py) -> Self { - Self(Some(value)) - } +pub(crate) struct PyCustomNode { + callback: Option>, + task_locals: Option, } -impl From> for PyCustomNode { - fn from(value: Option) -> Self { - Self(value) +impl PyCustomNode { + pub fn new(callback: Option>, task_locals: Option) -> Self { + Self { + callback, + task_locals, + } } } @@ -33,7 +33,7 @@ fn extract_custom_node_response(py: Python<'_>, result: PyObject) -> NodeResult impl CustomNodeAdapter for PyCustomNode { async fn handle(&self, request: CustomNodeRequest) -> NodeResult { - let Some(callable) = &self.0 else { + let Some(callable) = &self.callback else { return Err(anyhow!("Custom node handler not provided")); }; @@ -45,8 +45,16 @@ impl CustomNodeAdapter for PyCustomNode { return Ok(Either::Left(extract_custom_node_response(py, result))); } - let result_future = tokio::into_future(result.into_bound_py_any(py)?)?; - return Ok(Either::Right(result_future)); + let Some(task_locals) = &self.task_locals else { + Err(anyhow!("Task locals are required in async context"))? + }; + + let result_future = pyo3_async_runtimes::into_future_with_locals( + task_locals, + result.into_bound_py_any(py)?, + )?; + + Ok(Either::Right(result_future)) }); match maybe_result? { diff --git a/bindings/python/src/decision.rs b/bindings/python/src/decision.rs index f00c4088..a653ea94 100644 --- a/bindings/python/src/decision.rs +++ b/bindings/python/src/decision.rs @@ -4,6 +4,8 @@ use anyhow::{anyhow, Context}; use pyo3::types::PyDict; use pyo3::{pyclass, pymethods, Bound, IntoPyObjectExt, Py, PyAny, PyResult, Python}; use pyo3_async_runtimes::tokio; +use pyo3_async_runtimes::tokio::get_current_locals; +use pyo3_async_runtimes::tokio::re_exports::runtime::Runtime; use pythonize::depythonize; use serde_json::Value; use zen_engine::{Decision, EvaluationOptions}; @@ -11,6 +13,7 @@ use zen_engine::{Decision, EvaluationOptions}; use crate::custom_node::PyCustomNode; use crate::engine::PyZenEvaluateOptions; use crate::loader::PyDecisionLoader; +use crate::mt::worker_pool; use crate::value::PyValue; #[pyclass] @@ -40,16 +43,19 @@ impl PyZenDecision { }; let decision = self.0.clone(); - let result = futures::executor::block_on(decision.evaluate_with_opts( - context.into(), - 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 rt = Runtime::new()?; + let result = rt + .block_on(decision.evaluate_with_opts( + context.into(), + 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("Fail")?; PyValue(value).into_py_any(py) @@ -70,19 +76,26 @@ impl PyZenDecision { }; let decision = self.0.clone(); - let result = tokio::future_into_py(py, async move { - let result = futures::executor::block_on(decision.evaluate_with_opts( - context.into(), - 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")?; + let result = tokio::future_into_py_with_locals(py, get_current_locals(py)?, async move { + let value = worker_pool() + .spawn_pinned(move || async move { + decision + .evaluate_with_opts( + context.into(), + EvaluationOptions { + max_depth: options.max_depth, + trace: options.trace, + }, + ) + .await + .map(serde_json::to_value) + }) + .await + .context("Failed to join threads")? + .map_err(|e| { + anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) + })? + .context("Failed to serialize result")?; Python::with_gil(|py| PyValue(value).into_py_any(py)) })?; diff --git a/bindings/python/src/engine.rs b/bindings/python/src/engine.rs index 409b68e0..50fd969c 100644 --- a/bindings/python/src/engine.rs +++ b/bindings/python/src/engine.rs @@ -4,7 +4,8 @@ use anyhow::{anyhow, Context}; use pyo3::prelude::PyDictMethods; use pyo3::types::PyDict; use pyo3::{pyclass, pymethods, Bound, IntoPyObjectExt, Py, PyAny, PyResult, Python}; -use pyo3_async_runtimes::tokio; +use pyo3_async_runtimes::tokio::get_current_locals; +use pyo3_async_runtimes::{tokio, TaskLocals}; use pythonize::depythonize; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -14,12 +15,13 @@ use zen_engine::{DecisionEngine, EvaluationOptions}; use crate::custom_node::PyCustomNode; use crate::decision::PyZenDecision; use crate::loader::PyDecisionLoader; +use crate::mt::{block_on, worker_pool}; use crate::value::PyValue; #[pyclass] #[pyo3(name = "ZenEngine")] pub struct PyZenEngine { - graph: Arc>, + engine: Arc>, } #[derive(Serialize, Deserialize)] @@ -40,11 +42,10 @@ impl Default for PyZenEvaluateOptions { impl Default for PyZenEngine { fn default() -> Self { Self { - graph: DecisionEngine::new( + engine: Arc::new(DecisionEngine::new( Arc::new(PyDecisionLoader::default()), Arc::new(PyCustomNode::default()), - ) - .into(), + )), } } } @@ -53,7 +54,7 @@ impl Default for PyZenEngine { impl PyZenEngine { #[new] #[pyo3(signature = (maybe_options=None))] - pub fn new(maybe_options: Option<&Bound<'_, PyDict>>) -> PyResult { + pub fn new(py: Python, maybe_options: Option<&Bound<'_, PyDict>>) -> PyResult { let Some(options) = maybe_options else { return Ok(Default::default()); }; @@ -68,12 +69,16 @@ impl PyZenEngine { None => None, }; + let task_locals = TaskLocals::with_running_loop(py) + .ok() + .map(|s| s.copy_context(py).ok()) + .flatten(); + Ok(Self { - graph: DecisionEngine::new( + engine: Arc::new(DecisionEngine::new( Arc::new(PyDecisionLoader::from(loader)), - Arc::new(PyCustomNode::from(custom_node)), - ) - .into(), + Arc::new(PyCustomNode::new(custom_node, task_locals)), + )), }) } @@ -92,8 +97,7 @@ impl PyZenEngine { Default::default() }; - let graph = self.graph.clone(); - let result = futures::executor::block_on(graph.evaluate_with_opts( + let result = block_on(self.engine.evaluate_with_opts( key, context.into(), EvaluationOptions { @@ -124,21 +128,28 @@ impl PyZenEngine { Default::default() }; - let graph = self.graph.clone(); - let result = tokio::future_into_py(py, async move { - let result = futures::executor::block_on(graph.evaluate_with_opts( - key, - context.into(), - 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")?; + let engine = self.engine.clone(); + let result = tokio::future_into_py_with_locals(py, get_current_locals(py)?, async move { + let value = worker_pool() + .spawn_pinned(move || async move { + engine + .evaluate_with_opts( + key, + context.into(), + EvaluationOptions { + max_depth: options.max_depth, + trace: options.trace, + }, + ) + .await + .map(serde_json::to_value) + }) + .await + .context("Failed to join threads")? + .map_err(|e| { + anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) + })? + .context("Failed to serialize result")?; Python::with_gil(|py| PyValue(value).into_py_any(py)) })?; @@ -150,12 +161,12 @@ impl PyZenEngine { let decision_content: DecisionContent = serde_json::from_str(&content).context("Failed to serialize decision content")?; - let decision = self.graph.create_decision(decision_content.into()); + let decision = self.engine.create_decision(decision_content.into()); Ok(PyZenDecision::from(decision)) } pub fn get_decision<'py>(&'py self, _py: Python<'py>, key: String) -> PyResult { - let decision = futures::executor::block_on(self.graph.get_decision(&key)) + let decision = block_on(self.engine.get_decision(&key)) .context("Failed to find decision with given key")?; Ok(PyZenDecision::from(decision)) diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index 5cf6ef70..a7c5dfc4 100644 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -10,6 +10,7 @@ mod decision; mod engine; mod expression; mod loader; +mod mt; mod types; mod value; diff --git a/bindings/python/src/mt.rs b/bindings/python/src/mt.rs new file mode 100644 index 00000000..6309fba7 --- /dev/null +++ b/bindings/python/src/mt.rs @@ -0,0 +1,25 @@ +use pyo3_async_runtimes::tokio::re_exports::runtime::Runtime; +use std::sync::OnceLock; +use std::thread::available_parallelism; +use tokio_util::task::LocalPoolHandle; + +fn parallelism() -> usize { + available_parallelism().map(Into::into).unwrap_or(1) +} + +pub(crate) fn worker_pool() -> LocalPoolHandle { + static LOCAL_POOL: OnceLock = OnceLock::new(); + LOCAL_POOL + .get_or_init(|| LocalPoolHandle::new(parallelism())) + .clone() +} + +static RUNTIME: OnceLock = OnceLock::new(); + +fn get_runtime() -> &'static Runtime { + RUNTIME.get_or_init(|| Runtime::new().unwrap()) +} + +pub(crate) fn block_on(future: F) -> F::Output { + get_runtime().block_on(future) +} diff --git a/bindings/python/test_async.py b/bindings/python/test_async.py index 543d2a9e..65496dfb 100644 --- a/bindings/python/test_async.py +++ b/bindings/python/test_async.py @@ -2,6 +2,7 @@ import asyncio import glob import json import os.path +import time import unittest import zen @@ -26,7 +27,7 @@ def custom_handler(request): async def custom_async_handler(request): p1 = request.get_field("prop1") - await asyncio.sleep(0.25) + await asyncio.sleep(0.1) return { "output": {"sum": p1} } @@ -55,6 +56,18 @@ class AsyncZenEngine(unittest.IsolatedAsyncioTestCase): self.assertEqual(results[1]["result"]["sum"], 30) self.assertEqual(results[2]["result"]["sum"], 40) + async def test_async_sleep_function(self): + engine = zen.ZenEngine({"loader": loader, "customHandler": custom_async_handler}) + + await engine.async_evaluate("sleep-function.json", {}) + self.assertTrue(True) + + async def test_async_http_function(self): + engine = zen.ZenEngine({"loader": loader, "customHandler": custom_async_handler}) + + await engine.async_evaluate("http-function.json", {}) + self.assertTrue(True) + async def test_create_decisions_from_content(self): engine = zen.ZenEngine() with open("../../test-data/function.json", "r") as f: diff --git a/bindings/python/test_sync.py b/bindings/python/test_sync.py index 512de607..629f9a33 100644 --- a/bindings/python/test_sync.py +++ b/bindings/python/test_sync.py @@ -78,6 +78,18 @@ class ZenEngine(unittest.TestCase): result = zen.render_template("{{ a + b }}", {"a": 10, "b": 20}) self.assertEqual(result, 30) + def test_sleep_function(self): + engine = zen.ZenEngine({"loader": loader, "customHandler": custom_handler}) + + engine.evaluate("sleep-function.json", {}) + self.assertTrue(True) + + def test_http_function(self): + engine = zen.ZenEngine({"loader": loader, "customHandler": custom_handler}) + + engine.evaluate("http-function.json", {}) + self.assertTrue(True) + def test_evaluate_graphs(self): engine = zen.ZenEngine({"loader": graph_loader}) json_files = glob.glob("../../test-data/graphs/*.json") diff --git a/test-data/http-function.json b/test-data/http-function.json new file mode 100644 index 00000000..091e04d1 --- /dev/null +++ b/test-data/http-function.json @@ -0,0 +1,34 @@ +{ + "contentType": "application/vnd.gorules.decision", + "nodes": [ + { + "type": "inputNode", + "id": "2d560c7d-3528-43ed-88d8-28f4d8ef17be", + "name": "request", + "position": { + "x": 295, + "y": 175 + } + }, + { + "type": "functionNode", + "content": { + "source": "import zen from 'zen';\nimport http from 'http';\n\n/** @type {Handler} **/\nexport const handler = async (input) => {\n const r = await http.get('https://fakestoreapi.com/products/1');\n\n return r;\n};\n" + }, + "id": "66cb12f4-4cdd-4422-850b-4534f959407d", + "name": "function1", + "position": { + "x": 600, + "y": 175 + } + } + ], + "edges": [ + { + "id": "b65ca09a-a010-4fc2-bca1-f7cf21a0ddc3", + "sourceId": "2d560c7d-3528-43ed-88d8-28f4d8ef17be", + "type": "edge", + "targetId": "66cb12f4-4cdd-4422-850b-4534f959407d" + } + ] +} \ No newline at end of file diff --git a/test-data/sleep-function.json b/test-data/sleep-function.json new file mode 100644 index 00000000..eda6cfb6 --- /dev/null +++ b/test-data/sleep-function.json @@ -0,0 +1,34 @@ +{ + "contentType": "application/vnd.gorules.decision", + "nodes": [ + { + "type": "inputNode", + "id": "2d560c7d-3528-43ed-88d8-28f4d8ef17be", + "name": "request", + "position": { + "x": 295, + "y": 175 + } + }, + { + "type": "functionNode", + "content": { + "source": "import zen from 'zen';\nimport http from 'http';\n\n/** @type {Handler} **/\nexport const handler = async (input) => {\n await console.sleep(50);\n\n return { hello: 'world' };\n};\n" + }, + "id": "66cb12f4-4cdd-4422-850b-4534f959407d", + "name": "function1", + "position": { + "x": 600, + "y": 175 + } + } + ], + "edges": [ + { + "id": "b65ca09a-a010-4fc2-bca1-f7cf21a0ddc3", + "sourceId": "2d560c7d-3528-43ed-88d8-28f4d8ef17be", + "type": "edge", + "targetId": "66cb12f4-4cdd-4422-850b-4534f959407d" + } + ] +} \ No newline at end of file