diff --git a/langgraph-go/advancedgraph/graph.go b/langgraph-go/advancedgraph/graph.go index 620abece2..567a1cd41 100644 --- a/langgraph-go/advancedgraph/graph.go +++ b/langgraph-go/advancedgraph/graph.go @@ -111,6 +111,10 @@ func (c *Context) PublishToChannel(channel string, value any) error { return c.engine.Publish(channel, value) } +func (c *Context) SendCustomStreamEvent(value any) error { + return c.engine.SendCustomStreamEvent(value) +} + type Handler[StateT any] struct { engine *RustEngine done chan resultOrErr[StateT] @@ -130,7 +134,29 @@ func (h *Handler[StateT]) WaitForResult() (StateT, error) { return res.state, res.err } -func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT) (*Handler[StateT], error) { +func (h *Handler[StateT]) ReceiveStream() (any, error) { + event, hasEvent, err := h.engine.ReceiveStream() + if err != nil { + return nil, err + } + if !hasEvent { + return nil, nil + } + return event, nil +} + +func (h *Handler[StateT]) CloseStream() error { + return h.engine.CloseStream() +} + +func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT, streamMode ...string) (*Handler[StateT], error) { + resolvedStreamMode := "" + if len(streamMode) > 1 { + return nil, fmt.Errorf("start accepts at most one stream mode") + } + if len(streamMode) == 1 { + resolvedStreamMode = streamMode[0] + } engine := NewRustEngine() for _, ch := range g.asyncChannels { if err := engine.AddAsyncChannel(ch); err != nil { @@ -147,6 +173,7 @@ func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT) (*H rawState, err := engine.RunGraph( g.entryPoint, g.finishPoint, + resolvedStreamMode, initialState, initialInput, func(node string, nodeInput any, fallbackState map[string]any) (Command, error) { diff --git a/langgraph-go/advancedgraph/rust_engine.go b/langgraph-go/advancedgraph/rust_engine.go index 9d31d3396..0d078adde 100644 --- a/langgraph-go/advancedgraph/rust_engine.go +++ b/langgraph-go/advancedgraph/rust_engine.go @@ -128,6 +128,59 @@ func (e *RustEngine) AddAsyncChannel(channel string) error { return parseRustStatus(resp) } +func (e *RustEngine) StartStream(streamMode string) error { + var cmode *C.char + if streamMode != "" { + cmode = C.CString(streamMode) + defer C.free(unsafe.Pointer(cmode)) + } + resp := C.rc_start_stream(e.ptr, cmode) + return parseRustStatus(resp) +} + +func (e *RustEngine) ReceiveStream() (any, bool, error) { + resp := C.rc_receive_stream_json(e.ptr) + defer C.rc_string_free(resp) + + raw := C.GoString(resp) + var status struct { + OK bool `json:"ok"` + Error string `json:"error"` + HasEvent bool `json:"has_event"` + Event json.RawMessage `json:"event"` + } + if err := json.Unmarshal([]byte(raw), &status); err != nil { + return nil, false, fmt.Errorf("decode rust stream response: %w", err) + } + if !status.OK { + return nil, false, fmt.Errorf("rust stream failed: %s", status.Error) + } + if !status.HasEvent { + return nil, false, nil + } + var event any + if err := json.Unmarshal(status.Event, &event); err != nil { + return nil, false, fmt.Errorf("decode stream event: %w", err) + } + return coerceJSONValue(event), true, nil +} + +func (e *RustEngine) SendCustomStreamEvent(value any) error { + payload, err := json.Marshal(value) + if err != nil { + return fmt.Errorf("marshal stream event: %w", err) + } + cval := C.CString(string(payload)) + defer C.free(unsafe.Pointer(cval)) + resp := C.rc_send_custom_stream_event(e.ptr, cval) + return parseRustStatus(resp) +} + +func (e *RustEngine) CloseStream() error { + resp := C.rc_close_stream(e.ptr) + return parseRustStatus(resp) +} + func (e *RustEngine) Publish(channel string, value any) error { payload, err := json.Marshal(value) if err != nil { @@ -173,6 +226,7 @@ func (e *RustEngine) WaitAnyOf(cond AnyOfCondition) (WaitEvent, error) { func (e *RustEngine) RunGraph( entryPoint string, finishPoint string, + streamMode string, initialState any, initialInput any, exec func(node string, nodeInput any, state map[string]any) (Command, error), @@ -189,10 +243,17 @@ func (e *RustEngine) RunGraph( cfinish := C.CString(finishPoint) cinitial := C.CString(string(initialJSON)) cinitialInput := C.CString(string(initialInputJSON)) + var cstreamMode *C.char + if streamMode != "" { + cstreamMode = C.CString(streamMode) + } defer C.free(unsafe.Pointer(centry)) defer C.free(unsafe.Pointer(cfinish)) defer C.free(unsafe.Pointer(cinitial)) defer C.free(unsafe.Pointer(cinitialInput)) + if cstreamMode != nil { + defer C.free(unsafe.Pointer(cstreamMode)) + } callbackID := registerRunGraphCallbackCtx(&runGraphCallbackCtx{exec: exec}) defer unregisterRunGraphCallbackCtx(callbackID) @@ -203,6 +264,7 @@ func (e *RustEngine) RunGraph( cfinish, cinitial, cinitialInput, + cstreamMode, C.ulong(callbackID), (C.rc_node_callback_t)(C.goNodeCallback), ) diff --git a/langgraph-go/tests/test_streaming_test.go b/langgraph-go/tests/test_streaming_test.go new file mode 100644 index 000000000..6095d8203 --- /dev/null +++ b/langgraph-go/tests/test_streaming_test.go @@ -0,0 +1,102 @@ +package tests + +import ( + "strings" + "testing" + "time" + + ag "github.com/langchain-ai/langgraph/langgraph-go/advancedgraph" +) + +type streamState struct { + Done bool `json:"done"` +} + +type streamWorkflow struct{} + +func (w *streamWorkflow) startNode(ctx *ag.Context, _ any, state streamState) (ag.Command, error) { + if err := ctx.SendCustomStreamEvent(map[string]any{"step": "start", "value": 1}); err != nil { + return ag.Command{}, err + } + time.Sleep(80 * time.Millisecond) + if err := ctx.SendCustomStreamEvent(map[string]any{"step": "start", "value": 2}); err != nil { + return ag.Command{}, err + } + return ag.Command{ + Update: state, + Goto: []ag.Send{ + {Node: w.finishNode}, + }, + }, nil +} + +func (w *streamWorkflow) finishNode(ctx *ag.Context, _ any, state streamState) (ag.Command, error) { + state.Done = true + return ag.Command{Update: state}, nil +} + +func TestCustomStreamReceiveAndClose(t *testing.T) { + workflow := &streamWorkflow{} + graph := ag.NewAdvancedStateGraph[streamState]() + graph.AddEntryNode(workflow.startNode) + graph.AddFinishNode(workflow.finishNode) + + handler, err := graph.Compile().Start(nil, streamState{Done: false}, "custom") + if err != nil { + t.Fatalf("start failed: %v", err) + } + + event, err := handler.ReceiveStream() + if err != nil { + t.Fatalf("receive stream failed: %v", err) + } + if event == nil { + t.Fatalf("expected first stream event, got nil") + } + eventMap, ok := event.(map[string]any) + if !ok { + t.Fatalf("unexpected event type: %T", event) + } + if eventMap["step"] != "start" { + t.Fatalf("unexpected stream event payload: %#v", eventMap) + } + + if err := handler.CloseStream(); err != nil { + t.Fatalf("close stream failed: %v", err) + } + + closedEvent, err := handler.ReceiveStream() + if err != nil { + t.Fatalf("receive stream after close failed: %v", err) + } + if closedEvent != nil { + t.Fatalf("expected nil stream event after close, got %#v", closedEvent) + } + + result, err := handler.WaitForResult() + if err != nil { + t.Fatalf("result failed: %v", err) + } + if !result.Done { + t.Fatalf("expected final state done=true, got %#v", result) + } +} + +func TestOnlyCustomStreamModeSupported(t *testing.T) { + workflow := &streamWorkflow{} + graph := ag.NewAdvancedStateGraph[streamState]() + graph.AddEntryNode(workflow.startNode) + graph.AddFinishNode(workflow.finishNode) + + handler, err := graph.Compile().Start(nil, streamState{Done: false}, "values") + if err != nil { + t.Fatalf("start failed: %v", err) + } + _, runErr := handler.WaitForResult() + if runErr == nil { + t.Fatalf("expected run error for unsupported stream mode") + } + if !strings.Contains(runErr.Error(), "only `custom` is supported") { + t.Fatalf("unexpected error: %v", runErr) + } +} diff --git a/rust-core/include/langgraph_rust_core.h b/rust-core/include/langgraph_rust_core.h index c00880bde..041a46d18 100644 --- a/rust-core/include/langgraph_rust_core.h +++ b/rust-core/include/langgraph_rust_core.h @@ -19,12 +19,17 @@ void rc_engine_free(Engine* ptr); char* rc_add_async_channel(Engine* ptr, const char* channel); char* rc_publish_json(Engine* ptr, const char* channel, const char* value_json); char* rc_wait_any_of_json(Engine* ptr, const char* any_of_json); +char* rc_start_stream(Engine* ptr, const char* stream_mode); +char* rc_receive_stream_json(Engine* ptr); +char* rc_send_custom_stream_event(Engine* ptr, const char* value_json); +char* rc_close_stream(Engine* ptr); char* rc_run_graph_json( Engine* ptr, const char* entry_point, const char* finish_point, const char* initial_state_json, const char* initial_input_json, + const char* stream_mode, unsigned long user_data, rc_node_callback_t callback ); diff --git a/rust-core/src/engine.rs b/rust-core/src/engine.rs index c9f82180f..724ef657e 100644 --- a/rust-core/src/engine.rs +++ b/rust-core/src/engine.rs @@ -10,7 +10,7 @@ use std::sync::OnceLock; use std::thread; use std::time::{Duration, Instant}; use tokio::runtime::Runtime; -use tokio::sync::Notify; +use tokio::sync::{mpsc as tokio_mpsc, Mutex as AsyncMutex, Notify}; #[derive(Clone, Debug, Serialize, Deserialize)] #[serde(tag = "kind")] @@ -283,6 +283,13 @@ pub fn merge_json_update(state: &mut Value, update: Option) { pub struct Engine { channels: Arc>>>, channel_notify: Arc, + stream: Arc>>, +} + +#[derive(Clone)] +struct StreamChannel { + sender: tokio_mpsc::UnboundedSender, + receiver: Arc>>, } impl Engine { @@ -297,6 +304,60 @@ impl Engine { channels.entry(name.to_owned()).or_default(); } + pub fn start_stream(&self, stream_mode: Option<&str>) -> Result<(), String> { + let mode = stream_mode.and_then(|m| { + let trimmed = m.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed) + } + }); + let mut stream = self.stream.lock().expect("stream mutex poisoned"); + if let Some(mode) = mode { + if mode != "custom" { + return Err(format!( + "unsupported stream_mode `{mode}`, only `custom` is supported" + )); + } + let (sender, receiver) = tokio_mpsc::unbounded_channel(); + *stream = Some(StreamChannel { + sender, + receiver: Arc::new(AsyncMutex::new(receiver)), + }); + } else { + *stream = None; + } + Ok(()) + } + + pub fn close_stream(&self) { + let mut stream = self.stream.lock().expect("stream mutex poisoned"); + *stream = None; + } + + pub fn send_custom_stream_event(&self, value: serde_json::Value) { + let sender = { + let stream = self.stream.lock().expect("stream mutex poisoned"); + stream.as_ref().map(|s| s.sender.clone()) + }; + if let Some(tx) = sender { + let _ = tx.send(value); + } + } + + pub async fn receive_stream_async(&self) -> Option { + let receiver = { + let stream = self.stream.lock().expect("stream mutex poisoned"); + stream.as_ref().map(|s| Arc::clone(&s.receiver)) + }; + let Some(receiver) = receiver else { + return None; + }; + let mut guard = receiver.lock().await; + guard.recv().await + } + pub fn publish_json(&self, channel: &str, value: serde_json::Value) -> Result<(), String> { debug_log(&format!("Engine::publish_json(channel={channel})")); let mut channels = self.channels.lock().expect("channels mutex poisoned"); @@ -339,7 +400,10 @@ impl Engine { } } - pub async fn wait_for_any_of_async(&self, any_of: &AnyOfCondition) -> Result { + pub async fn wait_for_any_of_async( + &self, + any_of: &AnyOfCondition, + ) -> Result { debug_log(&format!( "Engine::wait_for_any_of_async(conditions={})", any_of.conditions.len() @@ -398,11 +462,7 @@ impl Engine { run_loop_block_on(self.wait_for_any_of_async(any_of)) } - fn try_take_channel_event( - &self, - channel: &str, - n: usize, - ) -> Result, String> { + fn try_take_channel_event(&self, channel: &str, n: usize) -> Result, String> { let mut channels = self.channels.lock().expect("channels mutex poisoned"); let queue = channels .get_mut(channel) diff --git a/rust-core/src/lib_c.rs b/rust-core/src/lib_c.rs index 1760ff75a..b68bd08d7 100644 --- a/rust-core/src/lib_c.rs +++ b/rust-core/src/lib_c.rs @@ -230,7 +230,11 @@ async fn run_graph_scheduler_json( tokio::spawn(async move { match engine_for_wait.wait_request_async(&wait).await { Ok(event) => { - let _ = tx_wait.send(SchedulerEventJson::Resume { node, arg, event }); + let _ = tx_wait.send(SchedulerEventJson::Resume { + node, + arg, + event, + }); } Err(e) => { let _ = tx_wait.send(SchedulerEventJson::WaitError(e)); @@ -382,6 +386,86 @@ pub unsafe extern "C" fn rc_wait_any_of_json( } } +#[no_mangle] +/// # Safety +/// `ptr` must be a valid engine pointer from `rc_engine_new`. +/// `stream_mode` must be null or a valid null-terminated UTF-8 string pointer. +pub unsafe extern "C" fn rc_start_stream( + ptr: *mut Engine, + stream_mode: *const c_char, +) -> *mut c_char { + if ptr.is_null() { + return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string()); + } + let mode = if stream_mode.is_null() { + None + } else { + match cstr_to_str(stream_mode) { + Ok(v) => Some(v), + Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + } + }; + match (*ptr).start_stream(mode) { + Ok(()) => into_c_ptr("{\"ok\":true}".to_string()), + Err(e) => into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + } +} + +#[no_mangle] +/// # Safety +/// `ptr` must be a valid engine pointer from `rc_engine_new`. +pub unsafe extern "C" fn rc_receive_stream_json(ptr: *mut Engine) -> *mut c_char { + if ptr.is_null() { + return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string()); + } + let event = run_loop_block_on((*ptr).receive_stream_async()); + match event { + Some(value) => match serde_json::to_string(&value) { + Ok(s) => into_c_ptr(format!("{{\"ok\":true,\"has_event\":true,\"event\":{s}}}")), + Err(e) => into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + }, + None => into_c_ptr("{\"ok\":true,\"has_event\":false}".to_string()), + } +} + +#[no_mangle] +/// # Safety +/// `ptr` must be a valid engine pointer from `rc_engine_new`. +/// `value_json` must be a valid null-terminated UTF-8 string pointer. +pub unsafe extern "C" fn rc_send_custom_stream_event( + ptr: *mut Engine, + value_json: *const c_char, +) -> *mut c_char { + if ptr.is_null() { + return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string()); + } + let value_json = match cstr_to_str(value_json) { + Ok(v) => v, + Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + }; + let value: Value = match serde_json::from_str(value_json) { + Ok(v) => v, + Err(e) => { + return into_c_ptr(format!( + "{{\"ok\":false,\"error\":\"invalid JSON value: {e}\"}}" + )) + } + }; + (*ptr).send_custom_stream_event(value); + into_c_ptr("{\"ok\":true}".to_string()) +} + +#[no_mangle] +/// # Safety +/// `ptr` must be a valid engine pointer from `rc_engine_new`. +pub unsafe extern "C" fn rc_close_stream(ptr: *mut Engine) -> *mut c_char { + if ptr.is_null() { + return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string()); + } + (*ptr).close_stream(); + into_c_ptr("{\"ok\":true}".to_string()) +} + #[no_mangle] /// # Safety /// `ptr` must be a valid engine pointer from `rc_engine_new`. @@ -393,6 +477,7 @@ pub unsafe extern "C" fn rc_run_graph_json( finish_point: *const c_char, initial_state_json: *const c_char, initial_input_json: *const c_char, + stream_mode: *const c_char, user_data: libc::c_ulong, callback: Option, ) -> *mut c_char { @@ -434,6 +519,18 @@ pub unsafe extern "C" fn rc_run_graph_json( )) } }; + let stream_mode = if stream_mode.is_null() { + None + } else { + match cstr_to_str(stream_mode) { + Ok(v) => Some(v.to_string()), + Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + } + }; + + if let Err(e) = (*ptr).start_stream(stream_mode.as_deref()) { + return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")); + } let (tx, rx) = mpsc::channel::>(); let user_data = CUserData(user_data); @@ -444,11 +541,12 @@ pub unsafe extern "C" fn rc_run_graph_json( finish_point, initial_state, initial_input, - run_engine, + run_engine.clone(), user_data, callback, ) .await; + run_engine.close_stream(); let _ = tx.send(out); }); if let Err(e) = submit { diff --git a/rust-core/src/lib_py.rs b/rust-core/src/lib_py.rs index d260b271a..d93d23a67 100644 --- a/rust-core/src/lib_py.rs +++ b/rust-core/src/lib_py.rs @@ -107,6 +107,35 @@ impl PyRustEngine { .map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}"))) } + fn start_stream(&self, stream_mode: Option<&str>) -> PyResult<()> { + self.inner + .start_stream(stream_mode) + .map_err(PyValueError::new_err) + } + + fn receive_stream_obj(&self, py: Python<'_>) -> PyResult> { + match run_loop_block_on(self.inner.receive_stream_async()) { + Some(value) => { + let event_json = serde_json::to_string(&value) + .map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))?; + json_string_to_py_obj(py, &event_json) + } + None => Ok(py.None()), + } + } + + fn send_custom_stream_event_obj(&self, py: Python<'_>, 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.send_custom_stream_event(parsed); + Ok(()) + } + + fn close_stream(&self) { + self.inner.close_stream(); + } + fn run_graph_py( &self, py: Python<'_>, @@ -114,18 +143,28 @@ impl PyRustEngine { finish_point: &str, initial_state: Py, callback: Py, + stream_mode: Option<&str>, ) -> PyResult> { let state = Arc::new(initial_state); 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) - .await; + 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)?; @@ -222,7 +261,8 @@ async fn run_graph_scheduler( let outcome = engine_for_wait.wait_request_async(&wait).await; match outcome { Ok(event) => { - let _ = tx_wait.send(SchedulerEventPy::Resume { node, arg, event }); + let _ = + tx_wait.send(SchedulerEventPy::Resume { node, arg, event }); } Err(e) => { let _ = tx_wait.send(SchedulerEventPy::WaitError(e)); @@ -261,21 +301,19 @@ fn spawn_node_task( 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 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)); }) } @@ -294,8 +332,8 @@ fn parse_node_outcome( if let Some(wait_obj) = suspended_item { let wait_json = py_obj_to_json_string(py, &wait_obj) .map_err(|e| format!("failed to encode suspend payload: {e}"))?; - let wait: WaitRequest = - serde_json::from_str(&wait_json).map_err(|e| format!("invalid suspend payload: {e}"))?; + let wait: WaitRequest = serde_json::from_str(&wait_json) + .map_err(|e| format!("invalid suspend payload: {e}"))?; return Ok(NodeOutcome::Suspended { wait }); } @@ -345,10 +383,10 @@ fn wrap_resume_arg(arg: &Py, event: &WaitEvent) -> Result, Stri wrapper .set_item("__lg_resume_arg__", arg.clone_ref(py)) .map_err(|e| format!("failed to set resume arg: {e}"))?; - let event_json = - serde_json::to_string(event).map_err(|e| format!("failed to encode wait event: {e}"))?; - let event_obj = - json_string_to_py_obj(py, &event_json).map_err(|e| format!("failed to parse event: {e}"))?; + let event_json = serde_json::to_string(event) + .map_err(|e| format!("failed to encode wait event: {e}"))?; + let event_obj = json_string_to_py_obj(py, &event_json) + .map_err(|e| format!("failed to parse event: {e}"))?; wrapper .set_item("__lg_resume_event__", event_obj.bind(py)) .map_err(|e| format!("failed to set resume event: {e}"))?; 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 26440de99..b89985081 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 @@ -162,12 +162,15 @@ class CompiledGraphEngine(Generic[StateT]): handler = await self.astart(initial_state) return await handler - async def astart(self, initial_state: StateT) -> GraphRunHandler[StateT]: + async def astart( + self, initial_state: StateT, *, stream_mode: str | None = None + ) -> GraphRunHandler[StateT]: run = _GraphEngineRun( nodes=self._nodes, async_channel_specs=self._async_channels, entry_point=self._entry_point, finish_point=self._finish_point, + stream_mode=stream_mode, ) task = asyncio.create_task(run.run(initial_state)) return GraphRunHandler(run=run, task=task) @@ -191,6 +194,9 @@ class Context: async def apublish_to_channel(self, channel: str, value: Any) -> None: await self._run.publish(channel, value) + def send_custom_stream_event(self, value: Any) -> None: + self._run.send_custom_stream_event(value) + class GraphRunHandler(Generic[StateT]): """Handle for an active in-memory run.""" @@ -204,6 +210,16 @@ class GraphRunHandler(Generic[StateT]): raise RuntimeError("Run has already completed") 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, + ) + + def close_stream(self) -> None: + self._run.close_stream_sync() + async def aresult(self) -> StateT: return await self._task @@ -219,10 +235,12 @@ class _GraphEngineRun: async_channel_specs: dict[str, _ChannelSpec], entry_point: str, finish_point: str | None, + stream_mode: str | None, ) -> None: self._nodes = nodes self._entry_point = entry_point self._finish_point = finish_point + self._stream_mode = stream_mode self._rust_engine = PyRustEngine() for name in async_channel_specs: self._rust_engine.add_async_channel(name) @@ -242,6 +260,7 @@ class _GraphEngineRun: finish_point, initial_state, self._execute_node_for_rust, + self._stream_mode, ) self._state = result_obj return cast(StateT, self._state) @@ -305,6 +324,15 @@ class _GraphEngineRun: def _publish_sync(self, channel: str, value: Any) -> None: self._rust_engine.publish_obj(channel, value) + def send_custom_stream_event(self, value: Any) -> None: + self._rust_engine.send_custom_stream_event_obj(value) + + def receive_stream_sync(self) -> Any | None: + return self._rust_engine.receive_stream_obj() + + def close_stream_sync(self) -> None: + self._rust_engine.close_stream() + def _execute_node_for_rust( self, node_name: str, node_input: Any, state: Any ) -> dict[str, Any]: diff --git a/saf-python-sdk/tests/advanced-graph/test_streaming.py b/saf-python-sdk/tests/advanced-graph/test_streaming.py new file mode 100644 index 000000000..7666452b8 --- /dev/null +++ b/saf-python-sdk/tests/advanced-graph/test_streaming.py @@ -0,0 +1,56 @@ +import asyncio + +import pytest +from typing_extensions import TypedDict + +from saf_python_sdk.advanced_graph import AdvancedStateGraph, Context +from saf_python_sdk.types import Command, Send + +pytestmark = pytest.mark.anyio + + +class StreamState(TypedDict): + done: bool + + +async def test_custom_stream_receive_and_close() -> None: + graph: AdvancedStateGraph[StreamState] = AdvancedStateGraph(StreamState) + + async def start_node(ctx: Context, state: StreamState) -> Command: + ctx.send_custom_stream_event({"step": "start", "value": 1}) + await asyncio.sleep(0.08) + ctx.send_custom_stream_event({"step": "start", "value": 2}) + return Command(update=state, goto=Send("finish_node", None)) + + async def finish_node(state: StreamState) -> dict[str, bool]: + return {"done": True} + + graph.add_entry_node(start_node) + graph.add_finish_node(finish_node) + + handler = await graph.compile().astart({"done": False}, stream_mode="custom") + + event = await handler.receive_stream() + assert isinstance(event, dict) + assert event["step"] == "start" + assert event["value"] == 1 + + handler.close_stream() + assert await handler.receive_stream() is None + + result = await handler.aresult() + assert result["done"] is True + + +async def test_only_custom_stream_mode_supported() -> None: + graph: AdvancedStateGraph[StreamState] = AdvancedStateGraph(StreamState) + + async def start_node(ctx: Context, state: StreamState) -> Command: + ctx.send_custom_stream_event({"hello": "world"}) + return Command(update=state) + + graph.add_entry_node(start_node) + + handler = await graph.compile().astart({"done": False}, stream_mode="values") + with pytest.raises(Exception, match="only `custom` is supported"): + await handler.aresult()