This commit is contained in:
Quanzheng Long
2026-03-17 18:24:37 -07:00
parent e13004da77
commit 90865e2af9
8 changed files with 410 additions and 456 deletions
+7
View File
@@ -100,3 +100,10 @@ dmypy.json
.turbo
.editorconfig
.scratch
# macOS debug symbol bundles generated during local Rust builds
saf-python-sdk/python/saf_python_sdk/*.dSYM/
# Local PyO3 extension artifacts for saf-python-sdk
saf-python-sdk/python/saf_python_sdk/langgraph_rust_core*.so
saf-python-sdk/python/saf_python_sdk/langgraph_rust_core*.pyd
+20 -5
View File
@@ -116,8 +116,9 @@ func (c *Context) SendCustomStreamEvent(value any) error {
}
type Handler[StateT any] struct {
engine *RustEngine
done chan resultOrErr[StateT]
engine *RustEngine
done chan resultOrErr[StateT]
streamReadyC chan struct{}
}
type resultOrErr[StateT any] struct {
@@ -135,6 +136,7 @@ func (h *Handler[StateT]) WaitForResult() (StateT, error) {
}
func (h *Handler[StateT]) ReceiveStream() (any, error) {
<-h.streamReadyC
event, hasEvent, err := h.engine.ReceiveStream()
if err != nil {
return nil, err
@@ -146,6 +148,7 @@ func (h *Handler[StateT]) ReceiveStream() (any, error) {
}
func (h *Handler[StateT]) CloseStream() error {
<-h.streamReadyC
return h.engine.CloseStream()
}
@@ -165,15 +168,27 @@ func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT, str
}
handler := &Handler[StateT]{
engine: engine,
done: make(chan resultOrErr[StateT], 1),
engine: engine,
done: make(chan resultOrErr[StateT], 1),
streamReadyC: make(chan struct{}),
}
go func() {
defer engine.Close()
streamModeForRun := resolvedStreamMode
if resolvedStreamMode != "" {
if err := engine.StartStream(resolvedStreamMode); err != nil {
close(handler.streamReadyC)
handler.done <- resultOrErr[StateT]{err: err}
close(handler.done)
return
}
streamModeForRun = ""
}
close(handler.streamReadyC)
rawState, err := engine.RunGraph(
g.entryPoint,
g.finishPoint,
resolvedStreamMode,
streamModeForRun,
initialState,
initialInput,
func(node string, nodeInput any, fallbackState map[string]any) (Command, error) {
+250
View File
@@ -491,3 +491,253 @@ impl Engine {
}))
}
}
#[derive(Debug, Deserialize)]
struct CallbackSendPayloadJson {
node: String,
#[serde(default)]
arg: Value,
}
#[derive(Debug, Deserialize)]
struct CallbackNodeExecResultJsonWire {
update: Option<Value>,
#[serde(default)]
sends: Vec<CallbackSendPayloadJson>,
}
#[derive(Debug, Deserialize)]
struct CallbackEnvelopeIn {
ok: bool,
#[serde(default)]
payload: Option<CallbackNodeExecResultJsonWire>,
#[serde(default)]
suspend: Option<WaitRequest>,
#[serde(default)]
error: Option<String>,
}
enum SchedulerEvent<U, A> {
Node(Result<NodeExecution<U, A>, String>),
Resume {
node: String,
arg: A,
event: WaitEvent,
},
WaitError(String),
}
struct NodeExecution<U, A> {
node: String,
arg: A,
outcome: NodeOutcome<U, A>,
}
fn spawn_node_task<State, U, A, F>(
node: String,
arg: A,
state_snapshot: State,
tx: tokio_mpsc::UnboundedSender<SchedulerEvent<U, A>>,
callback: Arc<F>,
) -> Result<(), String>
where
State: Send + 'static,
U: Send + 'static,
A: Clone + Send + 'static,
F: Fn(String, A, State) -> Result<NodeOutcome<U, A>, String> + Send + Sync + 'static,
{
node_pool_execute(move || {
let node_for_result = node.clone();
let arg_for_result = arg.clone();
let result = callback(node, arg, state_snapshot).map(|outcome| NodeExecution {
node: node_for_result,
arg: arg_for_result,
outcome,
});
let _ = tx.send(SchedulerEvent::Node(result));
})
}
pub async fn run_graph_with_callback<State, U, A, FCallback, FMerge, FWrap>(
entry_point: String,
finish_point: String,
initial_state: State,
initial_input: A,
engine: Engine,
callback: FCallback,
merge_update: FMerge,
wrap_resume_arg: FWrap,
) -> Result<State, String>
where
State: Clone + Send + 'static,
U: Send + 'static,
A: Clone + Send + 'static,
FCallback: Fn(String, A, State) -> Result<NodeOutcome<U, A>, String> + Send + Sync + 'static,
FMerge: Fn(&mut State, Option<U>) -> Result<(), String> + Send + Sync + 'static,
FWrap: Fn(A, WaitEvent) -> Result<A, String> + Send + Sync + 'static,
{
let callback = Arc::new(callback);
let merge_update = Arc::new(merge_update);
let wrap_resume_arg = Arc::new(wrap_resume_arg);
let (tx, mut rx) = tokio_mpsc::unbounded_channel::<SchedulerEvent<U, A>>();
let state = Arc::new(StdMutex::new(initial_state));
let tx_for_spawn = tx.clone();
let state_for_spawn = Arc::clone(&state);
let mut active: usize = 1;
let mut waiting: usize = 0;
spawn_node_task(
entry_point,
initial_input,
state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone(),
tx_for_spawn.clone(),
Arc::clone(&callback),
)?;
while active > 0 || waiting > 0 {
let evt = rx
.recv()
.await
.ok_or_else(|| "scheduler event channel closed".to_string())?;
match evt {
SchedulerEvent::Node(result) => {
active = active.saturating_sub(1);
let exec = result?;
match exec.outcome {
NodeOutcome::Completed(node_result) => {
let mut guard = state.lock().expect("state mutex poisoned");
merge_update(&mut guard, node_result.update)?;
drop(guard);
if exec.node == finish_point {
break;
}
for send in node_result.sends {
active += 1;
let snapshot = state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone();
spawn_node_task(
send.node,
send.arg,
snapshot,
tx_for_spawn.clone(),
Arc::clone(&callback),
)?;
}
}
NodeOutcome::Suspended { wait } => {
waiting += 1;
let tx_wait = tx_for_spawn.clone();
let node = exec.node;
let arg = exec.arg;
let engine_for_wait = engine.clone();
tokio::spawn(async move {
match engine_for_wait.wait_request_async(&wait).await {
Ok(event) => {
let _ =
tx_wait.send(SchedulerEvent::Resume { node, arg, event });
}
Err(e) => {
let _ = tx_wait.send(SchedulerEvent::WaitError(e));
}
}
});
}
}
}
SchedulerEvent::Resume { node, arg, event } => {
waiting = waiting.saturating_sub(1);
active += 1;
let snapshot = state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone();
let resume_arg = wrap_resume_arg(arg, event)?;
spawn_node_task(
node,
resume_arg,
snapshot,
tx_for_spawn.clone(),
Arc::clone(&callback),
)?;
}
SchedulerEvent::WaitError(e) => return Err(e),
}
}
let final_state = state.lock().expect("state mutex poisoned").clone();
Ok(final_state)
}
pub async fn run_graph_json_with_callback<F>(
entry_point: String,
finish_point: String,
initial_state: Value,
initial_input: Value,
engine: Engine,
callback: F,
) -> Result<Value, String>
where
F: Fn(String, Value, Value) -> Result<NodeOutcome<Value, Value>, String>
+ Send
+ Sync
+ 'static,
{
run_graph_with_callback(
entry_point,
finish_point,
initial_state,
initial_input,
engine,
callback,
|state: &mut Value, update: Option<Value>| {
merge_json_update(state, update);
Ok(())
},
|arg: Value, event: WaitEvent| {
Ok(serde_json::json!({
"__lg_resume_arg__": arg,
"__lg_resume_event__": event,
}))
},
)
.await
}
pub fn parse_callback_envelope_json(
raw: &str,
node_name: &str,
) -> Result<NodeOutcome<Value, Value>, String> {
let parsed: CallbackEnvelopeIn = serde_json::from_str(raw)
.map_err(|e| format!("decode callback envelope for `{node_name}` failed: {e}"))?;
if !parsed.ok {
return Err(parsed
.error
.unwrap_or_else(|| format!("callback reported error for `{node_name}`")));
}
if let Some(wait) = parsed.suspend {
return Ok(NodeOutcome::Suspended { wait });
}
let payload = parsed
.payload
.ok_or_else(|| format!("callback payload missing for `{node_name}`"))?;
let sends = payload
.sends
.into_iter()
.map(|s| SendPayload {
node: s.node,
arg: s.arg,
})
.collect();
Ok(NodeOutcome::Completed(NodeExecResult {
update: payload.update,
sends,
}))
}
+42 -250
View File
@@ -1,39 +1,11 @@
use crate::engine::{
merge_json_update, node_pool_execute, run_loop_block_on, run_loop_spawn, AnyOfCondition,
Engine, NodeExecResult, NodeOutcome, SendPayload, WaitEvent, WaitRequest,
parse_callback_envelope_json, run_graph_json_with_callback, run_loop_block_on, run_loop_spawn,
AnyOfCondition, Engine, NodeOutcome,
};
use serde::Deserialize;
use serde_json::Value;
use std::ffi::{CStr, CString};
use std::os::raw::c_char;
use std::sync::mpsc;
use std::sync::{Arc, Mutex};
use tokio::sync::mpsc as tokio_mpsc;
#[derive(Debug, Deserialize)]
struct SendPayloadJson {
node: String,
#[serde(default)]
arg: Value,
}
#[derive(Debug, Deserialize)]
struct NodeExecResultJsonWire {
update: Option<Value>,
#[serde(default)]
sends: Vec<SendPayloadJson>,
}
#[derive(Debug, Deserialize)]
struct CallbackEnvelopeIn {
ok: bool,
#[serde(default)]
payload: Option<NodeExecResultJsonWire>,
#[serde(default)]
suspend: Option<WaitRequest>,
#[serde(default)]
error: Option<String>,
}
type CNodeCallback = unsafe extern "C" fn(
user_data: libc::c_ulong,
@@ -42,9 +14,6 @@ type CNodeCallback = unsafe extern "C" fn(
state_json: *mut c_char,
) -> *mut c_char;
#[derive(Clone, Copy)]
struct CUserData(libc::c_ulong);
fn cstr_to_str<'a>(ptr: *const c_char) -> Result<&'a str, String> {
if ptr.is_null() {
return Err("Received null pointer".to_string());
@@ -63,217 +32,6 @@ fn into_c_ptr(s: String) -> *mut c_char {
}
}
fn parse_c_callback_result(
raw: String,
node_name: &str,
) -> Result<NodeOutcome<Value, Value>, String> {
let parsed: CallbackEnvelopeIn = serde_json::from_str(&raw)
.map_err(|e| format!("decode callback envelope for `{node_name}` failed: {e}"))?;
if !parsed.ok {
return Err(parsed
.error
.unwrap_or_else(|| format!("callback reported error for `{node_name}`")));
}
if let Some(wait) = parsed.suspend {
return Ok(NodeOutcome::Suspended { wait });
}
let payload = parsed
.payload
.ok_or_else(|| format!("callback payload missing for `{node_name}`"))?;
let sends = payload
.sends
.into_iter()
.map(|s| SendPayload {
node: s.node,
arg: s.arg,
})
.collect();
Ok(NodeOutcome::Completed(NodeExecResult {
update: payload.update,
sends,
}))
}
enum SchedulerEventJson {
Node(Result<NodeExecutionJson, String>),
Resume {
node: String,
arg: Value,
event: WaitEvent,
},
WaitError(String),
}
struct NodeExecutionJson {
node: String,
arg: Value,
outcome: NodeOutcome<Value, Value>,
}
fn spawn_json_node_task(
node: String,
arg: Value,
state_snapshot: Value,
tx: tokio_mpsc::UnboundedSender<SchedulerEventJson>,
user_data_bits: libc::c_ulong,
callback: CNodeCallback,
) -> Result<(), String> {
node_pool_execute(move || {
let node_for_result = node.clone();
let arg_for_result = arg.clone();
let result = (|| -> Result<NodeExecutionJson, String> {
let node_c =
CString::new(node.clone()).map_err(|e| format!("invalid node name: {e}"))?;
let arg_json = serde_json::to_string(&arg)
.map_err(|e| format!("serialize arg for `{node}` failed: {e}"))?;
let state_json = serde_json::to_string(&state_snapshot)
.map_err(|e| format!("serialize state for `{node}` failed: {e}"))?;
let arg_c =
CString::new(arg_json).map_err(|e| format!("invalid arg JSON bytes: {e}"))?;
let state_c =
CString::new(state_json).map_err(|e| format!("invalid state JSON bytes: {e}"))?;
let out_ptr = unsafe {
callback(
user_data_bits,
node_c.as_ptr() as *mut c_char,
arg_c.as_ptr() as *mut c_char,
state_c.as_ptr() as *mut c_char,
)
};
if out_ptr.is_null() {
return Err(format!("callback returned null for `{node}`"));
}
let out_raw = unsafe { CStr::from_ptr(out_ptr) }
.to_string_lossy()
.into_owned();
unsafe {
libc::free(out_ptr.cast());
}
let payload = parse_c_callback_result(out_raw, &node)?;
Ok(NodeExecutionJson {
node: node_for_result,
arg: arg_for_result,
outcome: payload,
})
})();
let _ = tx.send(SchedulerEventJson::Node(result));
})
}
async fn run_graph_scheduler_json(
entry_point: String,
finish_point: String,
initial_state: Value,
initial_input: Value,
engine: Engine,
user_data: CUserData,
callback: CNodeCallback,
) -> Result<Value, String> {
let (tx, mut rx) = tokio_mpsc::unbounded_channel::<SchedulerEventJson>();
let state = Arc::new(Mutex::new(initial_state));
let user_data_bits = user_data.0;
let tx_for_spawn = tx.clone();
let state_for_spawn = Arc::clone(&state);
let state_for_merge = Arc::clone(&state);
let mut active: usize = 1;
let mut waiting: usize = 0;
spawn_json_node_task(
entry_point,
initial_input,
state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone(),
tx_for_spawn.clone(),
user_data_bits,
callback,
)?;
while active > 0 || waiting > 0 {
let evt = rx
.recv()
.await
.ok_or_else(|| "scheduler event channel closed".to_string())?;
match evt {
SchedulerEventJson::Node(result) => {
active = active.saturating_sub(1);
let exec = result?;
match exec.outcome {
NodeOutcome::Completed(node_result) => {
let mut guard = state_for_merge.lock().expect("state mutex poisoned");
merge_json_update(&mut guard, node_result.update);
drop(guard);
if exec.node == finish_point {
break;
}
for send in node_result.sends {
active += 1;
let snapshot = state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone();
spawn_json_node_task(
send.node,
send.arg,
snapshot,
tx_for_spawn.clone(),
user_data_bits,
callback,
)?;
}
}
NodeOutcome::Suspended { wait } => {
waiting += 1;
let tx_wait = tx_for_spawn.clone();
let node = exec.node;
let arg = exec.arg;
let engine_for_wait = engine.clone();
tokio::spawn(async move {
match engine_for_wait.wait_request_async(&wait).await {
Ok(event) => {
let _ = tx_wait.send(SchedulerEventJson::Resume {
node,
arg,
event,
});
}
Err(e) => {
let _ = tx_wait.send(SchedulerEventJson::WaitError(e));
}
}
});
}
}
}
SchedulerEventJson::Resume { node, arg, event } => {
waiting = waiting.saturating_sub(1);
active += 1;
let snapshot = state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone();
spawn_json_node_task(
node,
wrap_resume_arg(arg, event),
snapshot,
tx_for_spawn.clone(),
user_data_bits,
callback,
)?;
}
SchedulerEventJson::WaitError(e) => return Err(e),
}
}
let final_state = state.lock().expect("state mutex poisoned").clone();
Ok(final_state)
}
fn wrap_resume_arg(arg: Value, event: WaitEvent) -> Value {
serde_json::json!({
"__lg_resume_arg__": arg,
"__lg_resume_event__": event,
})
}
#[no_mangle]
pub extern "C" fn rc_engine_new() -> *mut Engine {
Box::into_raw(Box::new(Engine::new()))
@@ -528,22 +286,56 @@ pub unsafe extern "C" fn rc_run_graph_json(
}
};
if let Err(e) = (*ptr).start_stream(stream_mode.as_deref()) {
return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}"));
if let Some(mode) = stream_mode.as_deref() {
if let Err(e) = (*ptr).start_stream(Some(mode)) {
return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}"));
}
}
let (tx, rx) = mpsc::channel::<Result<Value, String>>();
let user_data = CUserData(user_data);
let user_data_bits = user_data;
let run_engine = (*ptr).clone();
let submit = run_loop_spawn(async move {
let out = run_graph_scheduler_json(
let callback_wrapper = move |node: String,
arg: Value,
state_snapshot: Value|
-> Result<NodeOutcome<Value, Value>, String> {
let node_c =
CString::new(node.clone()).map_err(|e| format!("invalid node name: {e}"))?;
let arg_json = serde_json::to_string(&arg)
.map_err(|e| format!("serialize arg for `{node}` failed: {e}"))?;
let state_json = serde_json::to_string(&state_snapshot)
.map_err(|e| format!("serialize state for `{node}` failed: {e}"))?;
let arg_c =
CString::new(arg_json).map_err(|e| format!("invalid arg JSON bytes: {e}"))?;
let state_c =
CString::new(state_json).map_err(|e| format!("invalid state JSON bytes: {e}"))?;
let out_ptr = unsafe {
callback(
user_data_bits,
node_c.as_ptr() as *mut c_char,
arg_c.as_ptr() as *mut c_char,
state_c.as_ptr() as *mut c_char,
)
};
if out_ptr.is_null() {
return Err(format!("callback returned null for `{node}`"));
}
let out_raw = unsafe { CStr::from_ptr(out_ptr) }
.to_string_lossy()
.into_owned();
unsafe {
libc::free(out_ptr.cast());
}
parse_callback_envelope_json(&out_raw, &node)
};
let out = run_graph_json_with_callback(
entry_point,
finish_point,
initial_state,
initial_input,
run_engine.clone(),
user_data,
callback,
callback_wrapper,
)
.await;
run_engine.close_stream();
+66 -186
View File
@@ -1,6 +1,6 @@
#[cfg(feature = "python-bindings")]
use crate::engine::{
node_pool_execute, run_loop_block_on, run_loop_spawn, AnyOfCondition, Engine, NodeExecResult,
run_graph_with_callback, run_loop_block_on, AnyOfCondition, Engine, NodeExecResult,
NodeOutcome, SendPayload, WaitCondition, WaitEvent, WaitRequest,
};
#[cfg(feature = "python-bindings")]
@@ -14,11 +14,7 @@ use pyo3::types::{PyDict, PyList, PyTuple};
#[cfg(feature = "python-bindings")]
use serde_json::Value;
#[cfg(feature = "python-bindings")]
use std::sync::mpsc;
#[cfg(feature = "python-bindings")]
use std::sync::Arc;
#[cfg(feature = "python-bindings")]
use tokio::sync::mpsc as tokio_mpsc;
#[cfg(feature = "python-bindings")]
#[pyclass]
@@ -50,11 +46,7 @@ impl PyRustEngine {
fn publish_obj(&self, py: Python<'_>, channel: &str, value: Py<PyAny>) -> PyResult<()> {
let value_json = py_obj_to_json_string(py, &value.bind(py))?;
let parsed: Value = serde_json::from_str(&value_json)
.map_err(|e| PyValueError::new_err(format!("Invalid Python JSON value: {e}")))?;
self.inner
.publish_json(channel, parsed)
.map_err(PyValueError::new_err)
self.publish_json(channel, &value_json)
}
fn wait_any_of_json(&self, any_of_json: &str) -> PyResult<String> {
@@ -89,12 +81,7 @@ impl PyRustEngine {
fn wait_any_of_obj(&self, py: Python<'_>, any_of_payload: Py<PyAny>) -> PyResult<Py<PyAny>> {
let payload_json = py_obj_to_json_string(py, &any_of_payload.bind(py))?;
let any_of: AnyOfCondition = serde_json::from_str(&payload_json)
.map_err(|e| PyValueError::new_err(format!("Invalid any_of payload: {e}")))?;
let event = run_loop_block_on(self.inner.wait_for_any_of_async(&any_of))
.map_err(PyValueError::new_err)?;
let event_json = serde_json::to_string(&event)
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))?;
let event_json = self.wait_any_of_json(&payload_json)?;
json_string_to_py_obj(py, &event_json)
}
@@ -107,6 +94,7 @@ impl PyRustEngine {
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))
}
#[pyo3(signature = (stream_mode=None))]
fn start_stream(&self, stream_mode: Option<&str>) -> PyResult<()> {
self.inner
.start_stream(stream_mode)
@@ -114,7 +102,8 @@ impl PyRustEngine {
}
fn receive_stream_obj(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
match run_loop_block_on(self.inner.receive_stream_async()) {
let event = py.allow_threads(|| run_loop_block_on(self.inner.receive_stream_async()));
match event {
Some(value) => {
let event_json = serde_json::to_string(&value)
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))?;
@@ -136,6 +125,7 @@ impl PyRustEngine {
self.inner.close_stream();
}
#[pyo3(signature = (entry_point, finish_point, initial_state, callback, stream_mode=None))]
fn run_graph_py(
&self,
py: Python<'_>,
@@ -145,184 +135,74 @@ impl PyRustEngine {
callback: Py<PyAny>,
stream_mode: Option<&str>,
) -> PyResult<Py<PyAny>> {
let state = Arc::new(initial_state);
if let Some(mode) = stream_mode {
self.inner
.start_stream(Some(mode))
.map_err(PyValueError::new_err)?;
}
let engine = self.inner.clone();
let callback = Arc::new(callback);
let entry_point = entry_point.to_string();
let finish_point = finish_point.to_string();
self.inner
.start_stream(stream_mode)
.map_err(PyValueError::new_err)?;
let (done_tx, done_rx) = mpsc::channel::<Result<(), String>>();
let state_for_run = Arc::clone(&state);
let engine = self.inner.clone();
run_loop_spawn(async move {
let run_result = run_graph_scheduler(
entry_point,
finish_point,
callback,
state_for_run,
engine.clone(),
)
.await;
engine.close_stream();
let _ = done_tx.send(run_result);
})
.map_err(PyValueError::new_err)?;
let run_result = py
.allow_threads(move || done_rx.recv())
.map_err(|e| PyValueError::new_err(format!("run-loop recv failed: {e}")))?;
run_result.map_err(PyValueError::new_err)?;
Ok((*state).clone_ref(py))
}
}
let initial_state = Arc::new(initial_state);
let initial_input = Arc::clone(&initial_state);
#[cfg(feature = "python-bindings")]
enum SchedulerEventPy {
Node(Result<NodeExecutionPy, String>),
Resume {
node: String,
arg: Py<PyAny>,
event: WaitEvent,
},
WaitError(String),
}
struct NodeExecutionPy {
node: String,
arg: Py<PyAny>,
outcome: NodeOutcome<Py<PyAny>, Py<PyAny>>,
}
#[cfg(feature = "python-bindings")]
async fn run_graph_scheduler(
entry_point: String,
finish_point: String,
callback: Arc<Py<PyAny>>,
state: Arc<Py<PyAny>>,
engine: Engine,
) -> Result<(), String> {
let (tx, mut rx) = tokio_mpsc::unbounded_channel::<SchedulerEventPy>();
let initial_arg = Python::with_gil(|py| (*state).clone_ref(py));
let callback_for_spawn = Arc::clone(&callback);
let state_for_spawn = Arc::clone(&state);
let tx_for_spawn = tx.clone();
let state_for_merge = Arc::clone(&state);
let mut active: usize = 1;
let mut waiting: usize = 0;
spawn_node_task(
entry_point,
initial_arg,
tx_for_spawn.clone(),
Arc::clone(&callback_for_spawn),
Arc::clone(&state_for_spawn),
)?;
while active > 0 || waiting > 0 {
let event = rx
.recv()
.await
.ok_or_else(|| "scheduler event channel closed".to_string())?;
match event {
SchedulerEventPy::Node(result) => {
active = active.saturating_sub(1);
let exec = result?;
match exec.outcome {
NodeOutcome::Completed(node_result) => {
let run_result =
py.allow_threads(move || {
run_loop_block_on(run_graph_with_callback(
entry_point,
finish_point,
initial_state,
initial_input,
engine.clone(),
{
let callback = Arc::clone(&callback);
move |node: String,
arg: Arc<Py<PyAny>>,
state_snapshot: Arc<Py<PyAny>>|
-> Result<NodeOutcome<Arc<Py<PyAny>>, Arc<Py<PyAny>>>, String> {
Python::with_gil(|py| -> Result<NodeOutcome<Arc<Py<PyAny>>, Arc<Py<PyAny>>>, String> {
let callback_bound = callback.as_ref().bind(py);
let payload_obj = callback_bound
.call1((
node.as_str(),
arg.as_ref().clone_ref(py),
state_snapshot.as_ref().clone_ref(py),
))
.map_err(|e| format!("callback failed for node `{node}`: {e}"))?;
parse_node_outcome_arc(py, &payload_obj).map_err(|e| {
format!("invalid callback payload for `{node}`: {e}")
})
})
}
},
|state: &mut Arc<Py<PyAny>>, update: Option<Arc<Py<PyAny>>>| -> Result<(), String> {
Python::with_gil(|py| -> Result<(), String> {
if let Some(update) = node_result.update {
apply_update_to_state(py, state_for_merge.as_ref(), &update)
.map_err(|e| {
format!("state merge failed for `{}`: {e}", exec.node)
})?;
if let Some(update) = update {
apply_update_to_state(py, state.as_ref(), update.as_ref())
.map_err(|e| format!("state merge failed: {e}"))?;
}
Ok(())
})?;
if exec.node == finish_point {
break;
}
for send in node_result.sends {
active += 1;
spawn_node_task(
send.node,
send.arg,
tx_for_spawn.clone(),
Arc::clone(&callback_for_spawn),
Arc::clone(&state_for_spawn),
)?;
}
}
NodeOutcome::Suspended { wait } => {
waiting += 1;
let tx_wait = tx_for_spawn.clone();
let node = exec.node;
let arg = exec.arg;
let engine_for_wait = engine.clone();
tokio::spawn(async move {
let outcome = engine_for_wait.wait_request_async(&wait).await;
match outcome {
Ok(event) => {
let _ =
tx_wait.send(SchedulerEventPy::Resume { node, arg, event });
}
Err(e) => {
let _ = tx_wait.send(SchedulerEventPy::WaitError(e));
}
}
});
}
}
}
SchedulerEventPy::Resume { node, arg, event } => {
waiting = waiting.saturating_sub(1);
active += 1;
let resume_arg = wrap_resume_arg(&arg, &event)?;
spawn_node_task(
node,
resume_arg,
tx_for_spawn.clone(),
Arc::clone(&callback_for_spawn),
Arc::clone(&state_for_spawn),
)?;
}
SchedulerEventPy::WaitError(e) => return Err(e),
}
})
},
|arg: Arc<Py<PyAny>>, event: WaitEvent| -> Result<Arc<Py<PyAny>>, String> {
wrap_resume_arg(arg.as_ref(), &event).map(Arc::new)
},
))
});
self.inner.close_stream();
let out = run_result.map_err(PyValueError::new_err)?;
Ok(out.as_ref().clone_ref(py))
}
Ok(())
}
#[cfg(feature = "python-bindings")]
fn spawn_node_task(
node: String,
arg: Py<PyAny>,
tx: tokio_mpsc::UnboundedSender<SchedulerEventPy>,
callback: Arc<Py<PyAny>>,
state_for_task: Arc<Py<PyAny>>,
) -> Result<(), String> {
node_pool_execute(move || {
let node_for_result = node.clone();
let arg_for_result = Python::with_gil(|py| arg.clone_ref(py));
let outcome = Python::with_gil(|py| -> Result<NodeExecutionPy, String> {
let callback_bound = callback.as_ref().bind(py);
let payload_obj = callback_bound
.call1((node.as_str(), arg, (*state_for_task).clone_ref(py)))
.map_err(|e| format!("callback failed for node `{node}`: {e}"))?;
let payload = parse_node_outcome(py, &payload_obj)
.map_err(|e| format!("invalid callback payload for `{node}`: {e}"))?;
Ok(NodeExecutionPy {
node: node_for_result,
arg: arg_for_result,
outcome: payload,
})
});
let _ = tx.send(SchedulerEventPy::Node(outcome));
})
}
#[cfg(feature = "python-bindings")]
fn parse_node_outcome(
fn parse_node_outcome_arc(
py: Python<'_>,
payload_obj: &Bound<'_, PyAny>,
) -> Result<NodeOutcome<Py<PyAny>, Py<PyAny>>, String> {
) -> Result<NodeOutcome<Arc<Py<PyAny>>, Arc<Py<PyAny>>>, String> {
let payload_dict = payload_obj
.downcast::<PyDict>()
.map_err(|_| "payload must be a dict".to_string())?;
@@ -341,7 +221,7 @@ fn parse_node_outcome(
.get_item("update")
.map_err(|e| format!("failed to read update: {e}"))?;
let update = match update_item {
Some(v) if !v.is_none() => Some(v.unbind()),
Some(v) if !v.is_none() => Some(Arc::new(v.unbind())),
_ => None,
};
@@ -365,9 +245,9 @@ fn parse_node_outcome(
let node = node_obj
.extract::<String>()
.map_err(|e| format!("send.node must be string: {e}"))?;
let arg: Py<PyAny> = match send_dict.get_item("arg") {
Ok(Some(v)) => v.unbind(),
Ok(None) => Python::with_gil(|py| py.None()),
let arg: Arc<Py<PyAny>> = match send_dict.get_item("arg") {
Ok(Some(v)) => Arc::new(v.unbind()),
Ok(None) => Arc::new(Python::with_gil(|py| py.None())),
Err(e) => return Err(format!("failed to read send.arg: {e}")),
};
sends.push(SendPayload { node, arg });
@@ -211,11 +211,9 @@ class GraphRunHandler(Generic[StateT]):
await self._run.publish(channel, value)
async def receive_stream(self) -> Any | None:
loop = asyncio.get_running_loop()
return await loop.run_in_executor(
_advanced_graph_executor(),
self._run.receive_stream_sync,
)
# Use a separate thread pool from graph execution to avoid deadlock
# when LANGGRAPH_ADVANCED_GRAPH_PY_THREADS is configured to 1.
return await asyncio.to_thread(self._run.receive_stream_sync)
def close_stream(self) -> None:
self._run.close_stream_sync()
@@ -244,6 +242,9 @@ class _GraphEngineRun:
self._rust_engine = PyRustEngine()
for name in async_channel_specs:
self._rust_engine.add_async_channel(name)
self._stream_ready = threading.Event()
if self._stream_mode is None:
self._stream_ready.set()
self._tasks: set[asyncio.Task[list[Send]]] = set()
self._finished = False
self._state: Any = None
@@ -252,16 +253,24 @@ class _GraphEngineRun:
async def run(self, initial_state: StateT) -> StateT:
finish_point = self._finish_point or ""
if self._stream_mode is not None:
try:
self._rust_engine.start_stream(self._stream_mode)
finally:
self._stream_ready.set()
loop = asyncio.get_running_loop()
result_obj = await loop.run_in_executor(
_advanced_graph_executor(),
self._rust_engine.run_graph_py,
self._entry_point,
finish_point,
initial_state,
self._execute_node_for_rust,
self._stream_mode,
)
try:
result_obj = await loop.run_in_executor(
_advanced_graph_executor(),
self._rust_engine.run_graph_py,
self._entry_point,
finish_point,
initial_state,
self._execute_node_for_rust,
None,
)
finally:
self._stream_ready.set()
self._state = result_obj
return cast(StateT, self._state)
@@ -328,6 +337,7 @@ class _GraphEngineRun:
self._rust_engine.send_custom_stream_event_obj(value)
def receive_stream_sync(self) -> Any | None:
self._stream_ready.wait()
return self._rust_engine.receive_stream_obj()
def close_stream_sync(self) -> None:
+1 -1
View File
@@ -4,5 +4,5 @@ requires-python = ">=3.10"
[[package]]
name = "saf-python-sdk"
version = "0.1.1"
version = "0.1.2"
source = { editable = "." }