diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 9aab80b8a..2626abb08 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -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 diff --git a/libs/langgraph/langgraph/pregel/_remote_run_stream.py b/libs/langgraph/langgraph/pregel/_remote_run_stream.py new file mode 100644 index 000000000..06d8efd09 --- /dev/null +++ b/libs/langgraph/langgraph/pregel/_remote_run_stream.py @@ -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:`). + + `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. diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index 973768a4a..b5c4245f0 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -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( diff --git a/libs/langgraph/pyproject.toml b/libs/langgraph/pyproject.toml index c7bebbad5..d7fe15d98 100644 --- a/libs/langgraph/pyproject.toml +++ b/libs/langgraph/pyproject.toml @@ -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", diff --git a/libs/langgraph/tests/test_remote_graph_v3.py b/libs/langgraph/tests/test_remote_graph_v3.py new file mode 100644 index 000000000..db7f53e56 --- /dev/null +++ b/libs/langgraph/tests/test_remote_graph_v3.py @@ -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") diff --git a/libs/sdk-py/tests/integration/test_remote_graph_v3.py b/libs/sdk-py/tests/integration/test_remote_graph_v3.py new file mode 100644 index 000000000..42562133a --- /dev/null +++ b/libs/sdk-py/tests/integration/test_remote_graph_v3.py @@ -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.