Files
zen/bindings/python/src/custom_node.rs
T
stefan-gorulesandIvan Miletic 47812143a0 feat: passthrough nodes (#261)
* feat: passthrough nodes

* update bindings

* update model

* update exclusion

* fix merge strategy

* fix

* fix merge

* add tests

* fix walk

* fix merge

* fix

* fix: add tests

* fix types

---------

Co-authored-by: Ivan Miletic <vnmiletic@gmail.com>
2024-10-23 15:59:31 +02:00

61 lines
1.8 KiB
Rust

use anyhow::anyhow;
use either::Either;
use pyo3::types::PyDict;
use pyo3::{PyObject, PyResult, Python};
use pyo3_asyncio::tokio;
use pythonize::depythonize;
use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest};
use zen_engine::handler::node::{NodeResponse, NodeResult};
use crate::types::PyNodeRequest;
#[derive(Default)]
pub(crate) struct PyCustomNode(Option<PyObject>);
impl From<PyObject> for PyCustomNode {
fn from(value: PyObject) -> Self {
Self(Some(value))
}
}
impl From<Option<PyObject>> for PyCustomNode {
fn from(value: Option<PyObject>) -> Self {
Self(value)
}
}
fn extract_custom_node_response(py: Python<'_>, result: PyObject) -> 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 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(Either::Left(extract_custom_node_response(py, 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))
}
}
}
}