fix: python asyncio (#217)

* fix: python asyncio

* fix imports

* fix: pytho asyncio

* fix fmt

* remove println

* revert to futures executor

* remove extension module
This commit is contained in:
stefan-gorules
2024-07-17 17:52:17 +02:00
committed by GitHub
parent 0b6e43c21b
commit 9e60fdc882
10 changed files with 147 additions and 196 deletions
+16 -19
View File
@@ -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<Option<PyObject>> 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)
}
}
+23 -33
View File
@@ -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<()> {
+29 -43
View File
@@ -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<PyZenDecision> {
@@ -157,14 +148,9 @@ impl PyZenEngine {
Ok(PyZenDecision::from(decision))
}
pub fn get_decision(&self, key: String) -> PyResult<PyZenDecision> {
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<PyZenDecision> {
let decision = futures::executor::block_on(self.graph.get_decision(&key))
.context("Failed to find decision with given key")?;
Ok(PyZenDecision::from(decision))
}
-1
View File
@@ -9,7 +9,6 @@ mod decision;
mod engine;
mod expression;
mod loader;
mod mt;
mod types;
mod value;
-48
View File
@@ -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<LocalPoolHandle> = OnceLock::new();
LOCAL_POOL
.get_or_init(|| LocalPoolHandle::new(parallelism()))
.clone()
}
pub(crate) fn spawn_worker<F, Fut>(create_task: F) -> impl Future<Output = Fut::Output>
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<F, Fut>(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")
})
})
}