From c72c9851b33c4b2bf01939b91d7503c677969f5a Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 3 Jun 2026 00:25:12 -0700 Subject: [PATCH] fix: replay subgraph discovery after run start --- libs/sdk-py/langgraph_sdk/_async/stream.py | 40 ++++++++-- libs/sdk-py/langgraph_sdk/_sync/stream.py | 27 ++++++- libs/sdk-py/langgraph_sdk/stream/decoders.py | 42 ++++++++-- .../langgraph_sdk/stream/sync_controller.py | 42 ++++++++-- libs/sdk-py/tests/streaming/test_decoders.py | 53 +++++++++++++ .../tests/streaming/test_scoped_handles.py | 79 +++++++++++++++++++ .../streaming/test_sync_scoped_handles.py | 78 ++++++++++++++++++ 7 files changed, 339 insertions(+), 22 deletions(-) diff --git a/libs/sdk-py/langgraph_sdk/_async/stream.py b/libs/sdk-py/langgraph_sdk/_async/stream.py index 870ce2556..2ebc2d111 100644 --- a/libs/sdk-py/langgraph_sdk/_async/stream.py +++ b/libs/sdk-py/langgraph_sdk/_async/stream.py @@ -105,6 +105,8 @@ def _event_namespace(params_field: Any) -> list[str]: _ROOT_TERMINAL_LIFECYCLE_EVENTS = frozenset({"completed", "failed"}) +# `since` returns events after the given seq; -1 means replay from the first event. +_INITIAL_REPLAY_CURSOR = -1 def _is_root_terminal_lifecycle(event: Any) -> bool: @@ -126,6 +128,9 @@ def _is_root_terminal_lifecycle(event: Any) -> bool: data = params.get("data") or {} if not isinstance(data, dict): return False + payload_namespace = data.get("namespace") + if isinstance(payload_namespace, list) and payload_namespace: + return False return data.get("event") in _ROOT_TERMINAL_LIFECYCLE_EVENTS @@ -174,6 +179,7 @@ class RunModule: metadata: dict[str, Any] | None = None, ) -> dict[str, Any]: """Send `run.start` to the server. Returns the result (`{"run_id": ...}`).""" + replay_cursor = self._owner._cursor params: dict[str, Any] = {"assistant_id": self._owner.assistant_id} if input is not None: params["input"] = input @@ -186,6 +192,8 @@ class RunModule: self._owner._run_start_ready = gate try: result = await self._owner._send_command("run.start", params) + self._owner._run_start_replay_cursor = replay_cursor + self._owner._has_run_start_replay_cursor = True if not gate.done(): gate.set_result(None) self._owner._run_seen = True @@ -931,7 +939,10 @@ class _SubgraphsProjection: self._thread._activate_root_messages_inbox() if not self._scope else None ) try: - await self._thread._reconcile_stream(params) + await self._thread._reconcile_stream( + params, + use_run_start_replay_cursor=self._scope == (), + ) self._thread._ensure_fanout_running() while True: item = await sub.queue.get() @@ -1215,6 +1226,8 @@ class AsyncThreadStream: self._run_seen: bool = False self._run_done: asyncio.Future[_RunTerminal] | None = None self._cursor: int | None = None + self._run_start_replay_cursor: int | None = None + self._has_run_start_replay_cursor = False self._active_message_streams: set[AsyncChatModelStream] = set() self._active_tool_calls: set[ToolCallHandle] = set() # Root-scope inbox: populated by `_SubgraphsProjection` when it consumes @@ -1705,7 +1718,12 @@ class AsyncThreadStream: return True return False - async def _reconcile_stream(self, candidate_filter: SubscribeParams) -> None: + async def _reconcile_stream( + self, + candidate_filter: SubscribeParams, + *, + use_run_start_replay_cursor: bool = False, + ) -> None: """Ensure the shared SSE covers `candidate_filter`. Rotate if not. Open-new-before-close-old: any events buffered server-side between @@ -1721,8 +1739,17 @@ class AsyncThreadStream: if self._transport is None: raise RuntimeError("AsyncThreadStream not entered — use `async with`.") + replay_cursor = self._cursor + force_replay = False + if use_run_start_replay_cursor and self._has_run_start_replay_cursor: + replay_cursor = self._run_start_replay_cursor + if replay_cursor is None: + replay_cursor = _INITIAL_REPLAY_CURSOR + force_replay = replay_cursor != self._cursor + if ( - self._shared_stream is not None + not force_replay + and self._shared_stream is not None and self._shared_stream_filter is not None and filter_covers(self._shared_stream_filter, dict(candidate_filter)) ): @@ -1730,8 +1757,8 @@ class AsyncThreadStream: new_filter = self._compute_current_union(extra=candidate_filter) stream_params: dict[str, Any] = dict(new_filter) - if self._cursor is not None: - stream_params["since"] = self._cursor + if replay_cursor is not None: + stream_params["since"] = replay_cursor new_stream = self._transport.open_event_stream(stream_params) old_stream = self._shared_stream self._shared_stream = new_stream @@ -1964,6 +1991,9 @@ class AsyncThreadStream: params = event.get("params") or {} data = params.get("data") if isinstance(params, dict) else None phase = data.get("event") if isinstance(data, dict) else None + payload_namespace = data.get("namespace") if isinstance(data, dict) else None + if isinstance(payload_namespace, list) and payload_namespace: + return if phase in ("started", "running"): # Mark that we have observed an active run so thread.output # knows a run exists (handles reattach without run.start). diff --git a/libs/sdk-py/langgraph_sdk/_sync/stream.py b/libs/sdk-py/langgraph_sdk/_sync/stream.py index 9405e7f83..ff72fda5a 100644 --- a/libs/sdk-py/langgraph_sdk/_sync/stream.py +++ b/libs/sdk-py/langgraph_sdk/_sync/stream.py @@ -217,6 +217,8 @@ class SyncRunModule: metadata: dict[str, Any] | None = None, ) -> dict[str, Any]: """Send `run.start` to the server. Returns the result (`{"run_id": ...}`).""" + controller = self._owner._controller + replay_cursor = controller._cursor if controller is not None else None params: dict[str, Any] = {"assistant_id": self._owner.assistant_id} if input is not None: params["input"] = input @@ -226,8 +228,8 @@ class SyncRunModule: params["metadata"] = metadata result = self._owner._send_command("run.start", params) self._owner._run_seen = True - controller = self._owner._controller if controller is not None and controller._run_start_gate is not None: + controller.mark_run_start_replay_cursor(replay_cursor) controller._run_start_gate.set() return result @@ -974,7 +976,10 @@ class _SyncSubgraphsProjection: self._thread._activate_root_messages_inbox() if not self._scope else None ) try: - self._thread._reconcile_stream(params) + self._thread._reconcile_stream( + params, + use_run_start_replay_cursor=self._scope == (), + ) self._thread._ensure_fanout_running() while True: item = sub.queue.get() @@ -1215,10 +1220,21 @@ class SyncThreadStream: if self._controller is not None: self._controller.ensure_fanout_running() - def _reconcile_stream(self, candidate_filter: SubscribeParams) -> None: + def _reconcile_stream( + self, + candidate_filter: SubscribeParams, + *, + use_run_start_replay_cursor: bool = False, + ) -> None: if self._controller is None: raise RuntimeError("SyncThreadStream not entered — use `with`.") - self._controller.reconcile_stream(candidate_filter) + if not use_run_start_replay_cursor: + self._controller.reconcile_stream(candidate_filter) + return + self._controller.reconcile_stream( + candidate_filter, + use_run_start_replay_cursor=use_run_start_replay_cursor, + ) def _activate_root_messages_inbox(self) -> queue.Queue[Event | None]: if self._root_messages_inbox is None: @@ -1608,6 +1624,9 @@ class SyncThreadStream: params = event.get("params") or {} data = params.get("data") if isinstance(params, dict) else None phase = data.get("event") if isinstance(data, dict) else None + payload_namespace = data.get("namespace") if isinstance(data, dict) else None + if isinstance(payload_namespace, list) and payload_namespace: + return if phase in ("started", "running"): self._run_seen = True elif phase in ("completed", "failed"): diff --git a/libs/sdk-py/langgraph_sdk/stream/decoders.py b/libs/sdk-py/langgraph_sdk/stream/decoders.py index 45959bf46..af07884d3 100644 --- a/libs/sdk-py/langgraph_sdk/stream/decoders.py +++ b/libs/sdk-py/langgraph_sdk/stream/decoders.py @@ -56,6 +56,15 @@ def _event_namespace(params_field: Any) -> list[str]: return list(namespace) if isinstance(namespace, list) else [] +def _lifecycle_payload_namespace( + params_field: Mapping[str, Any], data: Mapping[str, Any] +) -> list[str]: + namespace = data.get("namespace") + if isinstance(namespace, list): + return list(namespace) + return _event_namespace(params_field) + + def _message_event_id(data: dict[str, Any]) -> str | None: message_id = data.get("id") or data.get("message_id") return str(message_id) if message_id is not None else None @@ -275,11 +284,15 @@ class SubgraphsDecoder: def feed(self, event: Mapping[str, Any]) -> Iterable[Any]: params = event.get("params") or {} - namespace = _event_namespace(params) data = params.get("data") if not isinstance(data, dict): return method = event.get("method") + namespace = ( + _lifecycle_payload_namespace(params, data) + if method == "lifecycle" + else _event_namespace(params) + ) # 1. Fanout: first active child whose path prefixes this namespace. ns_tuple = tuple(namespace) @@ -295,12 +308,12 @@ class SubgraphsDecoder: self._apply_tasks_result(namespace, data) elif _is_direct_child(namespace, self._scope): yield from self._discover(namespace) - elif ( - method == "lifecycle" - and data.get("event") == "started" - and _is_direct_child(namespace, self._scope) - ): - yield from self._discover(namespace) + elif method == "lifecycle": + event_type = data.get("event") + if event_type == "started" and _is_direct_child(namespace, self._scope): + yield from self._discover(namespace) + elif event_type in ("completed", "failed", "interrupted"): + self._apply_lifecycle_terminal(namespace, data) def _discover(self, namespace: list[str]) -> Iterable[Any]: path = tuple(namespace) @@ -330,6 +343,21 @@ class SubgraphsDecoder: handle._finish(status, error) del self._active[child_path] + def _apply_lifecycle_terminal( + self, namespace: list[str], data: dict[str, Any] + ) -> None: + path = tuple(namespace) + handle = self._active.pop(path, None) + if handle is None: + return + event_type = data.get("event") + if event_type == "failed": + handle._finish("failed", data.get("error")) + elif event_type == "interrupted": + handle._finish("interrupted", None) + else: + handle._finish("completed", None) + class ExtensionsDecoder: """Yields `params.data` from one named custom channel. diff --git a/libs/sdk-py/langgraph_sdk/stream/sync_controller.py b/libs/sdk-py/langgraph_sdk/stream/sync_controller.py index 165c7bfdf..34dadce61 100644 --- a/libs/sdk-py/langgraph_sdk/stream/sync_controller.py +++ b/libs/sdk-py/langgraph_sdk/stream/sync_controller.py @@ -20,9 +20,12 @@ from langgraph_sdk.stream.transport import ( ) _logger = logging.getLogger(__name__) +_CURSOR_UNSET = object() _ROOT_TERMINAL_LIFECYCLE_EVENTS = frozenset({"completed", "failed"}) +# `since` returns events after the given seq; -1 means replay from the first event. +_INITIAL_REPLAY_CURSOR = -1 def _is_root_terminal_lifecycle(event: Any) -> bool: @@ -44,6 +47,9 @@ def _is_root_terminal_lifecycle(event: Any) -> bool: data = params.get("data") or {} if not isinstance(data, dict): return False + payload_namespace = data.get("namespace") + if isinstance(payload_namespace, list) and payload_namespace: + return False return data.get("event") in _ROOT_TERMINAL_LIFECYCLE_EVENTS @@ -91,6 +97,8 @@ class SyncStreamController: self._reconnect_backoff_base = reconnect_backoff_base self._reconnect_backoff_cap = reconnect_backoff_cap self._drain_threads: set[threading.Thread] = set() + self._run_start_replay_cursor: int | None = None + self._has_run_start_replay_cursor = False def register_subscription(self, params: SubscribeParams) -> _SyncSubscription: with self._lock: @@ -116,14 +124,32 @@ class SyncStreamController: for sub in subs: sub.queue.put(None) - def reconcile_stream(self, candidate_filter: SubscribeParams) -> None: + def mark_run_start_replay_cursor(self, cursor: int | None) -> None: + with self._lock: + self._run_start_replay_cursor = cursor + self._has_run_start_replay_cursor = True + + def reconcile_stream( + self, + candidate_filter: SubscribeParams, + *, + use_run_start_replay_cursor: bool = False, + ) -> None: if self._run_start_gate is not None and not self._run_start_gate.wait( timeout=self._run_start_timeout ): raise TimeoutError("Sync run.start gate timeout.") with self._lock: + replay_cursor = self._cursor + force_replay = False + if use_run_start_replay_cursor and self._has_run_start_replay_cursor: + replay_cursor = self._run_start_replay_cursor + if replay_cursor is None: + replay_cursor = _INITIAL_REPLAY_CURSOR + force_replay = replay_cursor != self._cursor if ( - self._shared_stream is not None + not force_replay + and self._shared_stream is not None and self._shared_stream_filter is not None and filter_covers(self._shared_stream_filter, dict(candidate_filter)) ): @@ -131,7 +157,7 @@ class SyncStreamController: new_filter = self._compute_current_union(extra=candidate_filter) old_stream = self._shared_stream self._shared_stream = self._transport.open_event_stream( - self._filter_with_since(new_filter) + self._filter_with_since(new_filter, replay_cursor) ) self._shared_stream_filter = new_filter if old_stream is not None: @@ -243,10 +269,14 @@ class SyncStreamController: if isinstance(seq, int) and (self._cursor is None or seq > self._cursor): self._cursor = seq - def _filter_with_since(self, params: dict[str, Any]) -> dict[str, Any]: + def _filter_with_since( + self, params: dict[str, Any], cursor: Any = _CURSOR_UNSET + ) -> dict[str, Any]: out = dict(params) - if self._cursor is not None: - out["since"] = self._cursor + if cursor is _CURSOR_UNSET: + cursor = self._cursor + if cursor is not None: + out["since"] = cursor return out def _dedup_iter(self, source: Any) -> Any: diff --git a/libs/sdk-py/tests/streaming/test_decoders.py b/libs/sdk-py/tests/streaming/test_decoders.py index c1ebcd98d..94f9c10c7 100644 --- a/libs/sdk-py/tests/streaming/test_decoders.py +++ b/libs/sdk-py/tests/streaming/test_decoders.py @@ -340,6 +340,16 @@ def _scoped_factory(*, path, graph_name, trigger_call_id): ) +def _forwarded_lifecycle_event(seq: int, **data: Any) -> dict[str, Any]: + return { + "type": "event", + "method": "lifecycle", + "params": {"namespace": [], "data": data}, + "seq": seq, + "event_id": f"evt-{seq}", + } + + def test_subgraphs_decoder_discovers_on_lifecycle_started_once(): decoder = SubgraphsDecoder(scope=(), handle_factory=_scoped_factory) [h] = list(decoder.feed(lifecycle_started_event(seq=1, namespace=["child"]))) @@ -347,6 +357,49 @@ def test_subgraphs_decoder_discovers_on_lifecycle_started_once(): assert list(decoder.feed(lifecycle_started_event(seq=2, namespace=["child"]))) == [] +def test_subgraphs_decoder_discovers_on_forwarded_lifecycle_payload_namespace(): + decoder = SubgraphsDecoder(scope=(), handle_factory=_scoped_factory) + [h] = list( + decoder.feed( + _forwarded_lifecycle_event( + seq=1, + event="started", + namespace=["child:call-1"], + graph_name="child", + trigger_call_id="call-1", + ) + ) + ) + assert h.path == ("child:call-1",) + assert h.graph_name == "child" + assert h.trigger_call_id == "call-1" + + +def test_subgraphs_decoder_completes_on_forwarded_lifecycle_payload_namespace(): + decoder = SubgraphsDecoder(scope=(), handle_factory=_scoped_factory) + [h] = list( + decoder.feed( + _forwarded_lifecycle_event( + seq=1, + event="started", + namespace=["child:call-1"], + graph_name="child", + trigger_call_id="call-1", + ) + ) + ) + list( + decoder.feed( + _forwarded_lifecycle_event( + seq=2, + event="completed", + namespace=["child:call-1"], + ) + ) + ) + assert h.status == "completed" + + def test_subgraphs_decoder_discovers_on_tasks_start_without_result(): decoder = SubgraphsDecoder(scope=(), handle_factory=_scoped_factory) [h] = list(decoder.feed(tasks_start_event(seq=1, namespace=["child"]))) diff --git a/libs/sdk-py/tests/streaming/test_scoped_handles.py b/libs/sdk-py/tests/streaming/test_scoped_handles.py index 13955b225..a3ad52905 100644 --- a/libs/sdk-py/tests/streaming/test_scoped_handles.py +++ b/libs/sdk-py/tests/streaming/test_scoped_handles.py @@ -22,6 +22,16 @@ from streaming._events import ( from streaming._fake_server import FakeServer, _StreamScript +def _forwarded_lifecycle_event(seq: int, **data): + return { + "type": "event", + "method": "lifecycle", + "params": {"namespace": [], "data": data}, + "seq": seq, + "event_id": f"evt-{seq}", + } + + async def test_subgraphs_subscribes_to_tasks_channel(): fake = FakeServer() fake.script([lifecycle_completed_event(seq=1)]) @@ -64,6 +74,75 @@ async def test_subgraphs_yields_handle_and_completes_status(): assert handle.error is None +async def test_subgraphs_yields_handle_from_forwarded_lifecycle_namespace(): + fake = FakeServer() + fake.script( + [ + lifecycle_started_event(seq=0), + _forwarded_lifecycle_event( + seq=1, + event="started", + namespace=["worker:abc"], + graph_name="worker", + trigger_call_id="abc", + ), + _forwarded_lifecycle_event( + seq=2, + event="completed", + namespace=["worker:abc"], + ), + lifecycle_completed_event(seq=3), + ] + ) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + await thread.run.start(input={}) + handles = [handle async for handle in thread.subgraphs] + + assert len(handles) == 1 + handle = handles[0] + assert handle.path == ("worker:abc",) + assert handle.graph_name == "worker" + assert handle.trigger_call_id == "abc" + assert handle.status == "completed" + + +async def test_subgraphs_replays_from_before_run_start_cursor(): + fake = FakeServer() + fake.script_command_response( + { + "type": "success", + "id": None, + "result": {"run_id": "run-1"}, + "meta": {"applied_through_seq": 3}, + } + ) + fake.script( + [ + lifecycle_started_event(seq=1), + tasks_start_event(seq=2, namespace=["worker:abc"], task_id="t-child"), + lifecycle_completed_event(seq=3), + ] + ) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + await thread.run.start(input={}) + handles = [handle async for handle in thread.subgraphs] + + assert [handle.path for handle in handles] == [("worker:abc",)] + subgraph_requests = [ + body + for body in fake.stream_request_bodies + if "tasks" in (body.get("channels") or []) + ] + assert subgraph_requests + assert subgraph_requests[-1]["since"] == -1 + + async def test_subgraphs_failed_and_interrupted_statuses(): failed = FakeServer() failed.script( diff --git a/libs/sdk-py/tests/streaming/test_sync_scoped_handles.py b/libs/sdk-py/tests/streaming/test_sync_scoped_handles.py index 198ee2e8b..fdd6f3dbb 100644 --- a/libs/sdk-py/tests/streaming/test_sync_scoped_handles.py +++ b/libs/sdk-py/tests/streaming/test_sync_scoped_handles.py @@ -25,6 +25,17 @@ from streaming._events import ( ) from streaming._sync_fake_server import SyncFakeServer + +def _forwarded_lifecycle_event(seq: int, **data): + return { + "type": "event", + "method": "lifecycle", + "params": {"namespace": [], "data": data}, + "seq": seq, + "event_id": f"evt-{seq}", + } + + # --------------------------------------------------------------------------- # Task 11.1: _finish idempotency — double-finish must not double-close inboxes # --------------------------------------------------------------------------- @@ -122,6 +133,73 @@ def test_sync_subgraphs_yields_handle_and_completes_status(): assert handle.error is None +def test_sync_subgraphs_yields_handle_from_forwarded_lifecycle_namespace(): + fake = SyncFakeServer() + fake.script( + [ + lifecycle_started_event(seq=0), + _forwarded_lifecycle_event( + seq=1, + event="started", + namespace=["worker:abc"], + graph_name="worker", + trigger_call_id="abc", + ), + _forwarded_lifecycle_event( + seq=2, + event="completed", + namespace=["worker:abc"], + ), + lifecycle_completed_event(seq=3), + ] + ) + with httpx.Client(transport=fake.transport, base_url="http://test") as raw: + threads = SyncThreadsClient(SyncHttpClient(raw)) + with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + thread.run.start(input={}) + handles = list(thread.subgraphs) + + assert len(handles) == 1 + handle = handles[0] + assert handle.path == ("worker:abc",) + assert handle.graph_name == "worker" + assert handle.trigger_call_id == "abc" + assert handle.status == "completed" + + +def test_sync_subgraphs_replays_from_before_run_start_cursor(): + fake = SyncFakeServer() + fake.script_command_response( + { + "type": "success", + "id": None, + "result": {"run_id": "run-1"}, + "meta": {"applied_through_seq": 3}, + } + ) + fake.script( + [ + lifecycle_started_event(seq=1), + tasks_start_event(seq=2, namespace=["worker:abc"], task_id="t-child"), + lifecycle_completed_event(seq=3), + ] + ) + with httpx.Client(transport=fake.transport, base_url="http://test") as raw: + threads = SyncThreadsClient(SyncHttpClient(raw)) + with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + thread.run.start(input={}) + handles = list(thread.subgraphs) + + assert [handle.path for handle in handles] == [("worker:abc",)] + subgraph_requests = [ + body + for body in fake.stream_request_bodies + if "tasks" in (body.get("channels") or []) + ] + assert subgraph_requests + assert subgraph_requests[-1]["since"] == -1 + + def test_sync_subgraphs_failed_and_interrupted_statuses(): failed_fake = SyncFakeServer() failed_fake.script(