diff --git a/.gitignore b/.gitignore index 145295a29..ca5b0266b 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/langgraph-go/advancedgraph/graph.go b/langgraph-go/advancedgraph/graph.go index 567a1cd41..bf4eabf4d 100644 --- a/langgraph-go/advancedgraph/graph.go +++ b/langgraph-go/advancedgraph/graph.go @@ -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) { diff --git a/rust-core/src/engine.rs b/rust-core/src/engine.rs index 724ef657e..789086fc0 100644 --- a/rust-core/src/engine.rs +++ b/rust-core/src/engine.rs @@ -491,3 +491,253 @@ impl Engine { })) } } + +#[derive(Debug, Deserialize)] +struct CallbackSendPayloadJson { + node: String, + #[serde(default)] + arg: Value, +} + +#[derive(Debug, Deserialize)] +struct CallbackNodeExecResultJsonWire { + update: Option, + #[serde(default)] + sends: Vec, +} + +#[derive(Debug, Deserialize)] +struct CallbackEnvelopeIn { + ok: bool, + #[serde(default)] + payload: Option, + #[serde(default)] + suspend: Option, + #[serde(default)] + error: Option, +} + +enum SchedulerEvent { + Node(Result, String>), + Resume { + node: String, + arg: A, + event: WaitEvent, + }, + WaitError(String), +} + +struct NodeExecution { + node: String, + arg: A, + outcome: NodeOutcome, +} + +fn spawn_node_task( + node: String, + arg: A, + state_snapshot: State, + tx: tokio_mpsc::UnboundedSender>, + callback: Arc, +) -> Result<(), String> +where + State: Send + 'static, + U: Send + 'static, + A: Clone + Send + 'static, + F: Fn(String, A, State) -> Result, 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( + entry_point: String, + finish_point: String, + initial_state: State, + initial_input: A, + engine: Engine, + callback: FCallback, + merge_update: FMerge, + wrap_resume_arg: FWrap, +) -> Result +where + State: Clone + Send + 'static, + U: Send + 'static, + A: Clone + Send + 'static, + FCallback: Fn(String, A, State) -> Result, String> + Send + Sync + 'static, + FMerge: Fn(&mut State, Option) -> Result<(), String> + Send + Sync + 'static, + FWrap: Fn(A, WaitEvent) -> Result + 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::>(); + 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( + entry_point: String, + finish_point: String, + initial_state: Value, + initial_input: Value, + engine: Engine, + callback: F, +) -> Result +where + F: Fn(String, Value, Value) -> Result, String> + + Send + + Sync + + 'static, +{ + run_graph_with_callback( + entry_point, + finish_point, + initial_state, + initial_input, + engine, + callback, + |state: &mut Value, update: Option| { + 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, 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, + })) +} diff --git a/rust-core/src/lib_c.rs b/rust-core/src/lib_c.rs index b68bd08d7..82a88421f 100644 --- a/rust-core/src/lib_c.rs +++ b/rust-core/src/lib_c.rs @@ -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, - #[serde(default)] - sends: Vec, -} - -#[derive(Debug, Deserialize)] -struct CallbackEnvelopeIn { - ok: bool, - #[serde(default)] - payload: Option, - #[serde(default)] - suspend: Option, - #[serde(default)] - error: Option, -} 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, 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), - Resume { - node: String, - arg: Value, - event: WaitEvent, - }, - WaitError(String), -} - -struct NodeExecutionJson { - node: String, - arg: Value, - outcome: NodeOutcome, -} - -fn spawn_json_node_task( - node: String, - arg: Value, - state_snapshot: Value, - tx: tokio_mpsc::UnboundedSender, - 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 { - 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 { - let (tx, mut rx) = tokio_mpsc::unbounded_channel::(); - 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::>(); - 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, 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(); diff --git a/rust-core/src/lib_py.rs b/rust-core/src/lib_py.rs index d93d23a67..6f4e034af 100644 --- a/rust-core/src/lib_py.rs +++ b/rust-core/src/lib_py.rs @@ -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) -> 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 { @@ -89,12 +81,7 @@ impl PyRustEngine { fn wait_any_of_obj(&self, py: Python<'_>, any_of_payload: Py) -> PyResult> { 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> { - 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, stream_mode: Option<&str>, ) -> PyResult> { - 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::>(); - 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), - Resume { - node: String, - arg: Py, - event: WaitEvent, - }, - WaitError(String), -} - -struct NodeExecutionPy { - node: String, - arg: Py, - outcome: NodeOutcome, Py>, -} - -#[cfg(feature = "python-bindings")] -async fn run_graph_scheduler( - entry_point: String, - finish_point: String, - callback: Arc>, - state: Arc>, - engine: Engine, -) -> Result<(), String> { - let (tx, mut rx) = tokio_mpsc::unbounded_channel::(); - 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>, + state_snapshot: Arc>| + -> Result>, Arc>>, String> { + Python::with_gil(|py| -> Result>, Arc>>, 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>, update: Option>>| -> 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>, event: WaitEvent| -> Result>, 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, - tx: tokio_mpsc::UnboundedSender, - callback: Arc>, - state_for_task: Arc>, -) -> 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 { - 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, Py>, String> { +) -> Result>, Arc>>, String> { let payload_dict = payload_obj .downcast::() .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::() .map_err(|e| format!("send.node must be string: {e}"))?; - let arg: Py = match send_dict.get_item("arg") { - Ok(Some(v)) => v.unbind(), - Ok(None) => Python::with_gil(|py| py.None()), + let arg: Arc> = 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 }); diff --git a/saf-python-sdk/python/saf_python_sdk/advanced_graph/state.py b/saf-python-sdk/python/saf_python_sdk/advanced_graph/state.py index b89985081..d3cb448fe 100644 --- a/saf-python-sdk/python/saf_python_sdk/advanced_graph/state.py +++ b/saf-python-sdk/python/saf_python_sdk/advanced_graph/state.py @@ -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: diff --git a/saf-python-sdk/python/saf_python_sdk/langgraph_rust_core.cpython-313-darwin.so b/saf-python-sdk/python/saf_python_sdk/langgraph_rust_core.cpython-313-darwin.so deleted file mode 100755 index dcdb42fc7..000000000 Binary files a/saf-python-sdk/python/saf_python_sdk/langgraph_rust_core.cpython-313-darwin.so and /dev/null differ diff --git a/saf-python-sdk/uv.lock b/saf-python-sdk/uv.lock index ea0ad8d69..41e30f5f1 100644 --- a/saf-python-sdk/uv.lock +++ b/saf-python-sdk/uv.lock @@ -4,5 +4,5 @@ requires-python = ">=3.10" [[package]] name = "saf-python-sdk" -version = "0.1.1" +version = "0.1.2" source = { editable = "." }