mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 20:45:05 +02:00
fix: replay subgraph discovery after run start
This commit is contained in:
@@ -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).
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user