mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-22 09:35:07 +02:00
named-stream
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
+52
-17
@@ -289,7 +289,8 @@ pub fn merge_json_update(state: &mut Value, update: Option<Value>) {
|
||||
pub struct Engine {
|
||||
channels: Arc<StdMutex<HashMap<String, VecDeque<serde_json::Value>>>>,
|
||||
channel_notify: Arc<Notify>,
|
||||
stream: Arc<StdMutex<Option<StreamChannel>>>,
|
||||
custom_output_stream_names: Arc<StdMutex<Vec<String>>>,
|
||||
streams: Arc<StdMutex<Option<HashMap<String, StreamChannel>>>>,
|
||||
}
|
||||
|
||||
#[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<serde_json::Value> {
|
||||
pub async fn receive_stream_async(&self, stream_name: &str) -> Option<serde_json::Value> {
|
||||
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;
|
||||
|
||||
+41
-6
@@ -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 {
|
||||
|
||||
+19
-7
@@ -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<PyAny>) -> 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<Py<PyAny>> {
|
||||
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<Py<PyAny>> {
|
||||
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<PyAny>) -> PyResult<()> {
|
||||
fn send_custom_stream_event_obj(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
stream_name: &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.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))
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Generated
+1
-1
@@ -4,5 +4,5 @@ requires-python = ">=3.10"
|
||||
|
||||
[[package]]
|
||||
name = "saf-python-sdk"
|
||||
version = "0.1.3"
|
||||
version = "0.1.4"
|
||||
source = { editable = "." }
|
||||
|
||||
Reference in New Issue
Block a user