mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
feat(sdk-py): add async stream reconnect support (#7825)
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user