From 10b701cf4176e913454eea785098d67a18ba5cd0 Mon Sep 17 00:00:00 2001 From: Nick Hollon Date: Wed, 27 May 2026 13:30:01 -0400 Subject: [PATCH] feat(sdk-py): add async stream reconnect support (#7825) --- libs/sdk-py/langgraph_sdk/_async/stream.py | 44 ++++ .../sdk-py/langgraph_sdk/stream/controller.py | 117 +++++++++- .../langgraph_sdk/stream/transport/http.py | 9 +- .../stream/transport/sync_http.py | 138 ++++++++++++ libs/sdk-py/tests/streaming/_fake_server.py | 54 ++++- .../tests/streaming/_sync_fake_server.py | 116 ++++++++++ .../sdk-py/tests/streaming/test_controller.py | 199 ++++++++++++++++++ .../tests/streaming/test_shared_stream.py | 126 +++++++++++ .../streaming/test_sync_transport_http.py | 48 +++++ .../tests/streaming/test_transport_http.py | 36 ++++ 10 files changed, 864 insertions(+), 23 deletions(-) create mode 100644 libs/sdk-py/langgraph_sdk/stream/transport/sync_http.py create mode 100644 libs/sdk-py/tests/streaming/_sync_fake_server.py create mode 100644 libs/sdk-py/tests/streaming/test_sync_transport_http.py diff --git a/libs/sdk-py/langgraph_sdk/_async/stream.py b/libs/sdk-py/langgraph_sdk/_async/stream.py index e653fea06..abba4218f 100644 --- a/libs/sdk-py/langgraph_sdk/_async/stream.py +++ b/libs/sdk-py/langgraph_sdk/_async/stream.py @@ -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() diff --git a/libs/sdk-py/langgraph_sdk/stream/controller.py b/libs/sdk-py/langgraph_sdk/stream/controller.py index 186d4c6e8..107dbff89 100644 --- a/libs/sdk-py/langgraph_sdk/stream/controller.py +++ b/libs/sdk-py/langgraph_sdk/stream/controller.py @@ -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 diff --git a/libs/sdk-py/langgraph_sdk/stream/transport/http.py b/libs/sdk-py/langgraph_sdk/stream/transport/http.py index c29e388be..3cb6cde62 100644 --- a/libs/sdk-py/langgraph_sdk/stream/transport/http.py +++ b/libs/sdk-py/langgraph_sdk/stream/transport/http.py @@ -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) diff --git a/libs/sdk-py/langgraph_sdk/stream/transport/sync_http.py b/libs/sdk-py/langgraph_sdk/stream/transport/sync_http.py new file mode 100644 index 000000000..97a7f10af --- /dev/null +++ b/libs/sdk-py/langgraph_sdk/stream/transport/sync_http.py @@ -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() diff --git a/libs/sdk-py/tests/streaming/_fake_server.py b/libs/sdk-py/tests/streaming/_fake_server.py index d82ec2039..c602c5b17 100644 --- a/libs/sdk-py/tests/streaming/_fake_server.py +++ b/libs/sdk-py/tests/streaming/_fake_server.py @@ -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 diff --git a/libs/sdk-py/tests/streaming/_sync_fake_server.py b/libs/sdk-py/tests/streaming/_sync_fake_server.py new file mode 100644 index 000000000..eb1715c8d --- /dev/null +++ b/libs/sdk-py/tests/streaming/_sync_fake_server.py @@ -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}"}) diff --git a/libs/sdk-py/tests/streaming/test_controller.py b/libs/sdk-py/tests/streaming/test_controller.py index 35e8c3c84..e9c71670e 100644 --- a/libs/sdk-py/tests/streaming/test_controller.py +++ b/libs/sdk-py/tests/streaming/test_controller.py @@ -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 diff --git a/libs/sdk-py/tests/streaming/test_shared_stream.py b/libs/sdk-py/tests/streaming/test_shared_stream.py index 6eaeb1f30..9eb900e81 100644 --- a/libs/sdk-py/tests/streaming/test_shared_stream.py +++ b/libs/sdk-py/tests/streaming/test_shared_stream.py @@ -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 diff --git a/libs/sdk-py/tests/streaming/test_sync_transport_http.py b/libs/sdk-py/tests/streaming/test_sync_transport_http.py new file mode 100644 index 000000000..8b9536b72 --- /dev/null +++ b/libs/sdk-py/tests/streaming/test_sync_transport_http.py @@ -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) diff --git a/libs/sdk-py/tests/streaming/test_transport_http.py b/libs/sdk-py/tests/streaming/test_transport_http.py index 151b649ed..597136bc8 100644 --- a/libs/sdk-py/tests/streaming/test_transport_http.py +++ b/libs/sdk-py/tests/streaming/test_transport_http.py @@ -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)