From 221deee7743828703974fe5f927dcaca687fa9ea Mon Sep 17 00:00:00 2001 From: Nick Hollon Date: Wed, 27 May 2026 11:06:37 -0400 Subject: [PATCH] feat(sdk-py): wire lifecycle state and output prerequisites (#7821) --- libs/sdk-py/langgraph_sdk/_async/stream.py | 251 ++++++++++-- libs/sdk-py/langgraph_sdk/_async/threads.py | 14 +- .../langgraph_sdk/stream/transport/http.py | 21 +- libs/sdk-py/tests/streaming/_events.py | 25 ++ libs/sdk-py/tests/streaming/_fake_server.py | 39 ++ .../tests/streaming/test_lifecycle_watcher.py | 184 ++++++++- .../tests/streaming/test_thread_stream.py | 364 +++++++++++++++++- .../tests/streaming/test_transport_http.py | 99 +++++ 8 files changed, 943 insertions(+), 54 deletions(-) diff --git a/libs/sdk-py/langgraph_sdk/_async/stream.py b/libs/sdk-py/langgraph_sdk/_async/stream.py index 3a41e4d36..5c0fa529d 100644 --- a/libs/sdk-py/langgraph_sdk/_async/stream.py +++ b/libs/sdk-py/langgraph_sdk/_async/stream.py @@ -1,8 +1,11 @@ """Async thread-centric streaming surface for the v3 protocol. `AsyncThreadStream` is an async context manager that owns a -`ProtocolSseTransport` for one thread, dispatches `run.start` commands, -and exposes a raw `events` async iterable. +`ProtocolSseTransport` for one thread, dispatches commands (`run.start`, +`run.respond`), exposes typed subscriptions over a single shared SSE +(`subscribe`, `events`), and surfaces lifecycle state (`interrupted`, +`interrupts`) via an always-on lifecycle watcher SSE. Typed projections +(`thread.values`, `thread.messages`, etc.) mirror the v3 protocol surface. Direct port of `libs/sdk/src/client/stream/index.ts`. """ @@ -10,13 +13,14 @@ Direct port of `libs/sdk/src/client/stream/index.ts`. from __future__ import annotations import asyncio -from collections.abc import AsyncGenerator, AsyncIterator +import contextlib +from collections.abc import AsyncGenerator, AsyncIterator, Mapping from dataclasses import dataclass, field -from typing import Any, TypedDict +from typing import Any, Literal, TypedDict -import httpx from langchain_protocol import Event, SubscribeParams +from langgraph_sdk._async.http import HttpClient from langgraph_sdk.stream.transport import EventStreamHandle, ProtocolSseTransport @@ -28,6 +32,14 @@ class InterruptPayload(TypedDict): namespace: list[str] +@dataclass +class _RunTerminal: + """Terminal state record resolved into `_run_done` on lifecycle completion.""" + + status: Literal["completed", "errored"] + error: BaseException | None = None + + @dataclass class _Subscription: """Internal record for one active subscription on an `AsyncThreadStream`.""" @@ -77,8 +89,89 @@ class RunModule: params["config"] = config if metadata is not None: params["metadata"] = metadata - self._owner._ensure_lifecycle_watcher_running() - return await self._owner._send_command("run.start", params) + loop = asyncio.get_running_loop() + gate: asyncio.Future[None] = loop.create_future() + self._owner._run_start_ready = gate + try: + result = await self._owner._send_command("run.start", params) + if not gate.done(): + gate.set_result(None) + self._owner._run_seen = True + return result + except BaseException as err: + # Why: gate MUST reject on any exit type, including CancelledError, + # so awaiters see the failure rather than hanging indefinitely. + if not gate.done(): + gate.set_exception(err) + raise + finally: + # Why: concurrent run.start calls (multitask_strategy="enqueue") + # can replace _run_start_ready before our finally fires. + # Identity-check before clearing so the later call's gate isn't stomped. + if self._owner._run_start_ready is gate: + self._owner._run_start_ready = None + # Why: if the gate stored an exception that no awaiter consumed, + # retrieve it here to suppress asyncio's GC warning. The exception + # is already propagated to our caller via the `raise` above. + if gate.done() and not gate.cancelled(): + gate.exception() + + async def respond( + self, + response: Any, + *, + interrupt_id: str | None = None, + ) -> dict[str, Any]: + """Reply to a server-side interrupt and resume the run. + + Args: + response: the response value forwarded as `params.response` on the + wire (protocol field name). + interrupt_id: optional explicit id. When omitted, requires exactly + one outstanding interrupt and uses its id. + + Raises: + RuntimeError: no outstanding interrupts; `interrupt_id` is None but + multiple interrupts are outstanding; or the explicit + `interrupt_id` doesn't match any outstanding interrupt. + """ + # Why: take the `interrupts` snapshot AND dispatch the command under + # `_interrupts_lock`, so the lifecycle watcher's terminal-clear path + # cannot wipe `interrupts` between the snapshot and the wire send. + async with self._owner._interrupts_lock: + outstanding = list(self._owner.interrupts) + if interrupt_id is None: + if len(outstanding) == 0: + raise RuntimeError( + "thread.run.respond: no outstanding interrupt. Provide " + "an explicit `interrupt_id` or wait for " + "`thread.interrupted`." + ) + if len(outstanding) > 1: + ids = [p["interrupt_id"] for p in outstanding] + raise RuntimeError( + f"thread.run.respond: ambiguous — {len(outstanding)} " + f"outstanding interrupts ({ids!r}). Provide an explicit " + "`interrupt_id`." + ) + match = outstanding[0] + else: + match = next( + (p for p in outstanding if p["interrupt_id"] == interrupt_id), + None, + ) + if match is None: + raise RuntimeError( + f"thread.run.respond: interrupt_id {interrupt_id!r} does " + "not match any outstanding interrupt in " + "`thread.interrupts`." + ) + params = { + "interrupt_id": match["interrupt_id"], + "namespace": match["namespace"], + "response": response, + } + return await self._owner._send_command("input.respond", params) async def _close_after(handle: EventStreamHandle, *, delay: float = 0.0) -> None: @@ -101,15 +194,19 @@ class AsyncThreadStream: def __init__( self, *, - client: httpx.AsyncClient, + http: HttpClient, thread_id: str, assistant_id: str, + headers: Mapping[str, str] | None = None, max_queue_size: int = 1024, + run_start_timeout: float | None = None, ) -> None: - self._http_client = client + self._http = http + self._headers = dict(headers or {}) self.thread_id = thread_id self.assistant_id = assistant_id self._max_queue_size = max_queue_size + self._run_start_timeout = run_start_timeout self._closed = False self._transport: ProtocolSseTransport | None = None self._open_handles: list[EventStreamHandle] = [] @@ -122,18 +219,32 @@ class AsyncThreadStream: self._fanout_task: asyncio.Task[None] | None = None self.interrupted: bool = False self.interrupts: list[InterruptPayload] = [] + # Why: serialize the `interrupts` snapshot in `run.respond` with the + # terminal-clear path in `_apply_lifecycle_event`, so a respond() in + # flight cannot send a stale `interrupt_id` after the lifecycle watcher + # observes a `completed`/`errored` event. + self._interrupts_lock = asyncio.Lock() self._lifecycle_watcher_task: asyncio.Task[None] | None = None self._lifecycle_watcher_handle: EventStreamHandle | None = None + self._run_start_ready: asyncio.Future[None] | None = None + self._run_seen: bool = False + self._run_done: asyncio.Future[_RunTerminal] | None = None self.run = RunModule(self) async def __aenter__(self) -> AsyncThreadStream: if self._closed: raise RuntimeError("AsyncThreadStream is closed and cannot be re-entered.") self._transport = ProtocolSseTransport( - client=self._http_client, + client=self._http.client, thread_id=self.thread_id, + headers=self._headers, max_queue_size=self._max_queue_size, ) + # Create the run-done future here (async context guarantees a running loop). + self._run_done = asyncio.get_running_loop().create_future() + # Start the lifecycle watcher immediately so reattach and thread.output + # work without a preceding run.start call. + self._ensure_lifecycle_watcher_running() return self async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: @@ -167,6 +278,22 @@ class AsyncThreadStream: self._closed = True for handle in self._open_handles: await handle.close() + # Cancel _run_done so thread.output doesn't wait forever on close. + run_done = self._run_done + if run_done is not None and not run_done.done(): + run_done.cancel() + if self._lifecycle_watcher_task is not None: + self._lifecycle_watcher_task.cancel() + with contextlib.suppress(Exception, asyncio.CancelledError): + await self._lifecycle_watcher_task + if self._lifecycle_watcher_handle is not None: + await self._lifecycle_watcher_handle.close() + if self._fanout_task is not None: + self._fanout_task.cancel() + with contextlib.suppress(Exception, asyncio.CancelledError): + await self._fanout_task + if self._shared_stream is not None: + await self._shared_stream.close() if self._transport is not None: await self._transport.close() @@ -274,6 +401,7 @@ class AsyncThreadStream: that both old and new streams are simultaneously connected during rotation (enabling correct peak-count tracking and dedup correctness). """ + await self._await_run_start_gate(timeout=self._run_start_timeout) from langgraph_sdk.stream.subscription import filter_covers if self._transport is None: @@ -346,6 +474,21 @@ class AsyncThreadStream: raise RuntimeError(f"Protocol error [{code}]: {message}") return response.get("result", {}) + async def _await_run_start_gate(self, *, timeout: float | None = None) -> None: + """Wait for the current run.start to commit the thread server-side. + + No-op when no run.start is in flight. Re-raises if run.start failed. + Raises `asyncio.TimeoutError` if `timeout` is set and the gate does + not resolve in time; the gate itself is left intact for later callers. + """ + gate = self._run_start_ready + if gate is None or gate.done(): + return + if timeout is None: + await gate + else: + await asyncio.wait_for(asyncio.shield(gate), timeout=timeout) + def _ensure_lifecycle_watcher_running(self) -> None: # Why: this watcher is intentionally one-shot. If it crashes, it stays # dead until the AsyncThreadStream is closed. @@ -359,10 +502,9 @@ class AsyncThreadStream: """Always-on SSE consuming lifecycle + input channels. Independent of the union-filter shared stream so that interrupts - surface even when no other subscription is active. - - The watcher waits for the run-start gate before opening so it does not - race server-side thread creation. + surface even when no other subscription is active. Starts immediately + on session entry (before any run.start) so reattach and thread.output + work for existing runs. """ if self._transport is None: return @@ -375,15 +517,50 @@ class AsyncThreadStream: async for event in handle.events: if self._closed: return - self._apply_lifecycle_event(event) - except (Exception, asyncio.CancelledError): + await self._apply_lifecycle_event(event) + # Why: iterator exhausted without `_run_done` being resolved by a + # terminal lifecycle event. Surface any transport error captured + # on `handle.done`, otherwise treat the clean EOF as errored so + # awaiters of `_run_done` (e.g. `thread.output`) don't hang. + err = await handle.done + run_done = self._run_done + if run_done is not None and not run_done.done(): + if err is not None: + run_done.set_result( + _RunTerminal( + status="errored", + error=RuntimeError(f"Lifecycle transport failed: {err}"), + ) + ) + else: + run_done.set_result( + _RunTerminal( + status="errored", + error=RuntimeError( + "lifecycle stream ended before terminal event" + ), + ) + ) + return + except (Exception, asyncio.CancelledError) as exc: # Why: advisory-only watcher. Any error (HTTP failure, malformed # event in `_apply_lifecycle_event`, cancellation on close) must # not crash the caller; the watcher is one-shot best-effort. + # Resolve _run_done with an error so thread.output doesn't wait + # forever when the lifecycle transport fails. + run_done = self._run_done + if run_done is not None and not run_done.done(): + if not isinstance(exc, asyncio.CancelledError): + run_done.set_result( + _RunTerminal( + status="errored", + error=RuntimeError(f"Lifecycle transport failed: {exc}"), + ) + ) return - def _apply_lifecycle_event(self, event: Event) -> None: - """Update `interrupted` / `interrupts` state from a lifecycle or input event.""" + async def _apply_lifecycle_event(self, event: Event) -> None: + """Update `interrupted` / `interrupts` / `_run_done` from a lifecycle or input event.""" method = event.get("method") if method == "input.requested": params = event.get("params") or {} @@ -397,11 +574,41 @@ class AsyncThreadStream: if isinstance(params, dict) else [], } - self.interrupts.append(payload) - self.interrupted = True + async with self._interrupts_lock: + self.interrupts.append(payload) + self.interrupted = True elif method == "lifecycle": params = event.get("params") or {} data = params.get("data") if isinstance(params, dict) else None phase = data.get("phase") if isinstance(data, dict) else None - if phase in ("completed", "errored"): - self.interrupted = False + if phase in ("started", "running"): + # Mark that we have observed an active run so thread.output + # knows a run exists (handles reattach without run.start). + self._run_seen = True + elif phase in ("completed", "errored"): + # Why: interrupts describe current-run state. Clear on terminal + # lifecycle so a subsequent run.respond() can't fire against a + # stale prior-run interrupt_id. Acquire `_interrupts_lock` so + # any in-flight `run.respond` either completes against the + # pre-clear snapshot or sees the cleared state — never both. + async with self._interrupts_lock: + self.interrupted = False + self.interrupts = [] + run_done = self._run_done + if run_done is not None and not run_done.done(): + if phase == "errored": + error_msg = ( + data.get("error") if isinstance(data, dict) else None + ) + run_done.set_result( + _RunTerminal( + status="errored", + error=RuntimeError( + f"Run errored: {error_msg}" + if error_msg + else "Run errored" + ), + ) + ) + else: + run_done.set_result(_RunTerminal(status="completed")) diff --git a/libs/sdk-py/langgraph_sdk/_async/threads.py b/libs/sdk-py/langgraph_sdk/_async/threads.py index c56f77cb1..c5f722ac3 100644 --- a/libs/sdk-py/langgraph_sdk/_async/threads.py +++ b/libs/sdk-py/langgraph_sdk/_async/threads.py @@ -741,7 +741,8 @@ class ThreadsClient: thread_id: str | None = None, *, assistant_id: str, - headers: Mapping[str, str] | None = None, # noqa: ARG002 + headers: Mapping[str, str] | None = None, + run_start_timeout: float | None = None, ) -> AsyncThreadStream: """Open a v3 thread-centric streaming session. @@ -757,16 +758,21 @@ class ThreadsClient: thread_id: optional explicit thread identifier. Defaults to a fresh UUIDv4. assistant_id: assistant the run will use. Required. - headers: optional per-request headers. Reserved; not currently - forwarded. + headers: optional headers forwarded on every command and SSE + request for this stream session. + run_start_timeout: optional seconds to wait for an in-flight + `run.start` before subscribing operations raise + `asyncio.TimeoutError`. Defaults to `None` (wait forever). Returns: An `AsyncThreadStream` to use as an async context manager. """ return AsyncThreadStream( - client=self.http.client, + http=self.http, thread_id=thread_id if thread_id is not None else str(uuid.uuid4()), assistant_id=assistant_id, + headers=headers, + run_start_timeout=run_start_timeout, ) async def join_stream( diff --git a/libs/sdk-py/langgraph_sdk/stream/transport/http.py b/libs/sdk-py/langgraph_sdk/stream/transport/http.py index 705814c57..c29e388be 100644 --- a/libs/sdk-py/langgraph_sdk/stream/transport/http.py +++ b/libs/sdk-py/langgraph_sdk/stream/transport/http.py @@ -13,7 +13,7 @@ from __future__ import annotations import asyncio import contextlib -from collections.abc import AsyncIterator, Awaitable, Callable +from collections.abc import AsyncIterator, Awaitable, Callable, Mapping from dataclasses import dataclass from typing import Any, cast @@ -72,12 +72,14 @@ class ProtocolSseTransport: thread_id: str, commands_path: str | None = None, stream_path: str | None = None, + headers: Mapping[str, str] | None = None, max_queue_size: int = 1024, ) -> 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._max_queue_size = max_queue_size self._closed = False self._event_streams: set[asyncio.Task[None]] = set() @@ -92,10 +94,12 @@ class ProtocolSseTransport: """ if self._closed: raise RuntimeError("Protocol transport is closed.") + # Merge default headers first so content-type always wins. + merged_headers = {**self._default_headers, "content-type": "application/json"} response = await self._client.post( self._commands_url, content=orjson.dumps(command), - headers={"content-type": "application/json"}, + headers=merged_headers, ) response.raise_for_status() if response.status_code in (202, 204): @@ -134,15 +138,18 @@ class ProtocolSseTransport: async def pump() -> None: try: + # Merge default headers first so fixed SSE headers always win. + sse_headers = { + **self._default_headers, + "content-type": "application/json", + "accept": "text/event-stream", + "cache-control": "no-store", + } async with self._client.stream( "POST", self._stream_url, content=orjson.dumps(_build_event_stream_body(params)), - headers={ - "content-type": "application/json", - "accept": "text/event-stream", - "cache-control": "no-store", - }, + headers=sse_headers, ) as response: response.raise_for_status() if not ready.done(): diff --git a/libs/sdk-py/tests/streaming/_events.py b/libs/sdk-py/tests/streaming/_events.py index 5127d9808..86eb9fc82 100644 --- a/libs/sdk-py/tests/streaming/_events.py +++ b/libs/sdk-py/tests/streaming/_events.py @@ -27,6 +27,31 @@ def lifecycle_event( return _base(seq, "lifecycle", namespace or [], data or {"phase": "started"}) +def lifecycle_started_event( + seq: int = 0, namespace: list[str] | None = None +) -> dict[str, Any]: + """Lifecycle event with `phase="started"`.""" + return _base(seq, "lifecycle", namespace or [], {"phase": "started"}) + + +def lifecycle_completed_event( + seq: int = 0, namespace: list[str] | None = None +) -> dict[str, Any]: + """Lifecycle event with `phase="completed"`.""" + return _base(seq, "lifecycle", namespace or [], {"phase": "completed"}) + + +def lifecycle_errored_event( + seq: int = 0, + namespace: list[str] | None = None, + error: str = "run errored", +) -> dict[str, Any]: + """Lifecycle event with `phase="errored"` and an error message.""" + return _base( + seq, "lifecycle", namespace or [], {"phase": "errored", "error": error} + ) + + def values_event( seq: int = 0, namespace: list[str] | None = None, **data: Any ) -> dict[str, Any]: diff --git a/libs/sdk-py/tests/streaming/_fake_server.py b/libs/sdk-py/tests/streaming/_fake_server.py index b9dce4531..d82ec2039 100644 --- a/libs/sdk-py/tests/streaming/_fake_server.py +++ b/libs/sdk-py/tests/streaming/_fake_server.py @@ -5,6 +5,7 @@ just closely enough to validate the client: - POST /threads/{thread_id}/commands - POST /threads/{thread_id}/stream/events + - GET /threads/{thread_id}/state """ from __future__ import annotations @@ -27,22 +28,48 @@ class FakeServer: received_commands: every command body posted to /commands, in order. scripted_events: events the next /stream/events call will replay. stream_request_bodies: bodies posted to /stream/events, in order. + command_request_headers: headers from each POST to /commands, in order. + stream_request_headers_list: headers from each POST to /stream/events, in order. + state: the `ThreadState`-shaped dict returned by GET /threads/{thread_id}/state. + state_request_count: number of times the state endpoint has been called. + state_request_headers: headers from each GET to /threads/{thread_id}/state, in order. """ def __init__(self) -> None: self.received_commands: list[dict[str, Any]] = [] self.scripted_events: 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._stream_delay: float = 0.0 self._app: Starlette | None = None self.open_event_streams = 0 self._open_event_streams_max = 0 + self.state: dict[str, Any] = {} + self.state_request_count: int = 0 + self.state_request_headers: list[dict[str, str]] = [] def script(self, events: list[dict[str, Any]], *, delay: float = 0.0) -> None: """Set the events the next /stream/events call will replay.""" self.scripted_events = list(events) self._stream_delay = delay + def set_state( + self, + values: dict[str, Any], + next: list[Any] | None = None, + metadata: dict[str, Any] | None = None, + ) -> None: + """Store a `ThreadState`-shaped dict for GET /threads/{thread_id}/state.""" + 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, + } + @property def app(self) -> Starlette: if self._app is None: @@ -53,6 +80,7 @@ class FakeServer: async def commands(request: Request) -> Response: body = orjson.loads(await request.body()) self.received_commands.append(body) + self.command_request_headers.append(dict(request.headers)) command_id = body.get("id") return JSONResponse( { @@ -64,11 +92,17 @@ class FakeServer: async def stream_events(request: Request) -> Response: self.stream_request_bodies.append(orjson.loads(await request.body())) + self.stream_request_headers_list.append(dict(request.headers)) return StreamingResponse( self._sse_body(), media_type="text/event-stream", ) + async def thread_state(request: Request) -> Response: + self.state_request_count += 1 + self.state_request_headers.append(dict(request.headers)) + return JSONResponse(self.state) + return Starlette( routes=[ Route("/threads/{thread_id}/commands", commands, methods=["POST"]), @@ -77,6 +111,11 @@ class FakeServer: stream_events, methods=["POST"], ), + Route( + "/threads/{thread_id}/state", + thread_state, + methods=["GET"], + ), ] ) diff --git a/libs/sdk-py/tests/streaming/test_lifecycle_watcher.py b/libs/sdk-py/tests/streaming/test_lifecycle_watcher.py index 0422e27c6..94be14749 100644 --- a/libs/sdk-py/tests/streaming/test_lifecycle_watcher.py +++ b/libs/sdk-py/tests/streaming/test_lifecycle_watcher.py @@ -3,12 +3,14 @@ from __future__ import annotations import asyncio +import contextlib +from typing import Any import httpx from langgraph_sdk._async.http import HttpClient from langgraph_sdk._async.threads import ThreadsClient -from streaming._events import input_requested_event +from streaming._events import input_requested_event, lifecycle_event from streaming._fake_server import FakeServer @@ -36,3 +38,183 @@ async def test_interrupts_populated_from_input_requested_event(): assert thread.interrupted is True assert len(thread.interrupts) == 1 assert thread.interrupts[0]["interrupt_id"] == "i-1" + + +async def test_aenter_starts_lifecycle_watcher(): + """Entering AsyncThreadStream opens lifecycle/input SSE before run.start.""" + fake = FakeServer() + fake.script([lifecycle_event(seq=0, phase="started")]) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + # The lifecycle watcher task must be created on __aenter__, no run.start needed. + assert thread._lifecycle_watcher_task is not None + # Poll until the watcher has consumed the started event. + for _ in range(20): + if thread._run_seen: + break + await asyncio.sleep(0.05) + assert thread._run_seen is True + # No run.start was ever called — but the server still received a stream request. + assert len(fake.stream_request_bodies) >= 1 + + +async def test_reattach_observes_terminal_state(): + """Reattach (no run.start) consumes lifecycle replay and observes terminal state.""" + fake = FakeServer() + fake.script( + [ + lifecycle_event(seq=0, phase="running"), + lifecycle_event(seq=1, phase="completed"), + ] + ) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="existing", assistant_id="agent") as thread: + # Never call run.start — this is a reattach scenario. + # Poll until _run_done is resolved. + for _ in range(20): + run_done = thread._run_done + if run_done is not None and run_done.done(): + break + await asyncio.sleep(0.05) + assert thread._run_done is not None + assert thread._run_done.done() + terminal = thread._run_done.result() + assert terminal.status == "completed" + assert terminal.error is None + + +async def test_terminal_lifecycle_clears_interrupts(): + """Terminal lifecycle event clears interrupted/interrupts.""" + fake = FakeServer() + fake.script([lifecycle_event(seq=0, phase="completed")]) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + # Manually set interrupted state to simulate a prior interrupt. + thread.interrupted = True + thread.interrupts = [ + {"interrupt_id": "i-1", "value": None, "namespace": []} + ] + # Poll until the lifecycle watcher processes the completed event. + for _ in range(20): + if not thread.interrupted: + break + await asyncio.sleep(0.05) + assert thread.interrupted is False + assert thread.interrupts == [] + + +async def test_lifecycle_error_captured_for_output(): + """Lifecycle error terminal state is captured in _run_done with error set.""" + fake = FakeServer() + fake.script([lifecycle_event(seq=0, phase="errored", error="something went wrong")]) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + # Poll until _run_done is resolved. + for _ in range(20): + run_done = thread._run_done + if run_done is not None and run_done.done(): + break + await asyncio.sleep(0.05) + assert thread._run_done is not None + assert thread._run_done.done() + terminal = thread._run_done.result() + assert terminal.status == "errored" + assert terminal.error is not None + assert "something went wrong" in str(terminal.error) + + +async def test_run_start_sets_run_seen(): + """run.start() sets _run_seen to True (even without lifecycle event).""" + fake = FakeServer() + fake.script([]) # No events; the command response is sufficient. + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + assert thread._run_seen is False + await thread.run.start(input={}) + # _run_seen is set synchronously in run.start, before awaiting the result. + assert thread._run_seen is True + + +async def test_lifecycle_clean_eof_resolves_run_done_with_errored(): + """If the lifecycle SSE stream ends cleanly (server closes without a + terminal `completed` or `errored` event), `_run_done` must resolve with + an errored terminal so awaiters don't hang.""" + import pytest + + fake = FakeServer() + # Emit a non-terminal lifecycle event, then close cleanly without + # `completed` or `errored`. + fake.script([lifecycle_event(seq=0, phase="started")]) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + run_done = thread._run_done + assert run_done is not None + terminal = await asyncio.wait_for(run_done, timeout=2.0) + assert terminal.status == "errored" + assert terminal.error is not None + assert "ended before terminal" in str(terminal.error) + # Quiet unused-import warning under strict configs. + _ = pytest + + +async def test_lifecycle_mid_iteration_error_resolves_run_done_with_error( + monkeypatch: Any, +) -> None: + """If the transport reports an error via `handle.done` after iteration + exits without a terminal lifecycle event, `_run_done` propagates the + transport error rather than the generic clean-EOF message.""" + from langgraph_sdk.stream.transport import EventStreamHandle, ProtocolSseTransport + + def synthetic_handle() -> EventStreamHandle: + 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("simulated transport error")) + + async def empty_events() -> Any: + if False: + yield # pragma: no cover # make this an async generator + return + + async def noop_close() -> None: + return + + return EventStreamHandle( + events=empty_events(), + ready=ready, + done=done, + close=noop_close, + ) + + def patched_open(_self: ProtocolSseTransport, _params: Any) -> EventStreamHandle: + return synthetic_handle() + + monkeypatch.setattr(ProtocolSseTransport, "open_event_stream", patched_open) + + fake = FakeServer() + fake.script([]) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + run_done = thread._run_done + assert run_done is not None + terminal = await asyncio.wait_for(run_done, timeout=2.0) + assert terminal.status == "errored" + assert terminal.error is not None + assert "simulated transport error" in str(terminal.error) + # Quiet unused-import warnings under strict configs. + _ = contextlib diff --git a/libs/sdk-py/tests/streaming/test_thread_stream.py b/libs/sdk-py/tests/streaming/test_thread_stream.py index ff3009a55..bbbe1e702 100644 --- a/libs/sdk-py/tests/streaming/test_thread_stream.py +++ b/libs/sdk-py/tests/streaming/test_thread_stream.py @@ -16,7 +16,7 @@ from streaming._fake_server import FakeServer async def test_thread_stream_stores_thread_id_and_assistant_id(): async with httpx.AsyncClient(base_url="http://test") as client: stream = AsyncThreadStream( - client=client, + http=HttpClient(client), thread_id="t-1", assistant_id="agent", ) @@ -26,14 +26,18 @@ async def test_thread_stream_stores_thread_id_and_assistant_id(): async def test_aenter_returns_self(): async with httpx.AsyncClient(base_url="http://test") as client: - stream = AsyncThreadStream(client=client, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(client), thread_id="t-1", assistant_id="agent" + ) async with stream as entered: assert entered is stream async def test_aexit_marks_closed(): async with httpx.AsyncClient(base_url="http://test") as client: - stream = AsyncThreadStream(client=client, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(client), thread_id="t-1", assistant_id="agent" + ) async with stream: assert stream._closed is False assert stream._closed is True @@ -41,7 +45,9 @@ async def test_aexit_marks_closed(): async def test_close_is_idempotent(): async with httpx.AsyncClient(base_url="http://test") as client: - stream = AsyncThreadStream(client=client, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(client), thread_id="t-1", assistant_id="agent" + ) await stream.close() await stream.close() # must not raise assert stream._closed is True @@ -75,19 +81,73 @@ async def test_threads_stream_requires_assistant_id(): threads.stream(thread_id="t-1") # ty: ignore[missing-argument] -async def test_threads_stream_accepts_headers_kwarg(): - """`headers` is accepted as a kwarg even though it isn't forwarded yet.""" - from langgraph_sdk._async.http import HttpClient - from langgraph_sdk._async.threads import ThreadsClient - - async with httpx.AsyncClient(base_url="http://test") as raw: +async def test_threads_stream_headers_forwarded_to_commands(): + """Headers passed to `threads.stream()` are forwarded to /commands requests.""" + fake = FakeServer() + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as raw: threads = ThreadsClient(HttpClient(raw)) - stream = threads.stream( + async with threads.stream( thread_id="t-1", assistant_id="agent", - headers={"X-Foo": "bar"}, - ) - assert stream.thread_id == "t-1" + headers={"X-Custom-Header": "my-value"}, + ) as thread: + await thread.run.start(input={}) + assert fake.command_request_headers, "no command requests captured" + assert fake.command_request_headers[0].get("x-custom-header") == "my-value" + + +async def test_threads_stream_headers_forwarded_to_stream_events(): + """Headers passed to `threads.stream()` are forwarded to /stream/events requests.""" + fake = FakeServer() + fake.script([lifecycle_event(seq=0)]) + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream( + thread_id="t-1", + assistant_id="agent", + headers={"X-Custom-Header": "my-value"}, + ) as thread: + await thread.run.start(input={}) + _ = [e async for e in thread.subscribe(["lifecycle"])] + assert fake.stream_request_headers_list, "no stream/events requests captured" + assert fake.stream_request_headers_list[0].get("x-custom-header") == "my-value" + + +async def test_no_headers_by_default(): + """When `headers` is omitted, `_headers` is an empty dict and no custom + headers appear in command or stream requests. + """ + fake = FakeServer() + fake.script([lifecycle_event(seq=0)]) + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + assert thread._headers == {} + await thread.run.start(input={}) + _ = [e async for e in thread.subscribe(["lifecycle"])] + # No custom header keys beyond the protocol-required / transport-required ones. + protocol_keys = { + "content-type", + "accept", + "cache-control", + "host", + "user-agent", + "accept-encoding", + "connection", + "transfer-encoding", + "content-length", + } + extra_command = { + k for k in fake.command_request_headers[0] if k.lower() not in protocol_keys + } + extra_stream = { + k for k in fake.stream_request_headers_list[0] if k.lower() not in protocol_keys + } + assert extra_command == set(), f"unexpected command headers: {extra_command}" + assert extra_stream == set(), f"unexpected stream headers: {extra_stream}" async def test_aenter_constructs_transport_with_thread_id(): @@ -174,7 +234,9 @@ async def test_run_start_raises_outside_context_manager(): import pytest async with httpx.AsyncClient(base_url="http://test") as raw: - stream = AsyncThreadStream(client=raw, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(raw), thread_id="t-1", assistant_id="agent" + ) with pytest.raises(RuntimeError, match="async with"): await stream.run.start(input={"x": 1}) @@ -281,7 +343,9 @@ async def test_events_terminates_on_aexit(): async def test_events_raises_outside_context_manager(): async with httpx.AsyncClient(base_url="http://test") as raw: - stream = AsyncThreadStream(client=raw, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(raw), thread_id="t-1", assistant_id="agent" + ) with pytest.raises(RuntimeError, match="async with"): _ = stream.events @@ -291,7 +355,9 @@ async def test_aexit_preserves_original_exception_if_close_raises(): body's exception must propagate. close()'s error is suppressed (chained as context on close_err, but does not replace the original).""" async with httpx.AsyncClient(base_url="http://test") as raw: - thread = AsyncThreadStream(client=raw, thread_id="t-1", assistant_id="agent") + thread = AsyncThreadStream( + http=HttpClient(raw), thread_id="t-1", assistant_id="agent" + ) async def failing_close(): raise RuntimeError("close failed") @@ -358,7 +424,9 @@ async def test_fresh_thread_happy_path_end_to_end(): async def test_aenter_raises_after_close(): async with httpx.AsyncClient(base_url="http://test") as raw: - stream = AsyncThreadStream(client=raw, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(raw), thread_id="t-1", assistant_id="agent" + ) async with stream: pass # After exit, the stream is closed; re-entering must raise rather than @@ -370,7 +438,9 @@ async def test_aenter_raises_after_close(): async def test_register_subscription_assigns_monotonic_ids(): async with httpx.AsyncClient(base_url="http://test") as raw: - stream = AsyncThreadStream(client=raw, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(raw), thread_id="t-1", assistant_id="agent" + ) async with stream: sub_a = stream._register_subscription({"channels": ["values"]}) sub_b = stream._register_subscription({"channels": ["messages"]}) @@ -382,8 +452,262 @@ async def test_register_subscription_assigns_monotonic_ids(): async def test_unregister_subscription_removes_from_registry(): async with httpx.AsyncClient(base_url="http://test") as raw: - stream = AsyncThreadStream(client=raw, thread_id="t-1", assistant_id="agent") + stream = AsyncThreadStream( + http=HttpClient(raw), thread_id="t-1", assistant_id="agent" + ) async with stream: sub = stream._register_subscription({"channels": ["values"]}) stream._unregister_subscription(sub.id) assert sub.id not in stream._subscriptions + + +async def test_await_run_start_gate_honors_timeout(): + """Gate must raise asyncio.TimeoutError if run.start never completes + within the configured timeout.""" + import asyncio + + async with httpx.AsyncClient(base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + # Install a never-resolving gate to simulate an in-flight + # run.start that will not complete within the timeout window. + loop = asyncio.get_running_loop() + thread._run_start_ready = loop.create_future() + with pytest.raises(asyncio.TimeoutError): + await thread._await_run_start_gate(timeout=0.1) + # Gate must still be pending after the timeout (no side effects). + assert thread._run_start_ready is not None + assert not thread._run_start_ready.done() + + +async def test_await_run_start_gate_returns_when_gate_resolves_in_time(): + """With a generous timeout and a gate that resolves promptly, the + gate returns without raising.""" + import asyncio + + async with httpx.AsyncClient(base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + loop = asyncio.get_running_loop() + gate: asyncio.Future[None] = loop.create_future() + thread._run_start_ready = gate + loop.call_later(0.01, lambda: gate.set_result(None)) + await thread._await_run_start_gate(timeout=1.0) + + +async def test_run_start_timeout_constructor_kwarg_forwarded_to_gate(): + """`run_start_timeout` constructor kwarg is stored and consulted by + `_reconcile_stream` via `_await_run_start_gate`.""" + import asyncio + + async with httpx.AsyncClient(base_url="http://test") as raw: + stream = AsyncThreadStream( + http=HttpClient(raw), + thread_id="t-1", + assistant_id="agent", + run_start_timeout=0.1, + ) + async with stream as thread: + loop = asyncio.get_running_loop() + # Install a never-resolving gate. + thread._run_start_ready = loop.create_future() + with pytest.raises(asyncio.TimeoutError): + # Reconcile must surface the timeout from the gate. + await thread._reconcile_stream({"channels": ["lifecycle"]}) + + +async def test_subscribe_waits_for_run_start_to_commit(): + """Subscribing before run.start commits must not race the server. + + With the gate: subscribers wait for run.start to return before opening + their SSE. Without it, a fast subscribe would 404 against a thread the + server hasn't created yet. + """ + import asyncio + + fake = FakeServer() + fake.script([]) + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + # Kick run.start without awaiting — concurrently subscribe. + run_task = asyncio.create_task(thread.run.start(input={})) + sub_iter = thread.subscribe(["lifecycle"]) + # Drain one event or hit EOF. The iterator's first __anext__ + # awaits _reconcile_stream which awaits the gate. + async for _ in sub_iter: + break + # If the gate works, run.start completed before the subscription + # opened its SSE (and thus before iteration finished). + assert run_task.done() + + +async def test_run_respond_dispatches_input_respond_command(): + fake = FakeServer() + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + await thread.run.start(input={}) + # Simulate one outstanding interrupt. + thread.interrupts.append( + {"interrupt_id": "i-1", "value": None, "namespace": []} + ) + thread.interrupted = True + await thread.run.respond("yes") + command = fake.received_commands[-1] + assert command["method"] == "input.respond" + assert command["params"]["interrupt_id"] == "i-1" + assert command["params"]["response"] == "yes" + assert command["params"]["namespace"] == [] + + +async def test_run_respond_with_explicit_interrupt_id(): + fake = FakeServer() + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + await thread.run.start(input={}) + thread.interrupts.extend( + [ + {"interrupt_id": "a", "value": None, "namespace": []}, + {"interrupt_id": "b", "value": None, "namespace": []}, + ] + ) + thread.interrupted = True + await thread.run.respond("pick", interrupt_id="b") + assert fake.received_commands[-1]["params"]["interrupt_id"] == "b" + assert fake.received_commands[-1]["params"]["namespace"] == [] + + +async def test_run_respond_raises_when_no_outstanding_interrupts(): + async with httpx.AsyncClient(base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + with pytest.raises(RuntimeError, match="no outstanding interrupt"): + await thread.run.respond("yes") + + +async def test_run_respond_raises_when_ambiguous_interrupt_id(): + async with httpx.AsyncClient(base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + thread.interrupts.extend( + [ + {"interrupt_id": "a", "value": None, "namespace": []}, + {"interrupt_id": "b", "value": None, "namespace": []}, + ] + ) + thread.interrupted = True + with pytest.raises(RuntimeError, match=r"ambiguous|interrupt_id"): + await thread.run.respond("yes") + + +async def test_run_respond_snapshots_interrupts_under_lock(): + """`respond()` must take a snapshot of `interrupts` under the + `_interrupts_lock`, so a concurrent terminal-event clear cannot + invalidate the in-flight dispatch. + + Verifies: if `_interrupts_lock` is held when `respond()` is called, + `respond()` blocks until the lock is released — proving it serializes + with the terminal-clear path that takes the same lock. + """ + import asyncio + + fake = FakeServer() + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + await thread.run.start(input={}) + thread.interrupts.append( + {"interrupt_id": "i-1", "value": None, "namespace": []} + ) + thread.interrupted = True + # Take the interrupts lock externally to block `respond()`. + assert hasattr(thread, "_interrupts_lock"), ( + "AsyncThreadStream must expose _interrupts_lock" + ) + await thread._interrupts_lock.acquire() + try: + # `respond()` must NOT complete while we hold the lock. + task = asyncio.create_task(thread.run.respond("yes")) + # Give the task a chance to start and reach the lock. + await asyncio.sleep(0.05) + assert not task.done(), ( + "respond() should be blocked waiting for _interrupts_lock" + ) + finally: + thread._interrupts_lock.release() + # Now `respond()` should complete. + await asyncio.wait_for(task, timeout=1.0) + command = fake.received_commands[-1] + assert command["method"] == "input.respond" + assert command["params"]["interrupt_id"] == "i-1" + + +async def test_terminal_lifecycle_clear_acquires_interrupts_lock(): + """Terminal lifecycle event clears `interrupts` under the same lock + that `respond()` uses, preventing TOCTOU between snapshot and + dispatch.""" + import asyncio + + fake = FakeServer() + # No scripted events; we exercise `_apply_lifecycle_event` directly. + fake.script([]) + asgi = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + thread.interrupts.append( + {"interrupt_id": "i-1", "value": None, "namespace": []} + ) + thread.interrupted = True + # Hold the lock; a completion event must block on it before + # clearing interrupts. + await thread._interrupts_lock.acquire() + try: + from typing import cast + + from langchain_protocol import Event + + terminal_event = cast( + Event, + { + "type": "event", + "method": "lifecycle", + "params": { + "namespace": [], + "data": {"phase": "completed"}, + }, + "seq": 99, + "event_id": "evt-99", + }, + ) + clear_task = asyncio.create_task( + thread._apply_lifecycle_event(terminal_event) + ) + await asyncio.sleep(0.05) + # Interrupts must still be present — clear is blocked. + assert thread.interrupted is True + assert len(thread.interrupts) == 1 + assert not clear_task.done() + finally: + thread._interrupts_lock.release() + await asyncio.wait_for(clear_task, timeout=1.0) + assert thread.interrupted is False + assert thread.interrupts == [] + + +async def test_run_respond_raises_when_explicit_interrupt_id_not_outstanding(): + async with httpx.AsyncClient(base_url="http://test") as raw: + threads = ThreadsClient(HttpClient(raw)) + async with threads.stream(thread_id="t-1", assistant_id="agent") as thread: + thread.interrupts.append( + {"interrupt_id": "a", "value": None, "namespace": []} + ) + thread.interrupted = True + with pytest.raises(RuntimeError, match="does not match"): + await thread.run.respond("yes", interrupt_id="nonexistent") diff --git a/libs/sdk-py/tests/streaming/test_transport_http.py b/libs/sdk-py/tests/streaming/test_transport_http.py index 5904f81bb..151b649ed 100644 --- a/libs/sdk-py/tests/streaming/test_transport_http.py +++ b/libs/sdk-py/tests/streaming/test_transport_http.py @@ -434,3 +434,102 @@ async def test_transport_close_cancels_open_event_streams(): pass await asyncio.wait_for(drain(), timeout=1.0) + + +async def test_default_headers_forwarded_to_send_command(): + """Headers passed at construction are sent on every command request.""" + from streaming._fake_server import FakeServer + + fake = FakeServer() + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + sse = ProtocolSseTransport( + client=client, + thread_id="t-1", + headers={"X-Trace-Id": "abc123"}, + ) + await sse.send_command({"id": 1, "method": "run.start", "params": {}}) + assert fake.command_request_headers[0].get("x-trace-id") == "abc123" + # content-type must not be clobbered by default headers + assert "application/json" in fake.command_request_headers[0].get("content-type", "") + + +async def test_default_headers_forwarded_to_open_event_stream(): + """Headers passed at construction are sent on every SSE stream request.""" + from streaming._fake_server import FakeServer + + fake = FakeServer() + fake.script([]) + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + sse = ProtocolSseTransport( + client=client, + thread_id="t-1", + headers={"X-Trace-Id": "abc123"}, + ) + handle = sse.open_event_stream({"channels": ["lifecycle"]}) + await asyncio.wait_for(handle.ready, timeout=1.0) + _ = [e async for e in handle.events] + await handle.close() + assert fake.stream_request_headers_list[0].get("x-trace-id") == "abc123" + # Fixed SSE headers must not be clobbered by default headers + assert "text/event-stream" in fake.stream_request_headers_list[0].get("accept", "") + + +async def test_default_headers_cannot_override_sse_fixed_headers(): + """Caller-supplied default headers must not override content-type or accept.""" + from streaming._fake_server import FakeServer + + fake = FakeServer() + fake.script([]) + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + sse = ProtocolSseTransport( + client=client, + thread_id="t-1", + headers={ + "content-type": "text/plain", + "accept": "application/json", + "cache-control": "max-age=3600", + }, + ) + handle = sse.open_event_stream({"channels": ["lifecycle"]}) + await asyncio.wait_for(handle.ready, timeout=1.0) + _ = [e async for e in handle.events] + await handle.close() + hdrs = fake.stream_request_headers_list[0] + assert "application/json" in hdrs.get("content-type", "") + assert "text/event-stream" in hdrs.get("accept", "") + assert hdrs.get("cache-control") == "no-store" + + +async def test_fake_server_state_endpoint(): + """State endpoint returns the set state and increments the counter.""" + from streaming._fake_server import FakeServer + + fake = FakeServer() + fake.set_state({"foo": "bar"}, next=["node_a"]) + transport = httpx.ASGITransport(app=fake.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + resp = await client.get("/threads/t-1/state") + assert resp.status_code == 200 + body = resp.json() + assert body["values"] == {"foo": "bar"} + assert body["next"] == ["node_a"] + assert body["tasks"] == [] + assert body["metadata"] == {} + assert body["checkpoint"] is None + assert body["created_at"] is None + assert fake.state_request_count == 1 + assert len(fake.state_request_headers) == 1 + + +def test_values_event_builder_shape(): + """values_event produces the expected shape with params.data as the snapshot.""" + from streaming._events import values_event + + evt = values_event(seq=1, values={"foo": 1}) + assert evt["event_id"] == "evt-1" + assert evt["method"] == "values" + assert evt["params"]["data"] == {"values": {"foo": 1}} + assert evt["params"]["namespace"] == []