stream-draft

This commit is contained in:
Quanzheng Long
2026-03-17 16:12:16 -07:00
parent 5a9264d124
commit e13004da77
9 changed files with 512 additions and 36 deletions
+28 -1
View File
@@ -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) {
+62
View File
@@ -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),
)
+102
View File
@@ -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)
}
}
+5
View File
@@ -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
);
+67 -7
View File
@@ -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<Value>) {
pub struct Engine {
channels: Arc<StdMutex<HashMap<String, VecDeque<serde_json::Value>>>>,
channel_notify: Arc<Notify>,
stream: Arc<StdMutex<Option<StreamChannel>>>,
}
#[derive(Clone)]
struct StreamChannel {
sender: tokio_mpsc::UnboundedSender<serde_json::Value>,
receiver: Arc<AsyncMutex<tokio_mpsc::UnboundedReceiver<serde_json::Value>>>,
}
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<serde_json::Value> {
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<WaitEvent, String> {
pub async fn wait_for_any_of_async(
&self,
any_of: &AnyOfCondition,
) -> Result<WaitEvent, String> {
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<Option<WaitEvent>, String> {
fn try_take_channel_event(&self, channel: &str, n: usize) -> Result<Option<WaitEvent>, String> {
let mut channels = self.channels.lock().expect("channels mutex poisoned");
let queue = channels
.get_mut(channel)
+100 -2
View File
@@ -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<CNodeCallback>,
) -> *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::<Result<Value, String>>();
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 {
+63 -25
View File
@@ -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<Py<PyAny>> {
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<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);
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<PyAny>,
callback: Py<PyAny>,
stream_mode: Option<&str>,
) -> PyResult<Py<PyAny>> {
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::<Result<(), String>>();
let state_for_run = Arc::clone(&state);
let engine = self.inner.clone();
run_loop_spawn(async move {
let run_result =
run_graph_scheduler(entry_point, finish_point, callback, state_for_run, engine)
.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<NodeExecutionPy, String> {
let callback_bound = callback.as_ref().bind(py);
let payload_obj = callback_bound
.call1((node.as_str(), arg, (*state_for_task).clone_ref(py)))
.map_err(|e| format!("callback failed for node `{node}`: {e}"))?;
let payload = parse_node_outcome(py, &payload_obj)
.map_err(|e| format!("invalid callback payload for `{node}`: {e}"))?;
Ok(NodeExecutionPy {
node: node_for_result,
arg: arg_for_result,
outcome: payload,
})
},
);
let outcome = Python::with_gil(|py| -> Result<NodeExecutionPy, String> {
let callback_bound = callback.as_ref().bind(py);
let payload_obj = callback_bound
.call1((node.as_str(), arg, (*state_for_task).clone_ref(py)))
.map_err(|e| format!("callback failed for node `{node}`: {e}"))?;
let payload = parse_node_outcome(py, &payload_obj)
.map_err(|e| format!("invalid callback payload for `{node}`: {e}"))?;
Ok(NodeExecutionPy {
node: node_for_result,
arg: arg_for_result,
outcome: payload,
})
});
let _ = tx.send(SchedulerEventPy::Node(outcome));
})
}
@@ -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<PyAny>, event: &WaitEvent) -> Result<Py<PyAny>, 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}"))?;
@@ -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]:
@@ -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()