fix: replay subgraph discovery after run start

This commit is contained in:
Claude
2026-06-03 00:59:15 -07:00
parent 1ab3709bca
commit c72c9851b3
7 changed files with 339 additions and 22 deletions
+35 -5
View File
@@ -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).
+23 -4
View File
@@ -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"):
+35 -7
View File
@@ -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.
@@ -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:
@@ -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"])))
@@ -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(
@@ -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(