feat(sdk-py): add async stream reconnect support (#7825)

This commit is contained in:
Nick Hollon
2026-05-27 13:30:01 -04:00
committed by GitHub
parent bb9cfe7a22
commit 10b701cf41
10 changed files with 864 additions and 23 deletions
@@ -758,6 +758,50 @@ class _HandleSubgraphsProjection:
def __aiter__(self) -> AsyncIterator[ScopedStreamHandle]:
return self._subgraphs_iter()
def _route_sibling_inboxes_to_grandchildren(
self,
active: dict[tuple[str, ...], ScopedStreamHandle],
) -> None:
"""Drain non-blocking events from parent's messages/tools inboxes to grandchildren.
Called after each tasks event so grandchild handles receive events that
were enqueued in the parent handle's inboxes before (or just after) the
grandchild was discovered.
"""
parent = self._handle
for inbox_attr, grandchild_attr in (
("_messages_inbox", "_messages_inbox"),
("_tools_inbox", "_tools_inbox"),
):
inbox: asyncio.Queue[Event | None] = getattr(parent, inbox_attr)
staging: list[Event | None] = []
# Drain without blocking.
while not inbox.empty():
staging.append(inbox.get_nowait())
for event in staging:
if event is None:
# Re-queue the EOF sentinel — it belongs to the parent inbox consumer.
inbox.put_nowait(None)
continue
event_params = event.get("params") or {}
ns_tuple = tuple(_event_namespace(event_params))
routed = False
for _child_path, grandchild in active.items():
grandchild_len = len(grandchild.path)
if (
len(ns_tuple) >= grandchild_len
and ns_tuple[:grandchild_len] == grandchild.path
):
gc_inbox: asyncio.Queue[Event | None] = getattr(
grandchild, grandchild_attr
)
gc_inbox.put_nowait(event)
routed = True
break
if not routed:
# Not a grandchild event — put it back for the handle projection.
inbox.put_nowait(event)
async def _subgraphs_iter(self) -> AsyncGenerator[ScopedStreamHandle, None]:
self._handle._mark_iterated("tasks")
seen: set[tuple[str, ...]] = set()
+107 -10
View File
@@ -14,8 +14,10 @@ from __future__ import annotations
import asyncio
import contextlib
import logging
import random
from collections import OrderedDict
from collections.abc import AsyncGenerator, AsyncIterator
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable
from dataclasses import dataclass, field
from typing import Any
@@ -56,6 +58,8 @@ class _SeenEventIds:
# Per-subscription record
# ---------------------------------------------------------------------------
_logger = logging.getLogger(__name__)
@dataclass
class _Subscription:
@@ -108,8 +112,12 @@ class StreamController:
self,
*,
transport: Any,
run_start_gate: Callable[[], Awaitable[None]] | None = None, # noqa: ARG002
max_queue_size: int = 1024,
seen_event_ids_max: int = 10_000,
max_reconnect_attempts: int = 5,
reconnect_backoff_base: float = 0.1,
reconnect_backoff_cap: float = 2.0,
) -> None:
self._transport = transport
self._max_queue_size = max_queue_size
@@ -121,6 +129,10 @@ class StreamController:
self._fanout_task: asyncio.Task[None] | None = None
self._rotation_close_tasks: set[asyncio.Task[None]] = set()
self._closed = False
self._cursor: int | None = None
self._max_reconnect_attempts = max_reconnect_attempts
self._reconnect_backoff_base = reconnect_backoff_base
self._reconnect_backoff_cap = reconnect_backoff_cap
# ------------------------------------------------------------------
# Public API
@@ -179,6 +191,10 @@ class StreamController:
"""Remove a subscription from the registry. No-op if already absent."""
self._subscriptions.pop(subscription_id, None)
# Public aliases used by tests and external callers.
register_subscription = _register_subscription
unregister_subscription = _unregister_subscription
async def _subscription_iter(
self, params: SubscribeParams
) -> AsyncGenerator[Event, None]:
@@ -204,6 +220,9 @@ class StreamController:
if self._fanout_task is None or self._fanout_task.done():
self._fanout_task = asyncio.create_task(self._fanout())
# Public alias.
ensure_fanout_running = _ensure_fanout_running
async def _fanout(self) -> None:
"""Single consumer of the shared SSE; routes events to subscriptions.
@@ -211,6 +230,10 @@ class StreamController:
Re-read `self._shared_stream` on each outer iteration so we always
consume from the current handle. The old handle's iterator exhausts
naturally after `_close_after` closes it.
On a post-ready transport drop (non-cancelled error in `shared.done`),
attempts to reconnect up to `_max_reconnect_attempts` times before
giving up and closing subscriber queues.
"""
from langgraph_sdk.stream.subscription import matches_subscription
@@ -225,14 +248,21 @@ class StreamController:
for sub in list(self._subscriptions.values()):
if matches_subscription(event, sub.params):
sub.queue.put_nowait(event)
except Exception:
# Pump errored — close all subscription queues so consumers
# don't hang.
for sub in self._subscriptions.values():
sub.queue.put_nowait(None)
raise
except Exception as drop_err:
_logger.debug("transport drop in fanout: %r", drop_err)
if self._shared_stream is shared:
# No rotation happened; stream genuinely ended.
err = await shared.done
if (
err is not None
and not isinstance(err, asyncio.CancelledError)
and not self._closed
):
with contextlib.suppress(Exception):
await self._shared_stream.close()
reconnected = await self._reconnect_shared_stream()
if reconnected:
continue
break
# Rotation: loop again to pick up the new _shared_stream.
@@ -240,6 +270,45 @@ class StreamController:
for sub in self._subscriptions.values():
sub.queue.put_nowait(None)
async def _reconnect_sleep(self, attempt: int) -> None:
"""Sleep with exponential backoff and jitter for reconnect attempt *attempt*."""
base = self._reconnect_backoff_base
cap = self._reconnect_backoff_cap
delay = min(cap, base * (2**attempt))
jitter = random.uniform(0, delay * 0.25)
await asyncio.sleep(delay + jitter)
async def _reconnect_shared_stream(self) -> bool:
"""Attempt to reopen the shared stream after a transport drop.
Returns True if a new stream was successfully opened, False if all
reconnect attempts were exhausted or the controller was closed.
"""
# We intentionally use the *current* shared_stream_filter (the latest
# computed union of all live subscriptions), not the filter that was
# active when this stream was originally opened. If subscriptions were
# added or removed during the drop window, the reconnect picks up the
# new shape.
base_filter = self._shared_stream_filter
if base_filter is None:
return False
for attempt in range(self._max_reconnect_attempts):
if self._closed:
return False
try:
new_stream = self._transport.open_event_stream(
self._filter_with_since(base_filter)
)
await new_stream.ready
except asyncio.CancelledError:
raise
except Exception:
await self._reconnect_sleep(attempt)
continue
self._shared_stream = new_stream
return True
return False
# ------------------------------------------------------------------
# Stream rotation
# ------------------------------------------------------------------
@@ -262,7 +331,9 @@ class StreamController:
return # Existing stream is sufficient.
new_filter = self._compute_current_union(extra=candidate_filter)
new_stream = self._transport.open_event_stream(new_filter)
new_stream = self._transport.open_event_stream(
self._filter_with_since(new_filter)
)
old_stream = self._shared_stream
self._shared_stream = new_stream
self._shared_stream_filter = new_filter
@@ -272,6 +343,10 @@ class StreamController:
self._rotation_close_tasks.add(task)
task.add_done_callback(self._rotation_close_tasks.discard)
async def reconcile_stream(self, candidate_filter: SubscribeParams) -> None:
"""Public alias for `_reconcile_stream`."""
return await self._reconcile_stream(candidate_filter)
def _compute_current_union(
self, extra: SubscribeParams | None = None
) -> dict[str, Any]:
@@ -285,7 +360,28 @@ class StreamController:
return compute_union_filter(filters)
# ------------------------------------------------------------------
# Dedup
# Cursor tracking
# ------------------------------------------------------------------
def observe_applied_through_seq(self, seq: Any) -> None:
"""Advance the reconnect cursor from a command response meta sequence."""
self._observe_seq(seq)
def _observe_event(self, event: Event) -> None:
self._observe_seq(event.get("seq"))
def _observe_seq(self, seq: Any) -> None:
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]:
out = dict(params)
if self._cursor is not None:
out["since"] = self._cursor
return out
# ------------------------------------------------------------------
# Dedup iterator
# ------------------------------------------------------------------
async def _dedup_iter(self, source: AsyncIterator[Event]) -> AsyncIterator[Event]:
@@ -295,4 +391,5 @@ class StreamController:
if event_id in self._seen_event_ids:
continue
self._seen_event_ids.add(event_id)
self._observe_event(event)
yield event
@@ -45,8 +45,8 @@ class EventStreamHandle:
stream closes (server hangup or `close()`).
ready: resolves once HTTP response headers arrive; rejects on
connection failure before headers.
done: resolves with `None` on clean end or cancellation, or with
the exception on a mid-stream transport error.
done: resolves to `None` on clean end, or the post-ready stream
exception that ended the pump.
close: invoke to cancel the underlying task and free the
connection.
"""
@@ -174,16 +174,15 @@ class ProtocolSseTransport:
part = sse_decoder.decode(b"")
if part is not None and isinstance(part.data, dict):
await queue.put(cast("Event", part.data))
except asyncio.CancelledError:
except asyncio.CancelledError as err:
if not done.done():
done.set_result(None)
done.set_result(err)
raise
except BaseException as err:
if not ready.done():
ready.set_exception(err)
if not done.done():
done.set_result(err)
# Do not re-raise; the error is surfaced via `done`.
finally:
if not done.done():
done.set_result(None)
@@ -0,0 +1,138 @@
"""Synchronous HTTP/SSE transport for the v3 thread-centric protocol."""
from __future__ import annotations
import contextlib
from collections.abc import Callable, Iterator, Mapping
from dataclasses import dataclass
from typing import Any, cast
import httpx
import orjson
from langchain_protocol import Event
from langgraph_sdk.sse import BytesLineDecoder, SSEDecoder
from langgraph_sdk.stream.transport.http import _build_event_stream_body
@dataclass
class SyncEventStreamHandle:
"""Handle for one filtered synchronous SSE stream."""
events: Iterator[Event]
error: Callable[[], BaseException | None]
close: Callable[[], None]
class SyncProtocolSseTransport:
"""Sync v3 protocol transport bound to one thread id."""
def __init__(
self,
*,
client: httpx.Client,
thread_id: str,
commands_path: str | None = None,
stream_path: str | None = None,
headers: Mapping[str, str] | None = None,
) -> None:
self._client = client
self.thread_id = thread_id
self._commands_url = commands_path or f"/threads/{thread_id}/commands"
self._stream_url = stream_path or f"/threads/{thread_id}/stream/events"
self._default_headers: dict[str, str] = dict(headers or {})
self._closed = False
self._open_responses: list[httpx.Response] = []
def send_command(self, command: dict[str, Any]) -> dict[str, Any] | None:
if self._closed:
raise RuntimeError("Protocol transport is closed.")
merged_headers = {**self._default_headers, "content-type": "application/json"}
response = self._client.post(
self._commands_url,
content=orjson.dumps(command),
headers=merged_headers,
)
response.raise_for_status()
if response.status_code in (202, 204):
return None
payload = orjson.loads(response.content)
if not isinstance(payload, dict) or "id" not in payload:
raise RuntimeError("Protocol command did not return a valid response.")
return payload
def open_event_stream(self, params: dict[str, Any]) -> SyncEventStreamHandle:
if self._closed:
raise RuntimeError("Protocol transport is closed.")
sse_headers = {
**self._default_headers,
"content-type": "application/json",
"accept": "text/event-stream",
"cache-control": "no-store",
}
request = self._client.build_request(
"POST",
self._stream_url,
content=orjson.dumps(_build_event_stream_body(params)),
headers=sse_headers,
)
stream_cm = self._client.send(request, stream=True)
stream_cm.raise_for_status()
content_type = stream_cm.headers.get("content-type", "").partition(";")[0]
if "text/event-stream" not in content_type:
stream_cm.close()
raise httpx.TransportError(
"Expected response header Content-Type to contain "
f"'text/event-stream', got {content_type!r}"
)
self._open_responses.append(stream_cm)
closed = False
stream_error: BaseException | None = None
def events() -> Iterator[Event]:
nonlocal stream_error
line_decoder = BytesLineDecoder()
sse_decoder = SSEDecoder()
try:
for chunk in stream_cm.iter_bytes():
if closed:
return
for line in line_decoder.decode(chunk):
part = sse_decoder.decode(bytes(line))
if part is not None and isinstance(part.data, dict):
yield cast("Event", part.data)
for line in line_decoder.flush():
part = sse_decoder.decode(bytes(line))
if part is not None and isinstance(part.data, dict):
yield cast("Event", part.data)
part = sse_decoder.decode(b"")
if part is not None and isinstance(part.data, dict):
yield cast("Event", part.data)
except BaseException as exc:
if not closed:
stream_error = exc
raise
finally:
with contextlib.suppress(ValueError):
self._open_responses.remove(stream_cm)
stream_cm.close()
def error() -> BaseException | None:
return stream_error
def close() -> None:
nonlocal closed
closed = True
with contextlib.suppress(Exception):
stream_cm.close()
return SyncEventStreamHandle(events=events(), error=error, close=close)
def close(self) -> None:
if self._closed:
return
self._closed = True
for response in list(self._open_responses):
with contextlib.suppress(Exception):
response.close()
self._open_responses.clear()
+46 -8
View File
@@ -12,6 +12,7 @@ from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from dataclasses import dataclass
from typing import Any
import orjson
@@ -21,6 +22,13 @@ from starlette.responses import JSONResponse, Response, StreamingResponse
from starlette.routing import Route
@dataclass
class _StreamScript:
events: list[dict[str, Any]]
delay: float = 0.0
fail_after: int | None = None
class FakeServer:
"""Holds scripted state for tests and exposes a Starlette app.
@@ -48,11 +56,31 @@ class FakeServer:
self.state: dict[str, Any] = {}
self.state_request_count: int = 0
self.state_request_headers: list[dict[str, str]] = []
self._stream_scripts: list[_StreamScript] = []
self._command_response: dict[str, Any] | None = None
def script(self, events: list[dict[str, Any]], *, delay: float = 0.0) -> None:
"""Set the events the next /stream/events call will replay."""
def script(
self,
events: list[dict[str, Any]],
*,
delay: float = 0.0,
fail_after: int | None = None,
) -> None:
"""Set the events the next /stream/events calls will replay."""
self.scripted_events = list(events)
self._stream_delay = delay
self._stream_scripts = [
_StreamScript(events=list(events), delay=delay, fail_after=fail_after)
]
def script_sequence(self, scripts: list[_StreamScript]) -> None:
"""Set per-open stream scripts consumed in order by /stream/events."""
self._stream_scripts = list(scripts)
self.scripted_events = []
def script_command_response(self, response: dict[str, Any]) -> None:
"""Set the command envelope returned by /commands."""
self._command_response = dict(response)
def set_state(
self,
@@ -82,6 +110,10 @@ class FakeServer:
self.received_commands.append(body)
self.command_request_headers.append(dict(request.headers))
command_id = body.get("id")
if self._command_response is not None:
response = dict(self._command_response)
response["id"] = command_id
return JSONResponse(response)
return JSONResponse(
{
"type": "success",
@@ -124,16 +156,22 @@ class FakeServer:
self._open_event_streams_max = max(
self._open_event_streams_max, self.open_event_streams
)
script = (
self._stream_scripts.pop(0)
if self._stream_scripts
else _StreamScript(
events=list(self.scripted_events), delay=self._stream_delay
)
)
try:
# Why: script() rebinds scripted_events; in-flight iterators retain
# a reference to the prior list and are unaffected by later
# script() calls.
for event in self.scripted_events:
if self._stream_delay:
await asyncio.sleep(self._stream_delay)
for index, event in enumerate(script.events, start=1):
if script.delay:
await asyncio.sleep(script.delay)
payload = orjson.dumps(event).decode()
yield f"id: {event.get('event_id', '')}\n".encode()
yield f"event: message\ndata: {payload}\n\n".encode()
if script.fail_after is not None and index >= script.fail_after:
raise RuntimeError("scripted async stream failure")
finally:
self.open_event_streams -= 1
@@ -0,0 +1,116 @@
"""Synchronous fake v3 protocol server for sync streaming tests."""
from __future__ import annotations
from collections.abc import Iterator
from dataclasses import dataclass
from typing import Any
import httpx
import orjson
@dataclass
class SyncStreamScript:
events: list[dict[str, Any]]
fail_after: int | None = None
class _SseByteStream(httpx.SyncByteStream):
def __init__(self, script: SyncStreamScript) -> None:
self._script = script
def __iter__(self) -> Iterator[bytes]:
for index, event in enumerate(self._script.events, start=1):
payload = orjson.dumps(event).decode()
yield f"id: {event.get('event_id', '')}\n".encode()
yield f"event: message\ndata: {payload}\n\n".encode()
if self._script.fail_after is not None and index >= self._script.fail_after:
raise httpx.ReadError("scripted sync stream failure")
class SyncFakeServer:
"""Synchronous fake for `/commands`, `/stream/events`, and `/state`."""
def __init__(self) -> None:
self.received_commands: list[dict[str, Any]] = []
self.stream_request_bodies: list[dict[str, Any]] = []
self.command_request_headers: list[dict[str, str]] = []
self.stream_request_headers_list: list[dict[str, str]] = []
self.state_request_headers: list[dict[str, str]] = []
self.state_request_count = 0
self.scripted_events: list[dict[str, Any]] = []
self.state: dict[str, Any] = {}
self.transport = httpx.MockTransport(self._handle)
self._stream_scripts: list[SyncStreamScript] = []
self._command_response: dict[str, Any] | None = None
def script(
self,
events: list[dict[str, Any]],
*,
fail_after: int | None = None,
) -> None:
self.scripted_events = list(events)
self._stream_scripts = [
SyncStreamScript(events=list(events), fail_after=fail_after)
]
def script_sequence(self, scripts: list[SyncStreamScript]) -> None:
self._stream_scripts = list(scripts)
self.scripted_events = []
def script_command_response(self, response: dict[str, Any]) -> None:
self._command_response = dict(response)
def set_state(
self,
values: dict[str, Any],
next: list[Any] | None = None,
metadata: dict[str, Any] | None = None,
) -> None:
self.state = {
"values": values,
"next": next if next is not None else [],
"tasks": [],
"metadata": metadata if metadata is not None else {},
"checkpoint": None,
"created_at": None,
}
def _handle(self, request: httpx.Request) -> httpx.Response:
path = request.url.path
if path.endswith("/commands"):
body = orjson.loads(request.content)
self.received_commands.append(body)
self.command_request_headers.append(dict(request.headers))
if self._command_response is not None:
response = dict(self._command_response)
response["id"] = body.get("id")
return httpx.Response(200, json=response)
return httpx.Response(
200,
json={
"type": "success",
"id": body.get("id"),
"result": {"run_id": "run-1"},
},
)
if path.endswith("/stream/events"):
self.stream_request_bodies.append(orjson.loads(request.content))
self.stream_request_headers_list.append(dict(request.headers))
script = (
self._stream_scripts.pop(0)
if self._stream_scripts
else SyncStreamScript(events=list(self.scripted_events))
)
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=_SseByteStream(script),
)
if path.endswith("/state"):
self.state_request_count += 1
self.state_request_headers.append(dict(request.headers))
return httpx.Response(200, json=self.state)
return httpx.Response(404, json={"error": f"unexpected path: {path}"})
@@ -2,9 +2,15 @@
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from typing import Any
from unittest.mock import AsyncMock
import pytest
from langgraph_sdk.stream.controller import StreamController, _SeenEventIds
from langgraph_sdk.stream.transport.http import EventStreamHandle
# ---------------------------------------------------------------------------
# Task 3.1: bounded subscription queues
@@ -152,3 +158,196 @@ async def test_close_awaits_pending_rotation_closes():
# close() must block until the rotation close completes.
await controller.close()
assert rotation_close_done.is_set()
# ---------------------------------------------------------------------------
# Reconnect helpers
# ---------------------------------------------------------------------------
def _make_handle(
*,
error: BaseException | None = None,
) -> EventStreamHandle:
"""Build a minimal EventStreamHandle that closes immediately."""
loop = asyncio.get_running_loop()
ready: asyncio.Future[None] = loop.create_future()
ready.set_result(None)
done: asyncio.Future[BaseException | None] = loop.create_future()
done.set_result(error)
async def _aiter() -> AsyncIterator[Any]:
if False:
yield # pragma: no cover
return EventStreamHandle(
events=_aiter(),
ready=ready,
done=done,
close=AsyncMock(),
)
def _always_error_transport(error_type: type[Exception] = RuntimeError) -> Any:
"""Return a fake transport whose open_event_stream always raises."""
class _Transport:
def open_event_stream(self, _params: dict[str, Any]) -> EventStreamHandle:
raise error_type("scripted transport error")
return _Transport()
def _error_then_succeed_transport(fail_count: int) -> Any:
"""Return a fake transport that fails *fail_count* times then succeeds."""
calls = [0]
class _Transport:
def open_event_stream(self, _params: dict[str, Any]) -> EventStreamHandle:
calls[0] += 1
if calls[0] <= fail_count:
raise RuntimeError(f"scripted error #{calls[0]}")
return _make_handle()
return _Transport()
# ---------------------------------------------------------------------------
# Task 8.1: Exp+jitter backoff
# ---------------------------------------------------------------------------
@pytest.mark.anyio
async def test_reconnect_uses_exp_backoff_with_jitter(monkeypatch):
"""Reconnect attempts should sleep increasing durations with jitter,
not a fixed 50ms."""
sleeps: list[float] = []
async def fake_sleep(d: float) -> None:
sleeps.append(d)
monkeypatch.setattr("asyncio.sleep", fake_sleep)
transport = _always_error_transport()
controller = StreamController(
transport=transport,
run_start_gate=AsyncMock(),
max_reconnect_attempts=3,
reconnect_backoff_base=0.1,
reconnect_backoff_cap=2.0,
)
# Seed filter so reconnect doesn't bail early.
controller._shared_stream_filter = {"channels": ["lifecycle"]}
await controller._reconnect_shared_stream()
# Should have slept once per attempt.
assert len(sleeps) == 3
# All sleeps within [base, cap + 25% jitter].
assert all(0.1 <= s <= 2.5 for s in sleeps)
@pytest.mark.anyio
async def test_reconnect_accepts_backoff_kwargs():
"""StreamController must accept reconnect_backoff_base and _cap kwargs."""
controller = StreamController(
transport=_always_error_transport(),
run_start_gate=AsyncMock(),
max_reconnect_attempts=1,
reconnect_backoff_base=0.05,
reconnect_backoff_cap=1.0,
)
assert controller._reconnect_backoff_base == 0.05
assert controller._reconnect_backoff_cap == 1.0
# ---------------------------------------------------------------------------
# Task 8.2: Close old handle before reconnect
# ---------------------------------------------------------------------------
@pytest.mark.anyio
async def test_transport_drop_exception_logged_with_type(monkeypatch, caplog):
"""Bare `pass` discarded exception types; the drop should at least log."""
import logging
monkeypatch.setattr("asyncio.sleep", AsyncMock())
loop = asyncio.get_running_loop()
ready: asyncio.Future[None] = loop.create_future()
ready.set_result(None)
done: asyncio.Future[BaseException | None] = loop.create_future()
done.set_result(RuntimeError("transport drop"))
async def _raises() -> AsyncIterator[Any]:
raise RuntimeError("transport drop")
yield # pragma: no cover
old_handle = EventStreamHandle(
events=_raises(),
ready=ready,
done=done,
close=AsyncMock(),
)
transport = _always_error_transport()
controller = StreamController(
transport=transport,
run_start_gate=AsyncMock(),
max_reconnect_attempts=1,
reconnect_backoff_base=0.0,
reconnect_backoff_cap=0.0,
)
controller._shared_stream = old_handle
controller._shared_stream_filter = {"channels": ["lifecycle"]}
with caplog.at_level(logging.DEBUG, logger="langgraph_sdk.stream.controller"):
await controller._fanout()
assert any("transport drop" in rec.message for rec in caplog.records)
@pytest.mark.anyio
async def test_reconnect_closes_old_handle_before_opening_new(monkeypatch):
"""When the shared stream errors and triggers reconnect, the old
EventStreamHandle's close() must be called."""
# Suppress actual sleeps.
monkeypatch.setattr("asyncio.sleep", AsyncMock())
close_calls: list[str] = []
loop = asyncio.get_running_loop()
old_ready: asyncio.Future[None] = loop.create_future()
old_ready.set_result(None)
# done resolves with an error to trigger the reconnect path in _fanout.
old_done: asyncio.Future[BaseException | None] = loop.create_future()
old_done.set_result(RuntimeError("transport drop"))
async def _empty() -> AsyncIterator[Any]:
# Raise on the first iteration so _fanout exits the inner loop.
raise RuntimeError("transport drop")
yield # pragma: no cover
old_handle = EventStreamHandle(
events=_empty(),
ready=old_ready,
done=old_done,
close=AsyncMock(side_effect=lambda: close_calls.append("old_closed")),
)
# Transport always errors so reconnect exhausts all attempts and _fanout exits.
transport = _always_error_transport()
controller = StreamController(
transport=transport,
run_start_gate=AsyncMock(),
max_reconnect_attempts=1,
reconnect_backoff_base=0.0,
reconnect_backoff_cap=0.0,
)
controller._shared_stream = old_handle
controller._shared_stream_filter = {"channels": ["lifecycle"]}
# _fanout drives reconnect; wait for it to complete.
await controller._fanout()
assert "old_closed" in close_calls
@@ -9,6 +9,7 @@ import httpx
from langgraph_sdk._async.http import HttpClient
from langgraph_sdk._async.threads import ThreadsClient
from langgraph_sdk.stream.controller import StreamController
from langgraph_sdk.stream.transport.http import EventStreamHandle
from streaming._events import lifecycle_event, values_event
from streaming._fake_server import FakeServer
@@ -201,3 +202,128 @@ async def test_values_projection_registers_via_delegation_not_controller_directl
thread_count, ctrl_count = counts_during[0]
assert thread_count == ctrl_count
assert thread_count >= 1
def _make_handle(
events: list[dict[str, Any]],
err: BaseException | None = None,
) -> tuple[EventStreamHandle, asyncio.Queue]:
"""Build a synthetic EventStreamHandle for reconnect tests.
Returns the handle and the underlying queue so callers can inject events
or the end sentinel directly from test code.
"""
loop = asyncio.get_running_loop()
ready: asyncio.Future[None] = loop.create_future()
done: asyncio.Future[BaseException | None] = loop.create_future()
queue: asyncio.Queue = asyncio.Queue()
async def _pump() -> None:
ready.set_result(None)
for event in events:
await queue.put(event)
done.set_result(err)
await queue.put(None) # sentinel
asyncio.create_task(_pump()) # noqa: RUF006
async def _aiter():
while True:
item = await queue.get()
if item is None:
return
yield item
async def _close() -> None:
if not done.done():
done.set_result(None)
await queue.put(None)
return EventStreamHandle(
events=_aiter(), ready=ready, done=done, close=_close
), queue
async def test_shared_stream_reconnects_with_since_after_transport_drop():
"""StreamController reopens the stream with `since` after a post-ready drop."""
opened_params: list[dict[str, Any]] = []
handle1, _ = _make_handle(
[values_event(seq=1, values={"counter": 1})],
err=RuntimeError("scripted async stream failure"),
)
handle2, _ = _make_handle([values_event(seq=2, values={"counter": 2})])
handles = [handle1, handle2]
from unittest.mock import MagicMock
from langgraph_sdk.stream.transport.http import ProtocolSseTransport
transport = MagicMock(spec=ProtocolSseTransport)
def _open(params: dict[str, Any]) -> EventStreamHandle:
opened_params.append(dict(params))
return handles.pop(0)
transport.open_event_stream.side_effect = _open
async def gate() -> None:
return None
controller = StreamController(transport=transport, run_start_gate=gate)
sub = controller.register_subscription({"channels": ["values"]})
await controller.reconcile_stream({"channels": ["values"]})
controller.ensure_fanout_running()
first = await asyncio.wait_for(sub.queue.get(), timeout=1.0)
second = await asyncio.wait_for(sub.queue.get(), timeout=1.0)
end = await asyncio.wait_for(sub.queue.get(), timeout=1.0)
await controller.close()
assert first["seq"] == 1
assert second["seq"] == 2
assert end is None
assert opened_params[0]["channels"] == ["values"]
assert "since" not in opened_params[0]
assert opened_params[1]["channels"] == ["values"]
assert opened_params[1]["since"] == 1
async def test_shared_stream_reconnect_dedupes_replayed_overlap():
"""StreamController deduplicates events replayed on reconnect."""
handle1, _ = _make_handle(
[values_event(seq=1, values={"counter": 1})],
err=RuntimeError("scripted async stream failure"),
)
handle2, _ = _make_handle(
[
values_event(seq=1, values={"counter": 1}), # replayed overlap
values_event(seq=2, values={"counter": 2}),
]
)
handles = [handle1, handle2]
from unittest.mock import MagicMock
from langgraph_sdk.stream.transport.http import ProtocolSseTransport
transport = MagicMock(spec=ProtocolSseTransport)
transport.open_event_stream.side_effect = lambda _params: handles.pop(0)
async def gate() -> None:
return None
controller = StreamController(transport=transport, run_start_gate=gate)
sub = controller.register_subscription({"channels": ["values"]})
await controller.reconcile_stream({"channels": ["values"]})
controller.ensure_fanout_running()
received = [
await asyncio.wait_for(sub.queue.get(), timeout=1.0),
await asyncio.wait_for(sub.queue.get(), timeout=1.0),
await asyncio.wait_for(sub.queue.get(), timeout=1.0),
]
await controller.close()
assert [event["seq"] for event in received if event is not None] == [1, 2]
assert received[-1] is None
@@ -0,0 +1,48 @@
"""Sync HTTP/SSE transport tests."""
from __future__ import annotations
import httpx
import pytest
from langgraph_sdk.stream.transport.sync_http import SyncProtocolSseTransport
from streaming._events import values_event
from streaming._sync_fake_server import SyncFakeServer
def test_sync_transport_sends_command():
fake = SyncFakeServer()
with httpx.Client(transport=fake.transport, base_url="http://test") as raw:
transport = SyncProtocolSseTransport(client=raw, thread_id="t-1")
result = transport.send_command(
{"id": 1, "method": "run.start", "params": {"assistant_id": "agent"}}
)
assert result == {"type": "success", "id": 1, "result": {"run_id": "run-1"}}
assert fake.received_commands[0]["method"] == "run.start"
def test_sync_transport_streams_events():
fake = SyncFakeServer()
fake.script([values_event(seq=1, counter=1)])
with httpx.Client(transport=fake.transport, base_url="http://test") as raw:
transport = SyncProtocolSseTransport(client=raw, thread_id="t-1")
handle = transport.open_event_stream({"channels": ["values"]})
events = list(handle.events)
assert events == [values_event(seq=1, counter=1)]
assert fake.stream_request_bodies == [{"channels": ["values"]}]
def test_sync_open_event_stream_records_post_ready_error():
fake = SyncFakeServer()
fake.script([values_event(seq=1)], fail_after=1)
with httpx.Client(transport=fake.transport, base_url="http://test") as raw:
sse = SyncProtocolSseTransport(client=raw, thread_id="t-1")
handle = sse.open_event_stream({"channels": ["values"]})
with pytest.raises(httpx.ReadError, match="scripted sync stream failure"):
list(handle.events)
err = handle.error()
handle.close()
assert isinstance(err, httpx.ReadError)
@@ -14,6 +14,7 @@ async def test_event_stream_handle_constructs_with_open_state():
loop = asyncio.get_running_loop()
ready: asyncio.Future[None] = loop.create_future()
done: asyncio.Future[BaseException | None] = loop.create_future()
done.set_result(None)
closed = False
async def aiter_events():
@@ -533,3 +534,38 @@ def test_values_event_builder_shape():
assert evt["method"] == "values"
assert evt["params"]["data"] == {"values": {"foo": 1}}
assert evt["params"]["namespace"] == []
async def test_open_event_stream_done_records_post_ready_error():
from streaming._events import values_event
event_data = values_event(seq=1)
class _FailAfterOneStream(httpx.AsyncByteStream):
async def __aiter__(self):
import orjson
payload = orjson.dumps(event_data).decode()
yield f"id: {event_data.get('event_id', '')}\n".encode()
yield f"event: message\ndata: {payload}\n\n".encode()
raise RuntimeError("scripted async stream failure")
async def handler(_request: httpx.Request) -> httpx.Response:
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=_FailAfterOneStream(),
)
transport = httpx.MockTransport(handler)
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
sse = ProtocolSseTransport(client=client, thread_id="t-1")
handle = sse.open_event_stream({"channels": ["values"]})
await asyncio.wait_for(handle.ready, timeout=1.0)
received = [event async for event in handle.events]
err = await asyncio.wait_for(handle.done, timeout=1.0)
await handle.close()
assert received == [values_event(seq=1)]
assert isinstance(err, RuntimeError)
assert "scripted async stream failure" in str(err)