diff --git a/langgraph-go/advancedgraph/graph.go b/langgraph-go/advancedgraph/graph.go index 904b97d87..debea5cf9 100644 --- a/langgraph-go/advancedgraph/graph.go +++ b/langgraph-go/advancedgraph/graph.go @@ -13,6 +13,7 @@ type nodeExecutor func(ctx *Context, input any, state map[string]any) (Command, type AdvancedStateGraph[StateT any] struct { nodes map[string]nodeExecutor asyncChannels []string + customStreams []string entryPoint string finishPoint string stateType reflect.Type @@ -53,6 +54,10 @@ func (g *AdvancedStateGraph[StateT]) AddAsyncChannel(name string) { g.asyncChannels = append(g.asyncChannels, name) } +func (g *AdvancedStateGraph[StateT]) AddCustomOutputStream(name string) { + g.customStreams = append(g.customStreams, name) +} + func (g *AdvancedStateGraph[StateT]) AddEntryNode(fn any) string { name := NodeName(fn) return g.AddEntryNodeAs(name, fn) @@ -79,6 +84,7 @@ func (g *AdvancedStateGraph[StateT]) Compile() *CompiledGraph[StateT] { return &CompiledGraph[StateT]{ nodes: g.nodes, asyncChannels: g.asyncChannels, + customStreams: g.customStreams, entryPoint: g.entryPoint, finishPoint: g.finishPoint, stateType: g.stateType, @@ -88,6 +94,7 @@ func (g *AdvancedStateGraph[StateT]) Compile() *CompiledGraph[StateT] { type CompiledGraph[StateT any] struct { nodes map[string]nodeExecutor asyncChannels []string + customStreams []string entryPoint string finishPoint string stateType reflect.Type @@ -111,8 +118,8 @@ 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) +func (c *Context) SendCustomStreamEvent(streamName string, value any) error { + return c.engine.SendCustomStreamEvent(streamName, value) } type Handler[StateT any] struct { @@ -135,9 +142,9 @@ func (h *Handler[StateT]) WaitForResult() (StateT, error) { return res.state, res.err } -func (h *Handler[StateT]) ReceiveStream() (any, error) { +func (h *Handler[StateT]) ReceiveStream(streamName string) (any, error) { <-h.streamReadyC - event, hasEvent, err := h.engine.ReceiveStream() + event, hasEvent, err := h.engine.ReceiveStream(streamName) if err != nil { return nil, err } @@ -147,9 +154,9 @@ func (h *Handler[StateT]) ReceiveStream() (any, error) { return event, nil } -func (h *Handler[StateT]) CloseStream() error { +func (h *Handler[StateT]) CloseAllStreams() error { <-h.streamReadyC - return h.engine.CloseStream() + return h.engine.CloseAllStreams() } func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT, streamMode ...string) (*Handler[StateT], error) { @@ -166,6 +173,11 @@ func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT, str return nil, err } } + for _, streamName := range g.customStreams { + if err := engine.AddCustomOutputStream(streamName); err != nil { + return nil, err + } + } handler := &Handler[StateT]{ engine: engine, diff --git a/langgraph-go/advancedgraph/rust_engine.go b/langgraph-go/advancedgraph/rust_engine.go index 0d078adde..adbff2177 100644 --- a/langgraph-go/advancedgraph/rust_engine.go +++ b/langgraph-go/advancedgraph/rust_engine.go @@ -128,6 +128,13 @@ func (e *RustEngine) AddAsyncChannel(channel string) error { return parseRustStatus(resp) } +func (e *RustEngine) AddCustomOutputStream(streamName string) error { + cname := C.CString(streamName) + defer C.free(unsafe.Pointer(cname)) + resp := C.rc_add_custom_output_stream(e.ptr, cname) + return parseRustStatus(resp) +} + func (e *RustEngine) StartStream(streamMode string) error { var cmode *C.char if streamMode != "" { @@ -138,8 +145,10 @@ func (e *RustEngine) StartStream(streamMode string) error { return parseRustStatus(resp) } -func (e *RustEngine) ReceiveStream() (any, bool, error) { - resp := C.rc_receive_stream_json(e.ptr) +func (e *RustEngine) ReceiveStream(streamName string) (any, bool, error) { + cname := C.CString(streamName) + defer C.free(unsafe.Pointer(cname)) + resp := C.rc_receive_stream_json(e.ptr, cname) defer C.rc_string_free(resp) raw := C.GoString(resp) @@ -165,19 +174,21 @@ func (e *RustEngine) ReceiveStream() (any, bool, error) { return coerceJSONValue(event), true, nil } -func (e *RustEngine) SendCustomStreamEvent(value any) error { +func (e *RustEngine) SendCustomStreamEvent(streamName string, value any) error { payload, err := json.Marshal(value) if err != nil { return fmt.Errorf("marshal stream event: %w", err) } + cname := C.CString(streamName) cval := C.CString(string(payload)) + defer C.free(unsafe.Pointer(cname)) defer C.free(unsafe.Pointer(cval)) - resp := C.rc_send_custom_stream_event(e.ptr, cval) + resp := C.rc_send_custom_stream_event(e.ptr, cname, cval) return parseRustStatus(resp) } -func (e *RustEngine) CloseStream() error { - resp := C.rc_close_stream(e.ptr) +func (e *RustEngine) CloseAllStreams() error { + resp := C.rc_close_all_streams(e.ptr) return parseRustStatus(resp) } diff --git a/langgraph-go/tests/test_streaming_test.go b/langgraph-go/tests/test_streaming_test.go index 6095d8203..76e6d1277 100644 --- a/langgraph-go/tests/test_streaming_test.go +++ b/langgraph-go/tests/test_streaming_test.go @@ -15,11 +15,11 @@ type streamState struct { 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 { + if err := ctx.SendCustomStreamEvent("high", 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 { + if err := ctx.SendCustomStreamEvent("regular", map[string]any{"step": "start", "value": 2}); err != nil { return ag.Command{}, err } return ag.Command{ @@ -38,6 +38,8 @@ func (w *streamWorkflow) finishNode(ctx *ag.Context, _ any, state streamState) ( func TestCustomStreamReceiveAndClose(t *testing.T) { workflow := &streamWorkflow{} graph := ag.NewAdvancedStateGraph[streamState]() + graph.AddCustomOutputStream("high") + graph.AddCustomOutputStream("regular") graph.AddEntryNode(workflow.startNode) graph.AddFinishNode(workflow.finishNode) @@ -46,7 +48,7 @@ func TestCustomStreamReceiveAndClose(t *testing.T) { t.Fatalf("start failed: %v", err) } - event, err := handler.ReceiveStream() + event, err := handler.ReceiveStream("high") if err != nil { t.Fatalf("receive stream failed: %v", err) } @@ -61,11 +63,23 @@ func TestCustomStreamReceiveAndClose(t *testing.T) { t.Fatalf("unexpected stream event payload: %#v", eventMap) } - if err := handler.CloseStream(); err != nil { - t.Fatalf("close stream failed: %v", err) + eventRegular, err := handler.ReceiveStream("regular") + if err != nil { + t.Fatalf("receive regular stream failed: %v", err) + } + eventRegularMap, ok := eventRegular.(map[string]any) + if !ok { + t.Fatalf("unexpected regular event type: %T", eventRegular) + } + if eventRegularMap["value"] != float64(2) { + t.Fatalf("unexpected regular stream payload: %#v", eventRegularMap) } - closedEvent, err := handler.ReceiveStream() + if err := handler.CloseAllStreams(); err != nil { + t.Fatalf("close all streams failed: %v", err) + } + + closedEvent, err := handler.ReceiveStream("high") if err != nil { t.Fatalf("receive stream after close failed: %v", err) } @@ -85,6 +99,7 @@ func TestCustomStreamReceiveAndClose(t *testing.T) { func TestOnlyCustomStreamModeSupported(t *testing.T) { workflow := &streamWorkflow{} graph := ag.NewAdvancedStateGraph[streamState]() + graph.AddCustomOutputStream("regular") graph.AddEntryNode(workflow.startNode) graph.AddFinishNode(workflow.finishNode) diff --git a/rust-core/include/langgraph_rust_core.h b/rust-core/include/langgraph_rust_core.h index 041a46d18..f01c3dbe7 100644 --- a/rust-core/include/langgraph_rust_core.h +++ b/rust-core/include/langgraph_rust_core.h @@ -17,12 +17,13 @@ Engine* rc_engine_new(void); void rc_engine_free(Engine* ptr); char* rc_add_async_channel(Engine* ptr, const char* channel); +char* rc_add_custom_output_stream(Engine* ptr, const char* stream_name); 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_receive_stream_json(Engine* ptr, const char* stream_name); +char* rc_send_custom_stream_event(Engine* ptr, const char* stream_name, const char* value_json); +char* rc_close_all_streams(Engine* ptr); char* rc_run_graph_json( Engine* ptr, const char* entry_point, diff --git a/rust-core/src/engine.rs b/rust-core/src/engine.rs index fe0a54a18..26d8dd442 100644 --- a/rust-core/src/engine.rs +++ b/rust-core/src/engine.rs @@ -289,7 +289,8 @@ pub fn merge_json_update(state: &mut Value, update: Option) { pub struct Engine { channels: Arc>>>, channel_notify: Arc, - stream: Arc>>, + custom_output_stream_names: Arc>>, + streams: Arc>>>, } #[derive(Clone)] @@ -310,6 +311,21 @@ impl Engine { channels.entry(name.to_owned()).or_default(); } + pub fn add_custom_output_stream(&self, name: &str) -> Result<(), String> { + let stream_name = name.trim(); + if stream_name.is_empty() { + return Err("custom output stream name cannot be empty".to_string()); + } + let mut names = self + .custom_output_stream_names + .lock() + .expect("custom output stream names mutex poisoned"); + if !names.iter().any(|n| n == stream_name) { + names.push(stream_name.to_string()); + } + Ok(()) + } + pub fn start_stream(&self, stream_mode: Option<&str>) -> Result<(), String> { let mode = stream_mode.and_then(|m| { let trimmed = m.trim(); @@ -319,43 +335,62 @@ impl Engine { Some(trimmed) } }); - let mut stream = self.stream.lock().expect("stream mutex poisoned"); + let mut streams = self.streams.lock().expect("streams 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)), - }); + let stream_names = { + self.custom_output_stream_names + .lock() + .expect("custom output stream names mutex poisoned") + .clone() + }; + let mut by_name = HashMap::new(); + for stream_name in stream_names { + let (sender, receiver) = tokio_mpsc::unbounded_channel(); + by_name.insert( + stream_name, + StreamChannel { + sender, + receiver: Arc::new(AsyncMutex::new(receiver)), + }, + ); + } + *streams = Some(by_name); } else { - *stream = None; + *streams = None; } Ok(()) } - pub fn close_stream(&self) { - let mut stream = self.stream.lock().expect("stream mutex poisoned"); - *stream = None; + pub fn close_all_streams(&self) { + let mut streams = self.streams.lock().expect("streams mutex poisoned"); + *streams = None; } - pub fn send_custom_stream_event(&self, value: serde_json::Value) { + pub fn send_custom_stream_event(&self, stream_name: &str, value: serde_json::Value) { let sender = { - let stream = self.stream.lock().expect("stream mutex poisoned"); - stream.as_ref().map(|s| s.sender.clone()) + let streams = self.streams.lock().expect("streams mutex poisoned"); + streams + .as_ref() + .and_then(|m| m.get(stream_name)) + .map(|s| s.sender.clone()) }; if let Some(tx) = sender { let _ = tx.send(value); } } - pub async fn receive_stream_async(&self) -> Option { + pub async fn receive_stream_async(&self, stream_name: &str) -> Option { let receiver = { - let stream = self.stream.lock().expect("stream mutex poisoned"); - stream.as_ref().map(|s| Arc::clone(&s.receiver)) + let streams = self.streams.lock().expect("streams mutex poisoned"); + streams + .as_ref() + .and_then(|m| m.get(stream_name)) + .map(|s| Arc::clone(&s.receiver)) }; let Some(receiver) = receiver else { return None; diff --git a/rust-core/src/lib_c.rs b/rust-core/src/lib_c.rs index 82a88421f..20e3dbaf9 100644 --- a/rust-core/src/lib_c.rs +++ b/rust-core/src/lib_c.rs @@ -144,6 +144,27 @@ 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_name` must be a valid null-terminated UTF-8 string pointer. +pub unsafe extern "C" fn rc_add_custom_output_stream( + ptr: *mut Engine, + stream_name: *const c_char, +) -> *mut c_char { + if ptr.is_null() { + return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string()); + } + let stream_name = match cstr_to_str(stream_name) { + Ok(v) => v, + Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + }; + match (*ptr).add_custom_output_stream(stream_name) { + 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`. @@ -172,11 +193,19 @@ pub unsafe extern "C" fn rc_start_stream( #[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 { +/// `stream_name` must be a valid null-terminated UTF-8 string pointer. +pub unsafe extern "C" fn rc_receive_stream_json( + ptr: *mut Engine, + stream_name: *const c_char, +) -> *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()); + let stream_name = match cstr_to_str(stream_name) { + Ok(v) => v, + Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + }; + let event = run_loop_block_on((*ptr).receive_stream_async(stream_name)); match event { Some(value) => match serde_json::to_string(&value) { Ok(s) => into_c_ptr(format!("{{\"ok\":true,\"has_event\":true,\"event\":{s}}}")), @@ -189,14 +218,20 @@ pub unsafe extern "C" fn rc_receive_stream_json(ptr: *mut Engine) -> *mut c_char #[no_mangle] /// # Safety /// `ptr` must be a valid engine pointer from `rc_engine_new`. +/// `stream_name` must be a valid null-terminated UTF-8 string pointer. /// `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, + stream_name: *const c_char, 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 stream_name = match cstr_to_str(stream_name) { + Ok(v) => v, + Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), + }; let value_json = match cstr_to_str(value_json) { Ok(v) => v, Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")), @@ -209,18 +244,18 @@ pub unsafe extern "C" fn rc_send_custom_stream_event( )) } }; - (*ptr).send_custom_stream_event(value); + (*ptr).send_custom_stream_event(stream_name, 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 { +pub unsafe extern "C" fn rc_close_all_streams(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(); + (*ptr).close_all_streams(); into_c_ptr("{\"ok\":true}".to_string()) } @@ -338,7 +373,7 @@ pub unsafe extern "C" fn rc_run_graph_json( callback_wrapper, ) .await; - run_engine.close_stream(); + run_engine.close_all_streams(); 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 c91d6eadc..9e37557d7 100644 --- a/rust-core/src/lib_py.rs +++ b/rust-core/src/lib_py.rs @@ -36,6 +36,12 @@ impl PyRustEngine { self.inner.add_async_channel(name); } + fn add_custom_output_stream(&self, stream_name: &str) -> PyResult<()> { + self.inner + .add_custom_output_stream(stream_name) + .map_err(PyValueError::new_err) + } + 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) @@ -105,8 +111,9 @@ impl PyRustEngine { .map_err(PyValueError::new_err) } - fn receive_stream_obj(&self, py: Python<'_>) -> PyResult> { - let event = py.allow_threads(|| run_loop_block_on(self.inner.receive_stream_async())); + fn receive_stream_obj(&self, py: Python<'_>, stream_name: &str) -> PyResult> { + let event = + py.allow_threads(|| run_loop_block_on(self.inner.receive_stream_async(stream_name))); match event { Some(value) => { let event_json = serde_json::to_string(&value) @@ -117,16 +124,21 @@ impl PyRustEngine { } } - fn send_custom_stream_event_obj(&self, py: Python<'_>, value: Py) -> PyResult<()> { + fn send_custom_stream_event_obj( + &self, + py: Python<'_>, + stream_name: &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.send_custom_stream_event(parsed); + self.inner.send_custom_stream_event(stream_name, parsed); Ok(()) } - fn close_stream(&self) { - self.inner.close_stream(); + fn close_all_streams(&self) { + self.inner.close_all_streams(); } #[pyo3(signature = (entry_point, finish_point, initial_state, callback, stream_mode=None))] @@ -196,7 +208,7 @@ impl PyRustEngine { )) }); - self.inner.close_stream(); + self.inner.close_all_streams(); let out = run_result.map_err(PyValueError::new_err)?; Ok(out.as_ref().clone_ref(py)) } diff --git a/saf-python-sdk/pyproject.toml b/saf-python-sdk/pyproject.toml index b2c2f126b..5b9c4a0c9 100644 --- a/saf-python-sdk/pyproject.toml +++ b/saf-python-sdk/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "maturin" [project] name = "saf-python-sdk" -version = "0.1.3" +version = "0.1.4" description = "Standalone advanced graph runtime powered by Rust engine" readme = "README.md" requires-python = ">=3.10" 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 fe8c83019..7d2aec017 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 @@ -93,6 +93,7 @@ class AdvancedStateGraph(Generic[StateT]): self.state_schema = state_schema self._nodes: dict[str, Callable[..., Any]] = {} self._async_channels: dict[str, _ChannelSpec] = {} + self._custom_output_streams: dict[str, _ChannelSpec] = {} self._entry_point: str | None = None self._finish_point: str | None = None @@ -122,6 +123,15 @@ class AdvancedStateGraph(Generic[StateT]): raise ValueError(f"Channel `{name}` already exists") self._async_channels[name] = _ChannelSpec(typ=typ) + def add_custom_outout_stream(self, name: str, typ: Any) -> None: + if name in self._custom_output_streams: + raise ValueError(f"Custom output stream `{name}` already exists") + self._custom_output_streams[name] = _ChannelSpec(typ=typ) + + # Alias with corrected spelling. + def add_custom_output_stream(self, name: str, typ: Any) -> None: + self.add_custom_outout_stream(name, typ) + def add_entry_node(self, node: Callable[..., Any]) -> str: node_name = self.add_node(node) self._entry_point = self._resolve_node_name(node_name) @@ -150,6 +160,7 @@ class AdvancedStateGraph(Generic[StateT]): return CompiledGraphEngine( nodes=dict(self._nodes), async_channels=dict(self._async_channels), + custom_output_streams=dict(self._custom_output_streams), entry_point=self._entry_point, finish_point=self._finish_point, ) @@ -163,11 +174,13 @@ class CompiledGraphEngine(Generic[StateT]): *, nodes: dict[str, Callable[..., Any]], async_channels: dict[str, _ChannelSpec], + custom_output_streams: dict[str, _ChannelSpec], entry_point: str, finish_point: str | None, ) -> None: self._nodes = nodes self._async_channels = async_channels + self._custom_output_streams = custom_output_streams self._entry_point = entry_point self._finish_point = finish_point @@ -181,6 +194,7 @@ class CompiledGraphEngine(Generic[StateT]): run = _GraphEngineRun( nodes=self._nodes, async_channel_specs=self._async_channels, + custom_output_stream_specs=self._custom_output_streams, entry_point=self._entry_point, finish_point=self._finish_point, stream_mode=stream_mode, @@ -207,8 +221,8 @@ 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) + def send_custom_stream_event(self, stream_name: str, value: Any) -> None: + self._run.send_custom_stream_event(stream_name, value) class GraphRunHandler(Generic[StateT]): @@ -223,13 +237,13 @@ class GraphRunHandler(Generic[StateT]): raise RuntimeError("Run has already completed") await self._run.publish(channel, value) - async def receive_stream(self) -> Any | None: + async def receive_stream(self, stream_name: str) -> Any | None: # 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) + return await asyncio.to_thread(self._run.receive_stream_sync, stream_name) - def close_stream(self) -> None: - self._run.close_stream_sync() + def close_all_streams(self) -> None: + self._run.close_all_streams_sync() async def aresult(self) -> StateT: return await self._task @@ -244,6 +258,7 @@ class _GraphEngineRun: *, nodes: dict[str, Callable[..., Any]], async_channel_specs: dict[str, _ChannelSpec], + custom_output_stream_specs: dict[str, _ChannelSpec], entry_point: str, finish_point: str | None, stream_mode: str | None, @@ -255,6 +270,8 @@ class _GraphEngineRun: self._rust_engine = PyRustEngine() for name in async_channel_specs: self._rust_engine.add_async_channel(name) + for stream_name in custom_output_stream_specs: + self._rust_engine.add_custom_output_stream(stream_name) self._stream_ready = threading.Event() if self._stream_mode is None: self._stream_ready.set() @@ -361,15 +378,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 send_custom_stream_event(self, stream_name: str, value: Any) -> None: + self._rust_engine.send_custom_stream_event_obj(stream_name, value) - def receive_stream_sync(self) -> Any | None: + def receive_stream_sync(self, stream_name: str) -> Any | None: self._stream_ready.wait() - return self._rust_engine.receive_stream_obj() + return self._rust_engine.receive_stream_obj(stream_name) - def close_stream_sync(self) -> None: - self._rust_engine.close_stream() + def close_all_streams_sync(self) -> None: + self._rust_engine.close_all_streams() def _execute_node_for_rust( self, node_name: str, node_input: Any, state: Any diff --git a/saf-python-sdk/tests/advanced-graph/test_streaming.py b/saf-python-sdk/tests/advanced-graph/test_streaming.py index 7666452b8..b7e38167e 100644 --- a/saf-python-sdk/tests/advanced-graph/test_streaming.py +++ b/saf-python-sdk/tests/advanced-graph/test_streaming.py @@ -15,11 +15,13 @@ class StreamState(TypedDict): async def test_custom_stream_receive_and_close() -> None: graph: AdvancedStateGraph[StreamState] = AdvancedStateGraph(StreamState) + graph.add_custom_outout_stream("high", dict[str, int | str]) + graph.add_custom_outout_stream("regular", dict[str, int | str]) async def start_node(ctx: Context, state: StreamState) -> Command: - ctx.send_custom_stream_event({"step": "start", "value": 1}) + ctx.send_custom_stream_event("high", {"step": "start", "value": 1}) await asyncio.sleep(0.08) - ctx.send_custom_stream_event({"step": "start", "value": 2}) + ctx.send_custom_stream_event("regular", {"step": "start", "value": 2}) return Command(update=state, goto=Send("finish_node", None)) async def finish_node(state: StreamState) -> dict[str, bool]: @@ -30,13 +32,18 @@ async def test_custom_stream_receive_and_close() -> None: handler = await graph.compile().astart({"done": False}, stream_mode="custom") - event = await handler.receive_stream() + event = await handler.receive_stream("high") assert isinstance(event, dict) assert event["step"] == "start" assert event["value"] == 1 - handler.close_stream() - assert await handler.receive_stream() is None + event_regular = await handler.receive_stream("regular") + assert isinstance(event_regular, dict) + assert event_regular["value"] == 2 + + handler.close_all_streams() + assert await handler.receive_stream("high") is None + assert await handler.receive_stream("regular") is None result = await handler.aresult() assert result["done"] is True @@ -44,9 +51,10 @@ async def test_custom_stream_receive_and_close() -> None: async def test_only_custom_stream_mode_supported() -> None: graph: AdvancedStateGraph[StreamState] = AdvancedStateGraph(StreamState) + graph.add_custom_outout_stream("regular", dict[str, str]) async def start_node(ctx: Context, state: StreamState) -> Command: - ctx.send_custom_stream_event({"hello": "world"}) + ctx.send_custom_stream_event("regular", {"hello": "world"}) return Command(update=state) graph.add_entry_node(start_node) diff --git a/saf-python-sdk/uv.lock b/saf-python-sdk/uv.lock index 4c924866b..d90e910ea 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.3" +version = "0.1.4" source = { editable = "." }