From 9e60fdc882b979b9a7c28fc51fdf41317d1d4796 Mon Sep 17 00:00:00 2001 From: stefan-gorules <127550877+stefan-gorules@users.noreply.github.com> Date: Wed, 17 Jul 2024 17:52:17 +0200 Subject: [PATCH] fix: python asyncio (#217) * fix: python asyncio * fix imports * fix: pytho asyncio * fix fmt * remove println * revert to futures executor * remove extension module --- Cargo.toml | 5 +- bindings/python/Cargo.toml | 4 +- bindings/python/src/custom_node.rs | 35 +++++------ bindings/python/src/decision.rs | 56 +++++++---------- bindings/python/src/engine.rs | 72 +++++++++------------- bindings/python/src/lib.rs | 1 - bindings/python/src/mt.rs | 48 --------------- bindings/python/test_async.py | 61 ++++++++++++++++++ bindings/python/{index.py => test_sync.py} | 59 ++++-------------- core/engine/Cargo.toml | 2 +- 10 files changed, 147 insertions(+), 196 deletions(-) delete mode 100644 bindings/python/src/mt.rs create mode 100644 bindings/python/test_async.py rename bindings/python/{index.py => test_sync.py} (54%) diff --git a/Cargo.toml b/Cargo.toml index eb6581d6..c58167db 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -35,4 +35,7 @@ thiserror = "1.0.50" [profile.release] lto = true codegen-units = 1 -strip = "symbols" \ No newline at end of file +strip = "symbols" + +[patch.crates-io] +rquickjs-core = { git = "https://github.com/stefan-gorules/rquickjs.git", branch = "master" } \ No newline at end of file diff --git a/bindings/python/Cargo.toml b/bindings/python/Cargo.toml index f937da8a..45ac24b2 100644 --- a/bindings/python/Cargo.toml +++ b/bindings/python/Cargo.toml @@ -11,14 +11,14 @@ crate-type = ["cdylib"] [dependencies] anyhow = { workspace = true } +either = "1.13" pyo3 = { version = "0.20", features = ["anyhow", "serde"] } pyo3-asyncio = { version = "0.20", features = ["tokio-runtime"] } pythonize = "0.20" json_dotpath = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } -tokio = { workspace = true, features = ["rt"] } -tokio-util = { workspace = true, features = ["rt"] } +futures = "0.3" 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 ac945bc4..fd23aff9 100644 --- a/bindings/python/src/custom_node.rs +++ b/bindings/python/src/custom_node.rs @@ -1,7 +1,8 @@ use anyhow::anyhow; +use either::Either; use pyo3::types::PyDict; use pyo3::{PyObject, PyResult, Python}; -use pyo3_asyncio::tokio::into_future; +use pyo3_asyncio::tokio; use pythonize::depythonize; use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; @@ -24,7 +25,7 @@ impl From> for PyCustomNode { } } -fn extract_custom_node_response(result: PyObject, py: Python<'_>) -> NodeResult { +fn extract_custom_node_response(py: Python<'_>, result: PyObject) -> NodeResult { let dict = result.extract::<&PyDict>(py)?; let response: NodeResponse = depythonize(dict)?; Ok(response) @@ -36,28 +37,24 @@ impl CustomNodeAdapter for PyCustomNode { return Err(anyhow!("Custom node handler not provided")); }; - let (future, result) = Python::with_gil(|py| -> PyResult<_> { + let maybe_result: PyResult<_> = Python::with_gil(|py| { let req = PyNodeRequest::from_request(py, request)?; let result = callable.call1(py, (req,))?; let is_coroutine = result.getattr(py, "__await__").is_ok(); - if is_coroutine { - return Ok((Some(into_future(result.as_ref(py))), None)); + if !is_coroutine { + return Ok(Either::Left(extract_custom_node_response(py, result))); } - Ok((None, Some(extract_custom_node_response(result, py)))) - })?; - if let Some(result) = result { - return result; + let result_future = tokio::into_future(result.as_ref(py))?; + return Ok(Either::Right(result_future)); + }); + + match maybe_result? { + Either::Left(result) => result, + Either::Right(future) => { + let result = future.await?; + Python::with_gil(|py| extract_custom_node_response(py, 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 9be366f3..5ef7d4c7 100644 --- a/bindings/python/src/decision.rs +++ b/bindings/python/src/decision.rs @@ -11,7 +11,6 @@ use zen_engine::{Decision, EvaluationOptions}; use crate::custom_node::PyCustomNode; use crate::engine::PyZenEvaluateOptions; use crate::loader::PyDecisionLoader; -use crate::mt::{spawn_worker, spawn_worker_blocking}; use crate::value::PyValue; #[pyclass] @@ -35,19 +34,15 @@ impl PyZenDecision { }; let decision = self.0.clone(); - let result = spawn_worker_blocking(move || async move { - decision - .evaluate_with_opts( - &context, - EvaluationOptions { - max_depth: options.max_depth, - trace: options.trace, - }, - ) - .await - .map_err(|e| { - anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) - }) + 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("Fail")?; @@ -68,27 +63,22 @@ impl PyZenDecision { }; let decision = self.0.clone(); + 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())) + })?; - tokio::future_into_py( - py, - spawn_worker(move || async move { - let result = decision - .evaluate_with_opts( - &context, - EvaluationOptions { - max_depth: options.max_depth, - trace: options.trace, - }, - ) - .await - .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 value = serde_json::to_value(result).context("Failed to serialize result")?; - Python::with_gil(|py| Ok(PyValue(value).to_object(py))) - }), - ) + Python::with_gil(|py| Ok(PyValue(value).to_object(py))) + }) } pub fn validate(&self) -> PyResult<()> { diff --git a/bindings/python/src/engine.rs b/bindings/python/src/engine.rs index 5f923a5f..2cfe8e9f 100644 --- a/bindings/python/src/engine.rs +++ b/bindings/python/src/engine.rs @@ -3,6 +3,7 @@ use std::sync::Arc; use anyhow::{anyhow, Context}; use pyo3::types::PyDict; use pyo3::{pyclass, pymethods, PyAny, PyObject, PyResult, Python, ToPyObject}; +use pyo3_asyncio::tokio; use pythonize::depythonize; use serde::{Deserialize, Serialize}; @@ -12,7 +13,6 @@ use zen_engine::{DecisionEngine, EvaluationOptions}; use crate::custom_node::PyCustomNode; use crate::decision::PyZenDecision; use crate::loader::PyDecisionLoader; -use crate::mt::{spawn_worker, spawn_worker_blocking}; use crate::value::PyValue; #[pyclass] @@ -90,20 +90,16 @@ impl PyZenEngine { }; let graph = self.graph.clone(); - let result = spawn_worker_blocking(move || async move { - graph - .evaluate_with_opts( - key, - &context, - EvaluationOptions { - max_depth: options.max_depth, - trace: options.trace, - }, - ) - .await - .map_err(|e| { - anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) - }) + 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")?; @@ -125,28 +121,23 @@ impl PyZenEngine { }; let graph = self.graph.clone(); - pyo3_asyncio::tokio::future_into_py( - py, - spawn_worker(move || async move { - let result = graph - .evaluate_with_opts( - key, - &context, - EvaluationOptions { - max_depth: options.max_depth, - trace: options.trace, - }, - ) - .await - .map_err(|e| { - anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) - })?; + 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")?; + let value = serde_json::to_value(result).context("Failed to serialize result")?; - Python::with_gil(|py| Ok(PyValue(value).to_object(py))) - }), - ) + Python::with_gil(|py| Ok(PyValue(value).to_object(py))) + }) } pub fn create_decision(&self, content: String) -> PyResult { @@ -157,14 +148,9 @@ impl PyZenEngine { Ok(PyZenDecision::from(decision)) } - pub fn get_decision(&self, key: String) -> PyResult { - let graph = self.graph.clone(); - let decision = spawn_worker_blocking(move || async move { - graph - .get_decision(&key) - .await - .context("Failed to find decision with given key") - })?; + pub fn get_decision<'py>(&'py self, py: Python<'py>, key: String) -> PyResult { + let decision = futures::executor::block_on(self.graph.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 2bca2d2a..29b377a4 100644 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -9,7 +9,6 @@ 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 deleted file mode 100644 index 24103fb8..00000000 --- a/bindings/python/src/mt.rs +++ /dev/null @@ -1,48 +0,0 @@ -use std::future::Future; -use std::sync::OnceLock; -use std::thread::available_parallelism; -use tokio::runtime::Handle; -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() -} - -pub(crate) fn spawn_worker(create_task: F) -> impl Future -where - F: FnOnce() -> Fut, - F: Send + 'static, - Fut: Future + 'static, - Fut::Output: Send + 'static, -{ - async move { - worker_pool() - .spawn_pinned(create_task) - .await - .expect("Thread panicked") - } -} - -pub(crate) fn spawn_worker_blocking(create_task: F) -> Fut::Output -where - F: FnOnce() -> Fut, - F: Send + 'static, - Fut: Future + 'static, - Fut::Output: Send + 'static, -{ - tokio::task::block_in_place(move || { - Handle::current().block_on(async move { - worker_pool() - .spawn_pinned(create_task) - .await - .expect("Thread panicked") - }) - }) -} diff --git a/bindings/python/test_async.py b/bindings/python/test_async.py new file mode 100644 index 00000000..5eaea028 --- /dev/null +++ b/bindings/python/test_async.py @@ -0,0 +1,61 @@ +import asyncio +import unittest + +import zen + + +def loader(key): + with open("../../test-data/" + key, "r") as f: + return f.read() + + +def custom_handler(request): + p1 = request.get_field("prop1") + return { + "output": {"sum": p1} + } + + +async def custom_async_handler(request): + p1 = request.get_field("prop1") + await asyncio.sleep(0.25) + return { + "output": {"sum": p1} + } + + +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) + + +if __name__ == '__main__': + unittest.main() diff --git a/bindings/python/index.py b/bindings/python/test_sync.py similarity index 54% rename from bindings/python/index.py rename to bindings/python/test_sync.py index 9d205108..f17b7ae6 100644 --- a/bindings/python/index.py +++ b/bindings/python/test_sync.py @@ -1,23 +1,19 @@ -import zen -import asyncio import unittest +import zen + + def loader(key): with open("../../test-data/" + key, "r") as f: return f.read() + def custom_handler(request): p1 = request.get_field("prop1") return { - "output": { "sum": p1 } + "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): @@ -55,7 +51,7 @@ class ZenEngine(unittest.TestCase): self.assertEqual(r["result"]["output"], 30) def test_engine_custom_handler(self): - engine = zen.ZenEngine({ "loader": loader, "customHandler": custom_handler }) + engine = zen.ZenEngine({"loader": loader, "customHandler": custom_handler}) r1 = engine.evaluate("custom.json", {"a": 10}) r2 = engine.evaluate("custom.json", {"a": 20}) r3 = engine.evaluate("custom.json", {"a": 30}) @@ -65,50 +61,17 @@ class ZenEngine(unittest.TestCase): self.assertEqual(r3["result"]["sum"], 40) def test_evaluate_expression(self): - result = zen.evaluate_expression("sum(a)", { "a": [1, 2, 3, 4] }) + result = zen.evaluate_expression("sum(a)", {"a": [1, 2, 3, 4]}) self.assertEqual(result, 10) def test_evaluate_unary_expression(self): - result = zen.evaluate_unary_expression("'FR', 'ES', 'GB'", { "$": "GB" }) + result = zen.evaluate_unary_expression("'FR', 'ES', 'GB'", {"$": "GB"}) self.assertEqual(result, True) def test_render_template(self): - result = zen.render_template("{{ a + b }}", { "a": 10, "b": 20 }) + 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() +if __name__ == '__main__': + unittest.main() diff --git a/core/engine/Cargo.toml b/core/engine/Cargo.toml index 9466cf4e..215b4767 100644 --- a/core/engine/Cargo.toml +++ b/core/engine/Cargo.toml @@ -22,7 +22,7 @@ json_dotpath = { workspace = true } fixedbitset = "0.4.2" tokio = { workspace = true, features = ["sync", "time"] } reqwest = { version = "0.12", features = ["json", "rustls-tls"], default-features = false } -rquickjs = { git = "https://github.com/stefan-gorules/rquickjs.git", branch = "master", features = ["macro", "loader", "rust-alloc", "futures", "either", "properties"] } +rquickjs = { version = "0.6.2", features = ["macro", "loader", "rust-alloc", "futures", "either", "properties"] } itertools = { workspace = true } zen-expression = { path = "../expression", version = "0.24.1" } zen-tmpl = { path = "../template", version = "0.24.1" }