feat(langgraph): add v3 streaming support to RemoteGraph (#7927)

## Summary

- Adds `stream_events(version="v3")` and `astream_events(version="v3")`
to `RemoteGraph`, matching the local `CompiledStateGraph` surface and
unblocking polymorphic v3 streaming over `Graph | RemoteGraph`.
- Implementation is a thin adapter
(`libs/langgraph/langgraph/pregel/_remote_run_stream.py`) that wraps the
v3 SDK's `AsyncThreadStream` / `SyncThreadStream` and duck-types
`GraphRunStream` / `AsyncGraphRunStream`. No coupling to local v3 mux
internals.
- `v1` / `v2` paths unchanged. `astream_events(version='v1'|'v2')` still
raises `NotImplementedError` (separate gap).

### Scope decisions baked into this PR

- Unsupported v3 kwargs hard-reject at dispatch with
`NotImplementedError`: `control`, `transformers`, `interrupt_before`,
`interrupt_after`, and any unknown `**kwargs`. Server / SDK don't plumb
these through v3 yet; easy to lift later.
- Sync `interleave()` raises `NotImplementedError` pointing callers at
`astream_events`. Real sync interleave would need drainer threads;
deferred since most sync RemoteGraph callers just iterate raw events.
- Async `interleave()` is best-effort ordering (client receive order),
documented as a divergence from local v3's monotonic stamp ordering.
- Adapter `interrupted` / `interrupts` properties are **non-blocking**
snapshots of the SDK's current state. This differs from local
`(Async)GraphRunStream.interrupted`, which pump-drives the run to
terminal before returning. Callers needing a wait-for-interrupt pattern
should drain a projection (e.g., `interleave('values')`) until the SDK's
paused sentinel fires. Documented in the adapter docstrings.

### Audit of impact

Existing RemoteGraph callers in this org all use the v2 `.stream()` /
`.astream()` path (deepagents production wrapper, langgraph-api test
graphs, langgraph-supervisor TS type guard). **Zero callers** use
`stream_events` / `astream_events` on RemoteGraph today, so the new v3
methods are net-new surface — no risk of breaking existing consumers.

### Out of scope (follow-ups)

- Bumping `libs/langgraph/pyproject.toml`'s `langgraph-sdk` constraint
from `<0.4.0` to `<0.5.0`. Deferred until 0.4.0 publishes to PyPI; dev
resolution unaffected via the editable workspace dep.
- Real `astream_events(version='v1'|'v2')` implementation.
- Server-side plumbing for `control` / `interrupt_before` /
`interrupt_after` on v3 runs.
- Sync `interleave()` via drainer threads.


## Test plan

- [x] \`make test\` in \`libs/langgraph/\`: 1874 passed, 4 skipped (43
new in \`test_remote_graph_v3.py\`)
- [x] \`make lint\` in \`libs/langgraph/\`: ruff + mypy clean
- [x] \`pytest -m integration
tests/integration/test_remote_graph_v3.py\` in \`libs/sdk-py/\` against
the docker stack: 4/4 passed in 1.35s
- [x] Manual smoke: \`RemoteGraph('tools_agent',
url='http://localhost:2024').astream_events(..., version='v3')\`
end-to-end against the v3 integration api
- [x] Existing RemoteGraph v2 tests untouched (31 passed, 3 skipped with
docker up)
- [x] Will need rebase after \`langgraph-sdk 0.4.0\` lands on PyPI and
the version constraint is bumped in a separate PR
This commit is contained in:
Nick Hollon
2026-05-29 17:13:43 -04:00
committed by GitHub
parent ac3f5b007b
commit 68fa011fc9
6 changed files with 1262 additions and 10 deletions
+2
View File
@@ -50,6 +50,8 @@ jobs:
- '**/uv.lock'
sdk_py:
- 'libs/sdk-py/**'
- 'libs/langgraph/langgraph/pregel/remote.py'
- 'libs/langgraph/langgraph/pregel/_remote_run_stream.py'
lint:
needs: changes
@@ -0,0 +1,379 @@
# libs/langgraph/langgraph/pregel/_remote_run_stream.py
from __future__ import annotations
import logging
import sys
from collections.abc import AsyncIterator, Iterator, Mapping
from types import TracebackType
from typing import Any, cast
from langchain_core.runnables import RunnableConfig
from langgraph_sdk._async.stream import AsyncThreadStream
from langgraph_sdk._sync.stream import SyncThreadStream
from langgraph_sdk.client import LangGraphClient, SyncLangGraphClient
from langgraph.types import Command
logger = logging.getLogger(__name__)
def _translate_command_input(input: Any) -> Any:
"""Translate a local `Command` into the v3 wire `input`, else passthrough.
The v3 server decides start-vs-resume from thread state (an interrupted
run or pending interrupts) and, on resume, wraps the whole `input` as
`{"resume": input}` itself. So a resume `Command` must surface its raw
`resume` value as the wire `input` (not the serialized dataclass, which
the server would double-wrap). The v3 `run.start` path has no `goto` /
`update` channel, so those are rejected.
`langgraph_sdk` is upstream of `langgraph`, so this `Command`-aware
marshalling lives here on the adapter (langgraph) side of the boundary.
"""
if isinstance(input, Command):
if input.goto or input.update:
raise NotImplementedError(
"RemoteGraph v3 streaming supports `Command(resume=...)` only; "
"`goto` / `update` are not supported by the v3 `run.start` path."
)
return input.resume
return input
class _ChannelProjection:
"""Decoded projection for a wire channel the SDK doesn't type natively.
Subscribes to `channel` and yields each event's `params["data"]` — the same
item shape the SDK's typed projections yield (`_ValuesProjection` etc.) and
that local's `UpdatesTransformer` / `CheckpointsTransformer` /
`TasksTransformer` / `CustomTransformer` push, so iterating this matches the
corresponding local projection. Iterate with `for` against a sync stream and
`async for` against an async stream (matching the underlying SDK). Opening
the subscription requires the stream to be entered (`with` / `async with`).
"""
def __init__(self, sdk: AsyncThreadStream | SyncThreadStream, channel: str) -> None:
self._sdk = sdk
self._channel = channel
@staticmethod
def _data(event: Any) -> Any:
"""Extract `params.data` from a protocol event, tolerating odd shapes."""
params = event.get("params") if isinstance(event, dict) else None
return params.get("data") if isinstance(params, dict) else None
def __iter__(self) -> Iterator[Any]:
# Sync lane: the sync adapter's SDK returns a sync iterator here.
events = cast(Iterator[Any], self._sdk.subscribe([self._channel]))
for event in events:
data = self._data(event)
if data is not None:
yield data
def __aiter__(self) -> AsyncIterator[Any]:
return self._aiter()
async def _aiter(self) -> AsyncIterator[Any]:
# Async lane: the async adapter's SDK returns an async iterator here.
events = cast(AsyncIterator[Any], self._sdk.subscribe([self._channel]))
async for event in events:
data = self._data(event)
if data is not None:
yield data
class _ProjectionRegistry(Mapping[str, Any]):
"""Read-only name -> projection registry mirroring local `GraphRunStream.extensions`.
Resolution follows the langchain-protocol wire channels, and every entry
yields the same decoded item shape local does (`params.data`):
- `values` / `messages` / `tool_calls` / `subgraphs` resolve to the SDK's
decoded typed projections. `tool_calls` is the `tools` channel — tool
*execution* events, distinct from the tool-call *inputs* inside `messages`.
- `updates` / `checkpoints` / `tasks` / `custom` have no typed SDK
projection, so they resolve to a `_ChannelProjection` that subscribes to
the channel and yields `params.data` — matching the local transformer
output for those channels.
- any other name is a specific custom-extension channel
(`thread.extensions[name]`, i.e. `custom:<name>`).
`lifecycle` is intentionally absent: local derives a status payload from it
rather than yielding `params.data`, and the SDK consumes it as control-plane
(driving `output` / `interrupted`), so its shape can't be matched — it
remains reachable via the raw `events` iterator. `debug` is absent too: it
is not a v3 wire channel.
"""
# Channels the SDK decodes into typed projections.
_TYPED = ("values", "messages", "tool_calls", "subgraphs")
# Wire channels with no typed SDK projection — decoded here to match local.
_DECODED = ("updates", "checkpoints", "tasks", "custom")
_NATIVE = _TYPED + _DECODED
def __init__(self, sdk: AsyncThreadStream | SyncThreadStream) -> None:
self._sdk = sdk
def __getitem__(self, name: str) -> Any:
if name in self._TYPED:
return getattr(self._sdk, name)
if name in self._DECODED:
return _ChannelProjection(self._sdk, name)
return self._sdk.extensions[name]
def __iter__(self) -> Iterator[str]:
return iter(self._NATIVE)
def __len__(self) -> int:
return len(self._NATIVE)
class _RemoteGraphRunStream:
"""Sync adapter: SyncThreadStream -> GraphRunStream surface."""
def __init__(
self,
*,
sync_client: SyncLangGraphClient,
sdk_thread: SyncThreadStream,
input: Any,
config: RunnableConfig | None,
metadata: dict[str, Any] | None,
) -> None:
self._client = sync_client
self._sdk = sdk_thread
self._start_kwargs: dict[str, Any] = {
"input": _translate_command_input(input),
"config": config,
"metadata": metadata,
}
self._run_id: str | None = None
self._closed = False
self._events_iter: Iterator[Any] | None = None
def __enter__(self) -> _RemoteGraphRunStream:
if self._closed:
raise RuntimeError("_RemoteGraphRunStream already closed")
self._sdk.__enter__()
try:
result = self._sdk.run.start(**self._start_kwargs)
except BaseException:
self._sdk.__exit__(*sys.exc_info())
raise
self._run_id = result["run_id"]
return self
def __exit__(
self,
exc_type: type[BaseException] | None,
exc: BaseException | None,
tb: TracebackType | None,
) -> None:
if self._closed:
return
self._closed = True
self._sdk.__exit__(exc_type, exc, tb)
@property
def output(self) -> Any:
return self._sdk.output
@property
def interrupted(self) -> bool:
"""Whether the remote run is currently paused at an interrupt.
Reads the SDK's current value without blocking. This differs from
local `GraphRunStream.interrupted`, which drives the run to terminal
before returning the flag. Sync callers needing a wait-for-interrupt
pattern should switch to the async API and drain a projection.
"""
return self._sdk.interrupted
@property
def interrupts(self) -> list[Any]:
"""Current outstanding interrupt payloads (non-blocking snapshot)."""
return list(self._sdk.interrupts)
@property
def values(self) -> Any:
"""Live state-snapshot projection (mirrors local `run.values`)."""
return self._sdk.values
@property
def messages(self) -> Any:
"""Live message-stream projection (mirrors local `run.messages`)."""
return self._sdk.messages
@property
def subgraphs(self) -> Any:
"""Subgraph-handle projection (mirrors local `run.subgraphs`)."""
return self._sdk.subgraphs
@property
def tool_calls(self) -> Any:
"""Tool-execution projection (the `tools` channel).
These are tool *execution* events (started / output / finished),
distinct from the tool-call *inputs* carried inside `messages`.
"""
return self._sdk.tool_calls
@property
def extensions(self) -> Mapping[str, Any]:
"""Name -> projection registry (mirrors local `run.extensions`)."""
return _ProjectionRegistry(self._sdk)
def abort(self) -> None:
if self._closed:
return
self._closed = True
if self._run_id is not None:
try:
self._client.runs.cancel(self._sdk.thread_id, self._run_id, wait=False)
except Exception:
logger.debug("abort: runs.cancel failed", exc_info=True)
try:
self._sdk.close()
except Exception:
logger.debug("abort: sdk.close failed", exc_info=True)
def __iter__(self) -> Iterator[Any]:
if self._events_iter is None:
self._events_iter = iter(self._sdk.events)
return self._events_iter
def interleave(self, *names: str) -> Iterator[tuple[str, Any]]:
raise NotImplementedError
class _AsyncRemoteGraphRunStream:
"""Async adapter: AsyncThreadStream -> AsyncGraphRunStream surface."""
def __init__(
self,
*,
client: LangGraphClient,
sdk_thread: AsyncThreadStream,
input: Any,
config: RunnableConfig | None,
metadata: dict[str, Any] | None,
) -> None:
self._client = client
self._sdk = sdk_thread
self._start_kwargs: dict[str, Any] = {
"input": _translate_command_input(input),
"config": config,
"metadata": metadata,
}
self._run_id: str | None = None
self._closed = False
self._events_aiter: AsyncIterator[Any] | None = None
async def __aenter__(self) -> _AsyncRemoteGraphRunStream:
if self._closed:
raise RuntimeError("_AsyncRemoteGraphRunStream already closed")
await self._sdk.__aenter__()
try:
result = await self._sdk.run.start(**self._start_kwargs)
except BaseException:
await self._sdk.__aexit__(*sys.exc_info())
raise
self._run_id = result["run_id"]
return self
async def __aexit__(
self,
exc_type: type[BaseException] | None,
exc: BaseException | None,
tb: TracebackType | None,
) -> None:
if self._closed:
return
self._closed = True
await self._sdk.__aexit__(exc_type, exc, tb)
async def output(self) -> Any:
"""Drive the remote run to completion and return the final state.
Awaits the SDK's terminal-state awaitable, matching local
`AsyncGraphRunStream.output()` (a method, not a property, so
`run.output` without `await` fails at type-check time rather than
silently yielding a coroutine).
"""
return await self._sdk.output
async def interrupted(self) -> bool:
"""Whether the remote run is currently paused at an interrupt.
Reads the SDK's current value without blocking. This differs from
local `AsyncGraphRunStream.interrupted()`, which drives the run to
terminal before returning the flag. Callers that need a
wait-for-interrupt pattern should drain a projection (e.g.,
`async for snap in stream._sdk.values`) until the SDK's paused
sentinel fires, then call this method.
"""
return self._sdk.interrupted
async def interrupts(self) -> list[Any]:
"""Current outstanding interrupt payloads.
Non-blocking; reads the SDK's current snapshot. See `interrupted`
for the divergence from local v3 semantics.
"""
return list(self._sdk.interrupts)
@property
def values(self) -> Any:
"""Live state-snapshot projection (mirrors local `run.values`)."""
return self._sdk.values
@property
def messages(self) -> Any:
"""Live message-stream projection (mirrors local `run.messages`)."""
return self._sdk.messages
@property
def subgraphs(self) -> Any:
"""Subgraph-handle projection (mirrors local `run.subgraphs`)."""
return self._sdk.subgraphs
@property
def tool_calls(self) -> Any:
"""Tool-execution projection (the `tools` channel).
These are tool *execution* events (started / output / finished),
distinct from the tool-call *inputs* carried inside `messages`.
"""
return self._sdk.tool_calls
@property
def extensions(self) -> Mapping[str, Any]:
"""Name -> projection registry (mirrors local `run.extensions`)."""
return _ProjectionRegistry(self._sdk)
async def abort(self) -> None:
if self._closed:
return
self._closed = True
if self._run_id is not None:
try:
await self._client.runs.cancel(
self._sdk.thread_id, self._run_id, wait=False
)
except Exception:
logger.debug("abort: runs.cancel failed", exc_info=True)
try:
await self._sdk.close()
except Exception:
logger.debug("abort: sdk.close failed", exc_info=True)
def __aiter__(self) -> AsyncIterator[Any]:
if self._events_aiter is None:
self._events_aiter = self._sdk.events.__aiter__()
return self._events_aiter
# Note: deliberately no `interleave()` on the async adapter. Local
# `AsyncGraphRunStream` doesn't have one either (async callers compose
# with `asyncio.gather` / `asyncio.as_completed`). The sync adapter
# provides `interleave()` because sync callers have no comparable
# primitive for iterating multiple iterators concurrently.
+126 -9
View File
@@ -55,6 +55,10 @@ from langgraph._internal._constants import (
NS_SEP,
)
from langgraph.errors import GraphInterrupt, ParentCommand
from langgraph.pregel._remote_run_stream import (
_AsyncRemoteGraphRunStream,
_RemoteGraphRunStream,
)
from langgraph.pregel.protocol import PregelProtocol, StreamProtocol
from langgraph.types import (
All,
@@ -80,6 +84,8 @@ _CONF_DROPLIST = frozenset(
),
)
_V3_SUPPORTED_KWARGS = frozenset({"metadata", "headers"})
def _sanitize_config_value(v: Any) -> Any:
"""Recursively sanitize a config value to ensure it contains only primitives."""
@@ -186,6 +192,34 @@ class RemoteGraph(PregelProtocol):
)
return self.sync_client
def _reject_v3_unsupported(
self,
*,
control: Any,
transformers: Any,
interrupt_before: Any,
interrupt_after: Any,
extra_kwargs: dict[str, Any],
) -> None:
"""Raise NotImplementedError for kwargs unsupported by the v3 streaming path."""
for name, value in (
("control", control),
("transformers", transformers),
("interrupt_before", interrupt_before),
("interrupt_after", interrupt_after),
):
if value:
raise NotImplementedError(
f"RemoteGraph.stream_events(version='v3') does not support `{name}=`."
)
unknown = set(extra_kwargs) - _V3_SUPPORTED_KWARGS
if unknown:
raise NotImplementedError(
f"RemoteGraph.stream_events(version='v3') does not support "
f"the following kwargs: {sorted(unknown)!r}. "
f"Supported: {sorted(_V3_SUPPORTED_KWARGS)!r}."
)
def copy(self, update: dict[str, Any]) -> Self:
attrs = {**self.__dict__, **update}
return self.__class__(attrs.pop("assistant_id"), **attrs)
@@ -996,21 +1030,104 @@ class RemoteGraph(PregelProtocol):
else:
yield chunk
def stream_events(
self,
input: Any,
config: RunnableConfig | None = None,
*,
version: Literal["v1", "v2", "v3"] = "v2",
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
control: Any = None,
transformers: Sequence[Any] | None = None,
headers: dict[str, str] | None = None,
**kwargs: Any,
) -> Any:
"""Stream events from this remote graph.
For `version="v3"`, returns a `_RemoteGraphRunStream` whose surface
matches the local `GraphRunStream`. For other versions, delegates to
`Runnable.stream_events`.
"""
if version != "v3":
return super().stream_events(input, config, version=version, **kwargs)
self._reject_v3_unsupported(
control=control,
transformers=transformers,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
extra_kwargs=kwargs,
)
sync_client = self._validate_sync_client()
sanitized = self._sanitize_config(merge_configs(self.config, config))
thread_id = sanitized.get("configurable", {}).pop("thread_id", None)
merged_headers = (
_merge_tracing_headers(headers) if self.distributed_tracing else headers
)
sdk_thread = sync_client.threads.stream(
thread_id=thread_id,
assistant_id=self.assistant_id,
headers=merged_headers,
)
return _RemoteGraphRunStream(
sync_client=sync_client,
sdk_thread=sdk_thread,
input=input,
config=sanitized,
metadata=kwargs.get("metadata"),
)
async def astream_events(
self,
input: Any,
config: RunnableConfig | None = None,
*,
version: Literal["v1", "v2"],
include_names: Sequence[All] | None = None,
include_types: Sequence[All] | None = None,
include_tags: Sequence[All] | None = None,
exclude_names: Sequence[All] | None = None,
exclude_types: Sequence[All] | None = None,
exclude_tags: Sequence[All] | None = None,
version: Literal["v1", "v2", "v3"] = "v2",
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
control: Any = None,
transformers: Sequence[Any] | None = None,
headers: dict[str, str] | None = None,
**kwargs: Any,
) -> AsyncIterator[dict[str, Any]]:
raise NotImplementedError
) -> Any:
"""Async-stream events from this remote graph.
For `version="v3"`, awaits to an `_AsyncRemoteGraphRunStream`, matching
the local `Pregel.astream_events(version="v3")` awaitable contract:
`async with await rg.astream_events(..., version="v3") as run`. For
`version="v1"`/`"v2"`, raises NotImplementedError (use `astream`).
"""
if version != "v3":
raise NotImplementedError(
f"RemoteGraph.astream_events(version={version!r}) is not "
"implemented; use astream() for v1/v2 streaming or "
"version='v3'."
)
self._reject_v3_unsupported(
control=control,
transformers=transformers,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
extra_kwargs=kwargs,
)
client = self._validate_client()
sanitized = self._sanitize_config(merge_configs(self.config, config))
thread_id = sanitized.get("configurable", {}).pop("thread_id", None)
merged_headers = (
_merge_tracing_headers(headers) if self.distributed_tracing else headers
)
sdk_thread = client.threads.stream(
thread_id=thread_id,
assistant_id=self.assistant_id,
headers=merged_headers,
)
return _AsyncRemoteGraphRunStream(
client=client,
sdk_thread=sdk_thread,
input=input,
config=sanitized,
metadata=kwargs.get("metadata"),
)
@overload
def invoke(
+1 -1
View File
@@ -26,7 +26,7 @@ classifiers = [
dependencies = [
"langchain-core>=1.4.0,<2",
"langgraph-checkpoint>=4.1.0,<5.0.0",
"langgraph-sdk>=0.3.0,<0.4.0",
"langgraph-sdk>=0.4.0,<0.5.0",
"langgraph-prebuilt>=1.1.0,<1.2.0",
"xxhash>=3.5.0",
"pydantic>=2.7.4",
@@ -0,0 +1,621 @@
from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock
import pytest
from langgraph.pregel._remote_run_stream import (
_AsyncRemoteGraphRunStream,
_ChannelProjection,
_ProjectionRegistry,
_RemoteGraphRunStream,
_translate_command_input,
)
from langgraph.pregel.remote import (
_V3_SUPPORTED_KWARGS,
RemoteGraph,
)
from langgraph.types import Command
def _make_sync_adapter(*, run_start_returns=None, run_start_raises=None):
sync_client = MagicMock()
sdk_thread = MagicMock()
sdk_thread.thread_id = "thread-abc"
sdk_thread.__enter__ = MagicMock(return_value=sdk_thread)
sdk_thread.__exit__ = MagicMock(return_value=None)
sdk_thread.run = MagicMock()
if run_start_raises is not None:
sdk_thread.run.start = MagicMock(side_effect=run_start_raises)
else:
sdk_thread.run.start = MagicMock(
return_value=run_start_returns or {"run_id": "run-xyz"}
)
adapter = _RemoteGraphRunStream(
sync_client=sync_client,
sdk_thread=sdk_thread,
input={"x": 1},
config={"configurable": {}},
metadata=None,
)
return adapter, sync_client, sdk_thread
def test_enter_calls_sdk_enter_then_run_start_and_captures_run_id():
adapter, _, sdk_thread = _make_sync_adapter()
with adapter as stream:
assert stream is adapter
sdk_thread.__enter__.assert_called_once()
sdk_thread.run.start.assert_called_once_with(
input={"x": 1}, config={"configurable": {}}, metadata=None
)
assert adapter._run_id == "run-xyz"
def test_exit_delegates_to_sdk_exit_with_exc_info():
adapter, _, sdk_thread = _make_sync_adapter()
with adapter:
pass
sdk_thread.__exit__.assert_called_once_with(None, None, None)
def test_enter_unwinds_sdk_cm_when_run_start_raises():
adapter, _, sdk_thread = _make_sync_adapter(
run_start_raises=RuntimeError("start boom")
)
with pytest.raises(RuntimeError, match="start boom"):
with adapter:
pytest.fail("body should not run")
sdk_thread.__enter__.assert_called_once()
sdk_thread.__exit__.assert_called_once()
exc_info = sdk_thread.__exit__.call_args.args
assert exc_info[0] is RuntimeError
assert isinstance(exc_info[1], RuntimeError)
assert adapter._run_id is None
def test_output_interrupted_interrupts_passthrough():
adapter, _, sdk_thread = _make_sync_adapter()
sdk_thread.output = {"foo": 1}
sdk_thread.interrupted = True
sdk_thread.interrupts = [{"interrupt_id": "i1", "namespace": [], "value": "v"}]
with adapter as stream:
assert stream.output == {"foo": 1}
assert stream.interrupted is True
assert stream.interrupts == [
{"interrupt_id": "i1", "namespace": [], "value": "v"}
]
def test_sync_projection_attrs_forward_to_sdk():
adapter, _, sdk_thread = _make_sync_adapter()
sdk_thread.values = object()
sdk_thread.messages = object()
sdk_thread.tool_calls = object()
sdk_thread.subgraphs = object()
with adapter as stream:
assert stream.values is sdk_thread.values
assert stream.messages is sdk_thread.messages
assert stream.tool_calls is sdk_thread.tool_calls
assert stream.subgraphs is sdk_thread.subgraphs
assert set(stream.extensions) == set(_ProjectionRegistry._NATIVE)
assert stream.extensions["values"] is sdk_thread.values
def test_projection_registry_typed_decoded_and_custom():
sdk = MagicMock()
sdk.values = object()
sdk.messages = object()
sdk.tool_calls = object()
sdk.subgraphs = object()
custom_named = object()
sdk.extensions = {"my_custom": custom_named}
registry = _ProjectionRegistry(sdk)
# Typed channels resolve to the SDK's decoded projections.
assert registry["values"] is sdk.values
assert registry["tool_calls"] is sdk.tool_calls
assert registry["subgraphs"] is sdk.subgraphs
# Channels without a typed projection resolve to a decoding _ChannelProjection.
ckpt = registry["checkpoints"]
assert isinstance(ckpt, _ChannelProjection)
assert ckpt._channel == "checkpoints"
assert isinstance(registry["updates"], _ChannelProjection)
# A non-protocol name is a specific custom-extension channel.
assert registry["my_custom"] is custom_named
# Enumerable set is the typed + decoded channels (no `lifecycle`, no `debug`).
assert list(registry) == [
"values",
"messages",
"tool_calls",
"subgraphs",
"updates",
"checkpoints",
"tasks",
"custom",
]
assert len(registry) == 8
def test_channel_projection_decodes_params_data():
sdk = MagicMock()
# Two events with data + one malformed/dataless event that must be skipped.
sdk.subscribe = MagicMock(
return_value=iter(
[
{"params": {"data": {"n": 1}}},
{"params": {}}, # no data -> skipped
{"params": {"data": {"n": 2}}},
{"unexpected": "shape"}, # not a params dict -> skipped
]
)
)
proj = _ChannelProjection(sdk, "checkpoints")
assert list(proj) == [{"n": 1}, {"n": 2}]
sdk.subscribe.assert_called_once_with(["checkpoints"])
def test_sync_adapter_translates_command_input():
sync_client = MagicMock()
sdk_thread = MagicMock()
adapter = _RemoteGraphRunStream(
sync_client=sync_client,
sdk_thread=sdk_thread,
input=Command(resume="go"),
config=None,
metadata=None,
)
assert adapter._start_kwargs["input"] == "go"
def test_iter_caches_first_subscription():
adapter, _, sdk_thread = _make_sync_adapter()
fake_events = [object(), object(), object()]
sdk_thread.events = iter(fake_events)
with adapter as stream:
first = iter(stream)
second = iter(stream)
assert first is second
assert list(first) == fake_events
def test_abort_cancels_run_and_closes_sdk():
adapter, sync_client, sdk_thread = _make_sync_adapter()
with adapter as stream:
stream.abort()
sync_client.runs.cancel.assert_called_once_with(
"thread-abc", "run-xyz", wait=False
)
sdk_thread.close.assert_called_once()
def test_abort_before_enter_skips_cancel_but_closes_sdk():
adapter, sync_client, sdk_thread = _make_sync_adapter()
adapter.abort()
sync_client.runs.cancel.assert_not_called()
sdk_thread.close.assert_called_once()
def test_abort_is_idempotent():
adapter, sync_client, sdk_thread = _make_sync_adapter()
with adapter as stream:
stream.abort()
stream.abort()
assert sync_client.runs.cancel.call_count == 1
assert sdk_thread.close.call_count == 1
def test_abort_swallows_cancel_failure_and_still_closes():
adapter, sync_client, sdk_thread = _make_sync_adapter()
sync_client.runs.cancel.side_effect = RuntimeError("cancel boom")
with adapter as stream:
stream.abort()
sdk_thread.close.assert_called_once()
def test_sync_interleave_raises_not_implemented():
adapter, _, _ = _make_sync_adapter()
with adapter as stream:
with pytest.raises(NotImplementedError):
list(stream.interleave("messages"))
def test_async_adapter_has_no_interleave():
"""Async adapter intentionally lacks `interleave` (mirrors local
`AsyncGraphRunStream`, which doesn't have one either). Async callers
compose with `asyncio.gather` / `asyncio.as_completed`.
"""
assert not hasattr(_AsyncRemoteGraphRunStream, "interleave")
def _make_async_adapter(*, run_start_returns=None, run_start_raises=None):
client = MagicMock()
client.runs.cancel = AsyncMock()
sdk_thread = MagicMock()
sdk_thread.thread_id = "thread-abc"
sdk_thread.__aenter__ = AsyncMock(return_value=sdk_thread)
sdk_thread.__aexit__ = AsyncMock(return_value=None)
sdk_thread.close = AsyncMock()
sdk_thread.run = MagicMock()
if run_start_raises is not None:
sdk_thread.run.start = AsyncMock(side_effect=run_start_raises)
else:
sdk_thread.run.start = AsyncMock(
return_value=run_start_returns or {"run_id": "run-xyz"}
)
adapter = _AsyncRemoteGraphRunStream(
client=client,
sdk_thread=sdk_thread,
input={"x": 1},
config={"configurable": {}},
metadata=None,
)
return adapter, client, sdk_thread
@pytest.mark.anyio
async def test_aenter_calls_sdk_aenter_then_run_start_and_captures_run_id():
adapter, _, sdk_thread = _make_async_adapter()
async with adapter as stream:
assert stream is adapter
sdk_thread.__aenter__.assert_awaited_once()
sdk_thread.run.start.assert_awaited_once_with(
input={"x": 1}, config={"configurable": {}}, metadata=None
)
assert adapter._run_id == "run-xyz"
@pytest.mark.anyio
async def test_aexit_delegates_to_sdk_aexit():
adapter, _, sdk_thread = _make_async_adapter()
async with adapter:
pass
sdk_thread.__aexit__.assert_awaited_once_with(None, None, None)
@pytest.mark.anyio
async def test_aenter_unwinds_sdk_cm_when_run_start_raises():
adapter, _, sdk_thread = _make_async_adapter(
run_start_raises=RuntimeError("start boom")
)
with pytest.raises(RuntimeError, match="start boom"):
async with adapter:
pytest.fail("body should not run")
sdk_thread.__aenter__.assert_awaited_once()
sdk_thread.__aexit__.assert_awaited_once()
assert adapter._run_id is None
@pytest.mark.anyio
async def test_async_output_interrupted_interrupts_passthrough():
adapter, _, sdk_thread = _make_async_adapter()
async def _fake_output_awaitable():
return {"foo": 1}
sdk_thread.output = _fake_output_awaitable()
sdk_thread.interrupted = True
sdk_thread.interrupts = [{"interrupt_id": "i1", "namespace": [], "value": "v"}]
async with adapter as stream:
assert await stream.output() == {"foo": 1}
assert await stream.interrupted() is True
assert await stream.interrupts() == [
{"interrupt_id": "i1", "namespace": [], "value": "v"}
]
@pytest.mark.anyio
async def test_async_projection_attrs_forward_to_sdk():
adapter, _, sdk_thread = _make_async_adapter()
sdk_thread.values = object()
sdk_thread.messages = object()
sdk_thread.tool_calls = object()
sdk_thread.subgraphs = object()
async with adapter as stream:
assert stream.values is sdk_thread.values
assert stream.messages is sdk_thread.messages
assert stream.tool_calls is sdk_thread.tool_calls
assert stream.subgraphs is sdk_thread.subgraphs
assert set(stream.extensions) == set(_ProjectionRegistry._NATIVE)
assert stream.extensions["messages"] is sdk_thread.messages
@pytest.mark.anyio
async def test_aiter_caches_first_subscription():
adapter, _, sdk_thread = _make_async_adapter()
class _FakeAsyncEvents:
def __init__(self, items):
self._items = list(items)
def __aiter__(self):
return self
async def __anext__(self):
if not self._items:
raise StopAsyncIteration
return self._items.pop(0)
sdk_thread.events = _FakeAsyncEvents([object(), object()])
async with adapter as stream:
first = stream.__aiter__()
second = stream.__aiter__()
assert first is second
@pytest.mark.anyio
async def test_async_abort_cancels_run_and_closes_sdk():
adapter, client, sdk_thread = _make_async_adapter()
async with adapter as stream:
await stream.abort()
client.runs.cancel.assert_awaited_once_with("thread-abc", "run-xyz", wait=False)
sdk_thread.close.assert_awaited_once()
@pytest.mark.anyio
async def test_async_abort_before_aenter_skips_cancel():
adapter, client, sdk_thread = _make_async_adapter()
await adapter.abort()
client.runs.cancel.assert_not_awaited()
sdk_thread.close.assert_awaited_once()
@pytest.mark.anyio
async def test_async_abort_swallows_cancel_failure():
adapter, client, sdk_thread = _make_async_adapter()
client.runs.cancel.side_effect = RuntimeError("cancel boom")
async with adapter as stream:
await stream.abort()
sdk_thread.close.assert_awaited_once()
def _make_remote_graph() -> RemoteGraph:
sync_client = MagicMock()
async_client = MagicMock()
rg = RemoteGraph(
"agent",
client=async_client,
sync_client=sync_client,
)
return rg
def test_reject_v3_unsupported_passes_when_all_clear():
rg = _make_remote_graph()
rg._reject_v3_unsupported(
control=None,
transformers=None,
interrupt_before=None,
interrupt_after=None,
extra_kwargs={},
)
@pytest.mark.parametrize(
"kwarg_name,kwarg_value",
[
("control", object()),
("transformers", [object()]),
("interrupt_before", ["node_a"]),
("interrupt_after", ["node_b"]),
],
)
def test_reject_v3_unsupported_raises_per_kwarg(kwarg_name, kwarg_value):
rg = _make_remote_graph()
kwargs = dict(
control=None,
transformers=None,
interrupt_before=None,
interrupt_after=None,
extra_kwargs={},
)
kwargs[kwarg_name] = kwarg_value
with pytest.raises(NotImplementedError, match=f"`{kwarg_name}=`"):
rg._reject_v3_unsupported(**kwargs)
def test_reject_v3_unsupported_raises_on_unknown_extra_kwarg():
rg = _make_remote_graph()
with pytest.raises(NotImplementedError, match="context"):
rg._reject_v3_unsupported(
control=None,
transformers=None,
interrupt_before=None,
interrupt_after=None,
extra_kwargs={"context": {}},
)
def test_reject_v3_unsupported_allows_metadata_and_headers():
rg = _make_remote_graph()
rg._reject_v3_unsupported(
control=None,
transformers=None,
interrupt_before=None,
interrupt_after=None,
extra_kwargs={"metadata": {"a": 1}, "headers": {"X": "y"}},
)
def test_translate_command_input_surfaces_raw_resume_value():
# The v3 server wraps the resume `input` as {"resume": input} itself, so the
# wire `input` must be the raw resume value, not the serialized dataclass.
assert _translate_command_input(Command(resume="go")) == "go"
assert _translate_command_input(Command(resume={"id": "v"})) == {"id": "v"}
def test_translate_command_input_rejects_goto_and_update():
with pytest.raises(NotImplementedError, match="goto"):
_translate_command_input(Command(goto="node_b"))
with pytest.raises(NotImplementedError, match="update"):
_translate_command_input(Command(update={"a": 1}))
def test_translate_command_input_passes_through_non_command():
assert _translate_command_input({"a": 1}) == {"a": 1}
assert _translate_command_input(None) is None
def test_v3_supported_kwargs_known_set():
assert _V3_SUPPORTED_KWARGS == frozenset({"metadata", "headers"})
def test_stream_events_v3_constructs_sdk_thread_with_sanitized_args():
sync_client = MagicMock()
sdk_thread = MagicMock()
sync_client.threads.stream.return_value = sdk_thread
rg = RemoteGraph(
"agent",
client=MagicMock(),
sync_client=sync_client,
)
result = rg.stream_events(
{"input_key": 1},
config={"configurable": {"thread_id": "t1", "user": "u"}},
version="v3",
)
assert isinstance(result, _RemoteGraphRunStream)
sync_client.threads.stream.assert_called_once()
call = sync_client.threads.stream.call_args
assert call.kwargs["thread_id"] == "t1"
assert call.kwargs["assistant_id"] == "agent"
assert call.kwargs["headers"] is None
def test_stream_events_v3_passes_none_thread_id_when_absent():
sync_client = MagicMock()
sync_client.threads.stream.return_value = MagicMock()
rg = RemoteGraph("agent", client=MagicMock(), sync_client=sync_client)
rg.stream_events({"x": 1}, version="v3")
call = sync_client.threads.stream.call_args
assert call.kwargs["thread_id"] is None
def test_stream_events_v3_rejects_unsupported_kwargs_before_sdk_call():
sync_client = MagicMock()
rg = RemoteGraph("agent", client=MagicMock(), sync_client=sync_client)
with pytest.raises(NotImplementedError, match="control"):
rg.stream_events({"x": 1}, version="v3", control=object())
sync_client.threads.stream.assert_not_called()
def test_stream_events_v3_translates_command_input():
sync_client = MagicMock()
sync_client.threads.stream.return_value = MagicMock()
rg = RemoteGraph("agent", client=MagicMock(), sync_client=sync_client)
# Resume Command surfaces its raw resume value as the wire `input`; the v3
# server wraps it as {"resume": input} once it detects the interrupt.
adapter = rg.stream_events(Command(resume="go"), version="v3")
assert adapter._start_kwargs["input"] == "go"
def test_stream_events_v3_rejects_goto_update_command():
rg = RemoteGraph("agent", client=MagicMock(), sync_client=MagicMock())
with pytest.raises(NotImplementedError, match="goto"):
rg.stream_events(Command(goto="node_b"), version="v3")
def test_stream_events_v3_strips_checkpoint_keys_from_configurable():
sync_client = MagicMock()
sync_client.threads.stream.return_value = MagicMock()
rg = RemoteGraph("agent", client=MagicMock(), sync_client=sync_client)
adapter = rg.stream_events(
{"x": 1},
config={
"configurable": {
"thread_id": "t1",
"checkpoint_id": "c1",
"checkpoint_ns": "ns",
"user": "u",
}
},
version="v3",
)
sent_config = adapter._start_kwargs["config"]
assert "checkpoint_id" not in sent_config["configurable"]
assert "checkpoint_ns" not in sent_config["configurable"]
assert sent_config["configurable"]["user"] == "u"
def test_stream_events_v3_merges_tracing_headers_when_distributed_tracing(
monkeypatch,
):
from langgraph.pregel import remote as remote_mod
sync_client = MagicMock()
sync_client.threads.stream.return_value = MagicMock()
rg = RemoteGraph(
"agent",
client=MagicMock(),
sync_client=sync_client,
distributed_tracing=True,
)
captured = {}
def fake_merge(headers):
captured["arg"] = headers
return {"x-ls-trace": "1", **(headers or {})}
monkeypatch.setattr(remote_mod, "_merge_tracing_headers", fake_merge)
rg.stream_events({"x": 1}, version="v3", headers={"X-Custom": "y"})
assert captured["arg"] == {"X-Custom": "y"}
sent_headers = sync_client.threads.stream.call_args.kwargs["headers"]
assert sent_headers["x-ls-trace"] == "1"
assert sent_headers["X-Custom"] == "y"
def test_stream_events_v3_passes_headers_unchanged_without_tracing():
sync_client = MagicMock()
sync_client.threads.stream.return_value = MagicMock()
rg = RemoteGraph(
"agent",
client=MagicMock(),
sync_client=sync_client,
distributed_tracing=False,
)
rg.stream_events({"x": 1}, version="v3", headers={"X-Custom": "y"})
sent_headers = sync_client.threads.stream.call_args.kwargs["headers"]
assert sent_headers == {"X-Custom": "y"}
def test_stream_events_non_v3_delegates_to_super():
rg = RemoteGraph("agent", client=MagicMock(), sync_client=MagicMock())
sync_client_attr = rg.sync_client
try:
rg.stream_events({"x": 1}, version="v2")
except Exception:
pass
sync_client_attr.threads.stream.assert_not_called()
@pytest.mark.anyio
async def test_astream_events_v3_constructs_sdk_thread():
client = MagicMock()
sdk_thread = MagicMock()
client.threads.stream.return_value = sdk_thread
rg = RemoteGraph("agent", client=client, sync_client=MagicMock())
result = await rg.astream_events(
{"x": 1},
config={"configurable": {"thread_id": "t1"}},
version="v3",
)
assert isinstance(result, _AsyncRemoteGraphRunStream)
call = client.threads.stream.call_args
assert call.kwargs["thread_id"] == "t1"
assert call.kwargs["assistant_id"] == "agent"
@pytest.mark.anyio
async def test_astream_events_v3_rejects_unsupported_kwargs():
client = MagicMock()
rg = RemoteGraph("agent", client=client, sync_client=MagicMock())
with pytest.raises(NotImplementedError, match="transformers"):
await rg.astream_events({"x": 1}, version="v3", transformers=[object()])
client.threads.stream.assert_not_called()
@pytest.mark.anyio
async def test_astream_events_non_v3_raises_not_implemented():
rg = RemoteGraph("agent", client=MagicMock(), sync_client=MagicMock())
with pytest.raises(NotImplementedError, match="not implemented"):
await rg.astream_events({"x": 1}, version="v2")
@@ -0,0 +1,133 @@
"""Integration tests for RemoteGraph v3 streaming.
Tests the end-to-end wiring: RemoteGraph -> langgraph_sdk client.threads.stream(...) ->
docker-running langgraph-api -> SSE projections -> adapter classes.
Run with: pytest tests/integration/test_remote_graph_v3.py -m integration
"""
from __future__ import annotations
import uuid
import pytest
from langchain_core.runnables import RunnableConfig
from langgraph.pregel.remote import RemoteGraph
from langgraph.types import Command
pytestmark = pytest.mark.integration
URL = "http://localhost:2024"
# Input shapes matching what the integration graphs expect.
# `agent` graph (streaming_graph.py): AgentState has messages, value, items.
# `tools_agent` graph (tools_agent.py): create_agent graph expects messages list.
_AGENT_INPUT = {"messages": [], "value": "init", "items": []}
_TOOLS_AGENT_INPUT = {"messages": [{"role": "user", "content": "search for v3"}]}
@pytest.fixture
def remote_agent() -> RemoteGraph:
return RemoteGraph("agent", url=URL)
@pytest.fixture
def remote_tools_agent() -> RemoteGraph:
return RemoteGraph("tools_agent", url=URL)
async def test_async_happy_path_yields_output(remote_tools_agent: RemoteGraph) -> None:
"""tools_agent completes without interrupt; ``await stream.output`` drives
the run to terminal via the lifecycle watcher (no explicit event iteration
needed — the SSE subscription stays open by design after run completion)."""
async with await remote_tools_agent.astream_events(
_TOOLS_AGENT_INPUT,
version="v3",
) as stream:
output = await stream.output()
assert output is not None
assert (await stream.interrupted()) is False
async def test_async_interrupt_path_surfaces_interrupts(
remote_agent: RemoteGraph,
) -> None:
"""agent graph hits ask_human; interrupted must be True with >= 1 interrupt.
Note: interrupts pause the run but DON'T resolve `_run_done` (only
`completed` / `failed` lifecycle phases do), so `await stream.output()`
would hang. The adapter doesn't expose `interleave()` on the async
side (mirrors local `AsyncGraphRunStream`), so drain the `values`
projection directly until the run reports it is interrupted.
"""
async with await remote_agent.astream_events(
_AGENT_INPUT,
version="v3",
) as stream:
async for _ in stream.values:
if await stream.interrupted():
break
assert (await stream.interrupted()) is True
interrupts = await stream.interrupts()
assert len(interrupts) >= 1
async def test_async_resume_after_interrupt(remote_agent: RemoteGraph) -> None:
"""Interrupt the agent at ask_human, then resume the SAME thread with
`Command(resume=...)`.
Validates the v3 resume path end-to-end. The client sends the raw resume
value as `input` (not a serialized Command); the server detects the
thread's pending interrupt from persisted state — which survives the first
session's close — and wraps it as `Command(resume=...)`, driving the run
past `ask_human` to completion (the graph interrupts only once).
"""
thread_id = str(uuid.uuid4())
config: RunnableConfig = {"configurable": {"thread_id": thread_id}}
# First session: drive until the agent pauses at the ask_human interrupt.
async with await remote_agent.astream_events(
_AGENT_INPUT,
config=config,
version="v3",
) as stream:
async for _ in stream.values:
if await stream.interrupted():
break
assert (await stream.interrupted()) is True
# Second session on the same thread: resume with the human's answer. The
# run continues past ask_human to completion with no further interrupt.
async with await remote_agent.astream_events(
Command(resume="yes"),
config=config,
version="v3",
) as stream:
output = await stream.output()
assert output is not None
assert (await stream.interrupted()) is False
def test_sync_happy_path_yields_output(remote_tools_agent: RemoteGraph) -> None:
"""Sync stream: tools_agent completes; ``stream.output`` (sync property)
blocks until terminal."""
with remote_tools_agent.stream_events(
_TOOLS_AGENT_INPUT,
version="v3",
) as stream:
output = stream.output
assert output is not None
assert stream.interrupted is False
async def test_abort_mid_run_cancels_server_side(
remote_tools_agent: RemoteGraph,
) -> None:
"""Abort immediately after run.start; reaching the end without exception
confirms abort + __aexit__ cleanup worked."""
async with await remote_tools_agent.astream_events(
_TOOLS_AGENT_INPUT,
version="v3",
) as stream:
await stream.abort()
# Reaching here without unhandled exceptions confirms abort + __aexit__ succeeded.