mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-07 18:27:52 +02:00
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:
@@ -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.
|
||||
@@ -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(
|
||||
|
||||
@@ -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.
|
||||
Reference in New Issue
Block a user