From f6286dce3845e2ce9a41d300d3384185d2fb570f Mon Sep 17 00:00:00 2001 From: Christian Bromann Date: Wed, 11 Mar 2026 21:58:31 -0700 Subject: [PATCH] feat(langgraph): add protocol improvements for better streaming --- libs/langgraph/langgraph/pregel/_messages.py | 11 +- libs/langgraph/langgraph/pregel/main.py | 2 + libs/langgraph/langgraph/pregel/remote.py | 97 +++++++++++-- libs/langgraph/langgraph/types.py | 9 +- libs/langgraph/tests/test_remote_graph.py | 141 +++++++++++++++++++ libs/langgraph/tests/test_stream_v2.py | 68 ++++++++- libs/sdk-py/langgraph_sdk/_async/runs.py | 12 ++ libs/sdk-py/langgraph_sdk/_sync/runs.py | 12 ++ libs/sdk-py/langgraph_sdk/schema.py | 23 +++ libs/sdk-py/tests/test_client_stream.py | 81 +++++++++++ 10 files changed, 435 insertions(+), 21 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/_messages.py b/libs/langgraph/langgraph/pregel/_messages.py index acab06098..8ab7574e2 100644 --- a/libs/langgraph/langgraph/pregel/_messages.py +++ b/libs/langgraph/langgraph/pregel/_messages.py @@ -25,7 +25,7 @@ except ImportError: _StreamingCallbackHandler = object # type: ignore T = TypeVar("T") -Meta = tuple[tuple[str, ...], dict[str, Any]] +Meta = tuple[tuple[str, ...], dict[str, Any] | None] def _state_values(obj: Any) -> Sequence[Any]: @@ -56,6 +56,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): subgraphs: bool, *, parent_ns: tuple[str, ...] | None = None, + dedupe_metadata: bool = False, ) -> None: """Configure the handler to stream messages from LLMs and nodes. @@ -84,8 +85,10 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): self.stream = stream self.subgraphs = subgraphs self.metadata: dict[UUID, Meta] = {} + self.emitted_metadata: set[UUID] = set() self.seen: set[int | str] = set() self.parent_ns = parent_ns + self.dedupe_metadata = dedupe_metadata def _emit(self, meta: Meta, message: BaseMessage, *, dedupe: bool = False) -> None: if dedupe and message.id in self.seen: @@ -155,6 +158,10 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): if not isinstance(chunk, ChatGenerationChunk): return if meta := self.metadata.get(run_id): + if self.dedupe_metadata and run_id in self.emitted_metadata: + meta = (meta[0], None) + else: + self.emitted_metadata.add(run_id) self._emit(meta, chunk.message) def on_llm_end( @@ -170,6 +177,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): gen = response.generations[0][0] if isinstance(gen, ChatGeneration): self._emit(meta, gen.message, dedupe=True) + self.emitted_metadata.discard(run_id) self.metadata.pop(run_id, None) def on_llm_error( @@ -180,6 +188,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): parent_run_id: UUID | None = None, **kwargs: Any, ) -> Any: + self.emitted_metadata.discard(run_id) self.metadata.pop(run_id, None) def on_chain_start( diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 3ecefc1dc..5bf5b4337 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -2614,6 +2614,7 @@ class Pregel( stream.put, subgraphs, parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None, + dedupe_metadata=version == "v2", ) ) @@ -2965,6 +2966,7 @@ class Pregel( stream_put, subgraphs, parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None, + dedupe_metadata=version == "v2", ) ) diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index 973768a4a..6c1e18a08 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -71,6 +71,8 @@ logger = logging.getLogger(__name__) __all__ = ("RemoteGraph", "RemoteException") +_STREAM_PROTOCOL_CONFIG_KEY = "__stream_protocol_version__" + _CONF_DROPLIST = frozenset( ( CONFIG_KEY_CHECKPOINT_MAP, @@ -109,6 +111,45 @@ class RemoteException(Exception): pass +def _restore_message_metadata( + data: Any, metadata_by_message_id: dict[str, dict[str, Any]] +) -> Any: + """Restore deduplicated message metadata using the message id as cache key.""" + if not (isinstance(data, list) and len(data) == 2): + return data + message, metadata = data + if not isinstance(message, dict): + return data + message_id = message.get("id") + if not isinstance(message_id, str): + return data + if isinstance(metadata, dict): + metadata_by_message_id[message_id] = metadata + return (message, metadata) + return (message, metadata_by_message_id.get(message_id)) + + +def _merge_values_patch( + ns: tuple[str, ...], + mode: str, + data: Any, + values_by_ns: dict[tuple[str, ...], dict[str, Any]], +) -> tuple[str, Any]: + """Merge `values-patch` events back into full values snapshots.""" + if mode != "values-patch" or not isinstance(data, dict): + return mode, data + values = data.get("values") + if not isinstance(values, dict): + return "values", values if values is not None else {} + merged = dict(values_by_ns.get(ns, {})) + merged.update(values) + for key in data.get("deleted_keys", ()): + if isinstance(key, str): + merged.pop(key, None) + values_by_ns[ns] = merged + return "values", merged + + class RemoteGraph(PregelProtocol): """The `RemoteGraph` class is a client implementation for calling remote APIs that implement the LangGraph Server API specification. @@ -760,12 +801,25 @@ class RemoteGraph(PregelProtocol): stream_modes, requested, req_single, stream = self._get_stream_modes( stream_mode, config ) + stream_protocol_version = cast( + Literal["v1", "v2"] | None, + kwargs.pop( + "stream_protocol_version", + sanitized_config.get("configurable", {}).get( + _STREAM_PROTOCOL_CONFIG_KEY + ), + ), + ) + if version == "v2" and stream_protocol_version is None: + stream_protocol_version = "v2" if isinstance(input, Command): command: CommandSDK | None = cast(CommandSDK, asdict(input)) input = None else: command = None thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None) + message_metadata_by_id: dict[str, dict[str, Any]] = {} + values_by_ns: dict[tuple[str, ...], dict[str, Any]] = {} for chunk in sync_client.runs.stream( thread_id=thread_id, @@ -783,6 +837,7 @@ class RemoteGraph(PregelProtocol): _merge_tracing_headers(headers) if self.distributed_tracing else headers ), params=params, + stream_protocol_version=stream_protocol_version, **kwargs, ): # split mode and ns @@ -798,6 +853,11 @@ class RemoteGraph(PregelProtocol): if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS): caller_ns = tuple(caller_ns.split(NS_SEP)) ns = caller_ns + ns + mode, data = _merge_values_patch(ns, mode, chunk.data, values_by_ns) + if mode != chunk.event: + chunk = chunk._replace(data=data) + elif data is not chunk.data: + chunk = chunk._replace(data=data) # stream to parent stream if stream is not None and mode in stream.modes: stream((ns, mode, chunk.data)) @@ -815,7 +875,9 @@ class RemoteGraph(PregelProtocol): continue if chunk.event.startswith("messages"): - chunk = chunk._replace(data=tuple(chunk.data)) + chunk = chunk._replace( + data=_restore_message_metadata(chunk.data, message_metadata_by_id) + ) # emit chunk if version == "v2": @@ -827,11 +889,6 @@ class RemoteGraph(PregelProtocol): ) yield {"type": mode, "ns": ns, "data": chunk.data, "interrupts": ints} elif subgraphs: - if NS_SEP in chunk.event: - mode, ns_ = chunk.event.split(NS_SEP, 1) - ns = tuple(ns_.split(NS_SEP)) - else: - mode, ns = chunk.event, () if req_single: yield ns, chunk.data else: @@ -915,12 +972,25 @@ class RemoteGraph(PregelProtocol): stream_modes, requested, req_single, stream = self._get_stream_modes( stream_mode, config ) + stream_protocol_version = cast( + Literal["v1", "v2"] | None, + kwargs.pop( + "stream_protocol_version", + sanitized_config.get("configurable", {}).get( + _STREAM_PROTOCOL_CONFIG_KEY + ), + ), + ) + if version == "v2" and stream_protocol_version is None: + stream_protocol_version = "v2" if isinstance(input, Command): command: CommandSDK | None = cast(CommandSDK, asdict(input)) input = None else: command = None thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None) + message_metadata_by_id: dict[str, dict[str, Any]] = {} + values_by_ns: dict[tuple[str, ...], dict[str, Any]] = {} async for chunk in client.runs.stream( thread_id=thread_id, @@ -938,6 +1008,7 @@ class RemoteGraph(PregelProtocol): _merge_tracing_headers(headers) if self.distributed_tracing else headers ), params=params, + stream_protocol_version=stream_protocol_version, **kwargs, ): # split mode and ns @@ -953,6 +1024,11 @@ class RemoteGraph(PregelProtocol): if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS): caller_ns = tuple(caller_ns.split(NS_SEP)) ns = caller_ns + ns + mode, data = _merge_values_patch(ns, mode, chunk.data, values_by_ns) + if mode != chunk.event: + chunk = chunk._replace(data=data) + elif data is not chunk.data: + chunk = chunk._replace(data=data) # stream to parent stream if stream is not None and mode in stream.modes: stream((ns, mode, chunk.data)) @@ -970,7 +1046,9 @@ class RemoteGraph(PregelProtocol): continue if chunk.event.startswith("messages"): - chunk = chunk._replace(data=tuple(chunk.data)) + chunk = chunk._replace( + data=_restore_message_metadata(chunk.data, message_metadata_by_id) + ) # emit chunk if version == "v2": @@ -982,11 +1060,6 @@ class RemoteGraph(PregelProtocol): ) yield {"type": mode, "ns": ns, "data": chunk.data, "interrupts": ints} elif subgraphs: - if NS_SEP in chunk.event: - mode, ns_ = chunk.event.split(NS_SEP, 1) - ns = tuple(ns_.split(NS_SEP)) - else: - mode, ns = chunk.event, () if req_single: yield ns, chunk.data else: diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index d04d82da7..980bc07df 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -275,13 +275,14 @@ class MessagesStreamPart(TypedDict): """Stream part emitted for `stream_mode="messages"`. `data` is a 2-tuple of `(message, metadata)` where `message` is a - `BaseMessage` (e.g. `AIMessageChunk`) and `metadata` is a dict containing - keys like `langgraph_step`, `langgraph_node`, `langgraph_triggers`, etc. + `BaseMessage` (e.g. `AIMessageChunk`) and `metadata` is either a dict containing + keys like `langgraph_step`, `langgraph_node`, `langgraph_triggers`, etc. or + `None` for deduplicated follow-up chunks in `version="v2"` streams. """ type: Literal["messages"] ns: tuple[str, ...] - data: tuple[AnyMessage, dict[str, Any]] + data: tuple[AnyMessage, dict[str, Any] | None] class CustomStreamPart(TypedDict): @@ -346,7 +347,7 @@ async for part in graph.astream(input, version="v2"): if part["type"] == "values": part["data"] # OutputT — full state (pydantic/dataclass/dict) elif part["type"] == "messages": - part["data"] # tuple[BaseMessage, dict] — (message, metadata) + part["data"] # tuple[BaseMessage, dict | None] — (message, metadata) elif part["type"] == "custom": part["data"] # Any — user-defined ``` diff --git a/libs/langgraph/tests/test_remote_graph.py b/libs/langgraph/tests/test_remote_graph.py index 4edb9325e..6d930dbfe 100644 --- a/libs/langgraph/tests/test_remote_graph.py +++ b/libs/langgraph/tests/test_remote_graph.py @@ -882,6 +882,75 @@ def test_stream_sanitizes_thread_id(): assert not passed_config["configurable"] +def test_stream_restores_messages_and_merges_values_patch(): + mock_sync_client = MagicMock() + mock_sync_client.runs.stream.return_value = [ + StreamPart( + event="messages|tools:call_1", + data=[ + {"id": "msg-1", "type": "AIMessageChunk", "content": "hel"}, + { + "langgraph_checkpoint_ns": "tools:call_1", + "langgraph_node": "agent", + }, + ], + ), + StreamPart( + event="messages|tools:call_1", + data=[ + {"id": "msg-1", "type": "AIMessageChunk", "content": "lo"}, + None, + ], + ), + StreamPart( + event="values|tools:call_1", + data={"messages": [{"type": "human", "content": "hi"}], "count": 1}, + ), + StreamPart( + event="values-patch|tools:call_1", + data={"values": {"count": 2}, "deleted_keys": ["messages"]}, + ), + ] + remote_pregel = RemoteGraph("test_graph_id", sync_client=mock_sync_client) + + parts = list( + remote_pregel.stream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode=["messages", "values"], + subgraphs=True, + version="v2", + ) + ) + + message_parts = [part for part in parts if part["type"] == "messages"] + assert message_parts[0]["data"][1] == { + "langgraph_checkpoint_ns": "tools:call_1", + "langgraph_node": "agent", + } + assert message_parts[1]["data"][1] == { + "langgraph_checkpoint_ns": "tools:call_1", + "langgraph_node": "agent", + } + value_parts = [part for part in parts if part["type"] == "values"] + assert value_parts == [ + { + "type": "values", + "ns": ("tools:call_1",), + "data": {"messages": [{"type": "human", "content": "hi"}], "count": 1}, + "interrupts": (), + }, + { + "type": "values", + "ns": ("tools:call_1",), + "data": {"count": 2}, + "interrupts": (), + }, + ] + _, kwargs = mock_sync_client.runs.stream.call_args + assert kwargs["stream_protocol_version"] == "v2" + + @pytest.mark.anyio async def test_ainvoke(): # set up test @@ -1092,6 +1161,78 @@ def test_stream_context_base_model(): assert kwargs["context"] == ctx +@pytest.mark.anyio +async def test_astream_restores_messages_and_merges_values_patch(): + mock_async_client = MagicMock() + async_iter = MagicMock() + async_iter.__aiter__.return_value = [ + StreamPart( + event="messages|tools:call_1", + data=[ + {"id": "msg-1", "type": "AIMessageChunk", "content": "hel"}, + { + "langgraph_checkpoint_ns": "tools:call_1", + "langgraph_node": "agent", + }, + ], + ), + StreamPart( + event="messages|tools:call_1", + data=[ + {"id": "msg-1", "type": "AIMessageChunk", "content": "lo"}, + None, + ], + ), + StreamPart( + event="values|tools:call_1", + data={"messages": [{"type": "human", "content": "hi"}], "count": 1}, + ), + StreamPart( + event="values-patch|tools:call_1", + data={"values": {"count": 2}, "deleted_keys": ["messages"]}, + ), + ] + mock_async_client.runs.stream.return_value = async_iter + remote_pregel = RemoteGraph("test_graph_id", client=mock_async_client) + + parts = [] + async for part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode=["messages", "values"], + subgraphs=True, + version="v2", + ): + parts.append(part) + + message_parts = [part for part in parts if part["type"] == "messages"] + assert message_parts[0]["data"][1] == { + "langgraph_checkpoint_ns": "tools:call_1", + "langgraph_node": "agent", + } + assert message_parts[1]["data"][1] == { + "langgraph_checkpoint_ns": "tools:call_1", + "langgraph_node": "agent", + } + value_parts = [part for part in parts if part["type"] == "values"] + assert value_parts == [ + { + "type": "values", + "ns": ("tools:call_1",), + "data": {"messages": [{"type": "human", "content": "hi"}], "count": 1}, + "interrupts": (), + }, + { + "type": "values", + "ns": ("tools:call_1",), + "data": {"count": 2}, + "interrupts": (), + }, + ] + _, kwargs = mock_async_client.runs.stream.call_args + assert kwargs["stream_protocol_version"] == "v2" + + @pytest.mark.skip( "Unskip this test to manually test the LangSmith Deployment integration" ) diff --git a/libs/langgraph/tests/test_stream_v2.py b/libs/langgraph/tests/test_stream_v2.py index 29d34e678..f2fdd9b15 100644 --- a/libs/langgraph/tests/test_stream_v2.py +++ b/libs/langgraph/tests/test_stream_v2.py @@ -90,6 +90,25 @@ def _make_messages_graph() -> StateGraph[ return builder +def _make_streaming_messages_graph() -> StateGraph[ + MessagesState, None, MessagesState, MessagesState +]: + model = FakeChatModel(messages=[AIMessage(content="hello world", id="ai-1")]) + + def call_model(state: MessagesState) -> dict[str, Any]: + streamed = model.stream(state["messages"]) + message = next(streamed) + for chunk in streamed: + message += chunk + return {"messages": message} + + builder = StateGraph(MessagesState, input_schema=MessagesState) + builder.add_node("call_model", call_model) + builder.add_edge(START, "call_model") + builder.add_edge("call_model", END) + return builder + + def _make_custom_graph() -> Any: @entrypoint() def graph(inputs: Any, *, writer: StreamWriter) -> Any: @@ -164,6 +183,13 @@ class TestV1BackwardsCompat: ns, _data = chunk assert isinstance(ns, tuple) + def test_stream_v1_messages_keep_metadata_on_every_chunk(self) -> None: + graph = _make_streaming_messages_graph().compile() + chunks = list(graph.stream(_MSG_INPUT, stream_mode="messages")) + metadata = [meta for _message, meta in chunks] + assert len(metadata) >= 3 + assert all(isinstance(meta, dict) for meta in metadata) + # --- v2 sync stream --- @@ -202,8 +228,22 @@ class TestV2Stream: assert isinstance(data, tuple) and len(data) == 2 message, metadata = data assert isinstance(message, BaseMessage) - assert isinstance(metadata, dict) - assert "langgraph_node" in metadata + assert metadata is None or isinstance(metadata, dict) + assert any(isinstance(c["data"][1], dict) for c in msg_chunks) + + def test_messages_streaming_dedupes_metadata(self) -> None: + graph = _make_streaming_messages_graph().compile() + chunks = list(graph.stream(_MSG_INPUT, stream_mode="messages", version="v2")) + msg_chunks = [c for c in chunks if c["type"] == "messages"] + assert len(msg_chunks) >= 3 + first_message, first_metadata = msg_chunks[0]["data"] + assert isinstance(first_message, BaseMessage) + assert isinstance(first_metadata, dict) + assert "langgraph_node" in first_metadata + for chunk in msg_chunks[1:]: + message, metadata = chunk["data"] + assert isinstance(message, BaseMessage) + assert metadata is None def test_custom(self) -> None: graph = _make_custom_graph() @@ -541,8 +581,28 @@ class TestV2StreamAsync: assert isinstance(data, tuple) and len(data) == 2 message, metadata = data assert isinstance(message, BaseMessage) - assert isinstance(metadata, dict) - assert "langgraph_node" in metadata + assert metadata is None or isinstance(metadata, dict) + assert any(isinstance(c["data"][1], dict) for c in msg_chunks) + + @pytest.mark.anyio + async def test_messages_streaming_dedupes_metadata(self) -> None: + graph = _make_streaming_messages_graph().compile() + chunks = [ + c + async for c in graph.astream( + _MSG_INPUT, stream_mode="messages", version="v2" + ) + ] + msg_chunks = [c for c in chunks if c["type"] == "messages"] + assert len(msg_chunks) >= 3 + first_message, first_metadata = msg_chunks[0]["data"] + assert isinstance(first_message, BaseMessage) + assert isinstance(first_metadata, dict) + assert "langgraph_node" in first_metadata + for chunk in msg_chunks[1:]: + message, metadata = chunk["data"] + assert isinstance(message, BaseMessage) + assert metadata is None @NEEDS_CONTEXTVARS @pytest.mark.anyio diff --git a/libs/sdk-py/langgraph_sdk/_async/runs.py b/libs/sdk-py/langgraph_sdk/_async/runs.py index c6d1fa162..05d65cf0e 100644 --- a/libs/sdk-py/langgraph_sdk/_async/runs.py +++ b/libs/sdk-py/langgraph_sdk/_async/runs.py @@ -85,6 +85,7 @@ class RunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, feedback_keys: Sequence[str] | None = None, @@ -124,6 +125,7 @@ class RunsClient: multitask_strategy: MultitaskStrategy | None = None, if_not_exists: IfNotExists | None = None, after_seconds: int | None = None, + stream_protocol_version: StreamVersion | None = None, headers: Mapping[str, str] | None = None, params: QueryParamTypes | None = None, on_run_created: Callable[[RunCreateMetadata], None] | None = None, @@ -144,6 +146,7 @@ class RunsClient: metadata: Mapping[str, Any] | None = None, config: Config | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, feedback_keys: Sequence[str] | None = None, @@ -180,6 +183,7 @@ class RunsClient: if_not_exists: IfNotExists | None = None, webhook: str | None = None, after_seconds: int | None = None, + stream_protocol_version: StreamVersion | None = None, headers: Mapping[str, str] | None = None, params: QueryParamTypes | None = None, on_run_created: Callable[[RunCreateMetadata], None] | None = None, @@ -211,6 +215,7 @@ class RunsClient: multitask_strategy: MultitaskStrategy | None = None, if_not_exists: IfNotExists | None = None, after_seconds: int | None = None, + stream_protocol_version: StreamVersion | None = None, headers: Mapping[str, str] | None = None, params: QueryParamTypes | None = None, on_run_created: Callable[[RunCreateMetadata], None] | None = None, @@ -312,6 +317,7 @@ class RunsClient: "stream_mode": stream_mode, "stream_subgraphs": stream_subgraphs, "stream_resumable": stream_resumable, + "stream_protocol_version": stream_protocol_version, "assistant_id": assistant_id, "interrupt_before": interrupt_before, "interrupt_after": interrupt_after, @@ -393,6 +399,7 @@ class RunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -420,6 +427,7 @@ class RunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, # deprecated + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -557,6 +565,7 @@ class RunsClient: "stream_mode": stream_mode, "stream_subgraphs": stream_subgraphs, "stream_resumable": stream_resumable, + "stream_protocol_version": stream_protocol_version, "config": config, "context": context, "metadata": metadata, @@ -644,6 +653,7 @@ class RunsClient: config: Config | None = None, context: Context | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -670,6 +680,7 @@ class RunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, # deprecated + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -785,6 +796,7 @@ class RunsClient: "config": config, "context": context, "metadata": metadata, + "stream_protocol_version": stream_protocol_version, "assistant_id": assistant_id, "interrupt_before": interrupt_before, "interrupt_after": interrupt_after, diff --git a/libs/sdk-py/langgraph_sdk/_sync/runs.py b/libs/sdk-py/langgraph_sdk/_sync/runs.py index 1b52b8113..54f3976e4 100644 --- a/libs/sdk-py/langgraph_sdk/_sync/runs.py +++ b/libs/sdk-py/langgraph_sdk/_sync/runs.py @@ -84,6 +84,7 @@ class SyncRunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, feedback_keys: Sequence[str] | None = None, @@ -114,6 +115,7 @@ class SyncRunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, feedback_keys: Sequence[str] | None = None, @@ -143,6 +145,7 @@ class SyncRunsClient: config: Config | None = None, context: Context | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, feedback_keys: Sequence[str] | None = None, @@ -172,6 +175,7 @@ class SyncRunsClient: config: Config | None = None, context: Context | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, feedback_keys: Sequence[str] | None = None, @@ -202,6 +206,7 @@ class SyncRunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, # deprecated + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, feedback_keys: Sequence[str] | None = None, @@ -307,6 +312,7 @@ class SyncRunsClient: "stream_mode": stream_mode, "stream_subgraphs": stream_subgraphs, "stream_resumable": stream_resumable, + "stream_protocol_version": stream_protocol_version, "assistant_id": assistant_id, "interrupt_before": interrupt_before, "interrupt_after": interrupt_after, @@ -360,6 +366,7 @@ class SyncRunsClient: config: Config | None = None, context: Context | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -388,6 +395,7 @@ class SyncRunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -415,6 +423,7 @@ class SyncRunsClient: checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, checkpoint_during: bool | None = None, # deprecated + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -552,6 +561,7 @@ class SyncRunsClient: "stream_mode": stream_mode, "stream_subgraphs": stream_subgraphs, "stream_resumable": stream_resumable, + "stream_protocol_version": stream_protocol_version, "config": config, "context": context, "metadata": metadata, @@ -665,6 +675,7 @@ class SyncRunsClient: checkpoint_during: bool | None = None, # deprecated checkpoint: Checkpoint | None = None, checkpoint_id: str | None = None, + stream_protocol_version: StreamVersion | None = None, interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, webhook: str | None = None, @@ -782,6 +793,7 @@ class SyncRunsClient: "config": config, "context": context, "metadata": metadata, + "stream_protocol_version": stream_protocol_version, "assistant_id": assistant_id, "interrupt_before": interrupt_before, "interrupt_after": interrupt_after, diff --git a/libs/sdk-py/langgraph_sdk/schema.py b/libs/sdk-py/langgraph_sdk/schema.py index 6fffcec34..ada89ac24 100644 --- a/libs/sdk-py/langgraph_sdk/schema.py +++ b/libs/sdk-py/langgraph_sdk/schema.py @@ -532,6 +532,8 @@ class RunCreate(TypedDict): """List of node names to interrupt execution after.""" webhook: str | None """URL to send webhook notifications about the run's progress.""" + stream_protocol_version: StreamVersion | None + """Opt into an alternate stream wire protocol version.""" multitask_strategy: MultitaskStrategy | None """Strategy for handling concurrent runs on the same thread.""" @@ -733,6 +735,26 @@ class ValuesStreamPart(TypedDict): """List of interrupts that occurred during this step.""" +class ValuesPatchPayload(TypedDict): + """Incremental patch payload for subgraph `values` updates.""" + + values: dict[str, Any] + """Only the changed fields since the previous `values` event for this namespace.""" + deleted_keys: NotRequired[list[str]] + """Optional list of keys that were removed from the previous values snapshot.""" + + +class ValuesPatchStreamPart(TypedDict): + """Stream part emitted for incremental subgraph value patches (`values-patch`).""" + + type: Literal["values-patch"] + """Stream part type discriminator.""" + ns: list[str] + """Namespace path of the emitting node (empty for root graph).""" + data: ValuesPatchPayload + """Incremental state patch for the namespace.""" + + class UpdatesStreamPart(TypedDict): """Stream part emitted for `stream_mode="updates"`.""" @@ -845,6 +867,7 @@ class MetadataStreamPart(TypedDict): StreamPartV2 = ( ValuesStreamPart + | ValuesPatchStreamPart | UpdatesStreamPart | MessagesPartialStreamPart | MessagesCompleteStreamPart diff --git a/libs/sdk-py/tests/test_client_stream.py b/libs/sdk-py/tests/test_client_stream.py index 9ac455f9f..ac4908daa 100644 --- a/libs/sdk-py/tests/test_client_stream.py +++ b/libs/sdk-py/tests/test_client_stream.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from collections.abc import Iterator, Sequence from pathlib import Path from typing import Any @@ -8,7 +9,9 @@ import httpx import pytest from typing_extensions import assert_type +from langgraph_sdk._async.runs import RunsClient from langgraph_sdk._shared.utilities import _sse_to_v2_dict +from langgraph_sdk._sync.runs import SyncRunsClient from langgraph_sdk.client import HttpClient, SyncHttpClient from langgraph_sdk.schema import ( CheckpointPayload, @@ -24,6 +27,7 @@ from langgraph_sdk.schema import ( TaskResultPayload, TasksStreamPart, UpdatesStreamPart, + ValuesPatchStreamPart, ValuesStreamPart, ) from langgraph_sdk.sse import BytesLike, BytesLineDecoder, SSEDecoder @@ -375,6 +379,81 @@ def test_sse_to_v2_dict_values_with_interrupts() -> None: assert "__interrupt__" not in result["data"] +def test_sse_to_v2_dict_values_patch() -> None: + payload = {"values": {"count": 2}, "deleted_keys": ["stale"]} + result = _sse_to_v2_dict("values-patch|tools:call_1", payload) + assert result is not None + _assert_v2_shape(result) + assert result == { + "type": "values-patch", + "ns": ["tools:call_1"], + "data": {"values": {"count": 2}, "deleted_keys": ["stale"]}, + "interrupts": [], + } + + +@pytest.mark.asyncio +async def test_async_runs_stream_includes_stream_protocol_version(): + async def handler(request: httpx.Request) -> httpx.Response: + assert request.method == "POST" + assert request.url.path == "/runs/stream" + body = json.loads(request.content) + assert body["stream_protocol_version"] == "v2" + return httpx.Response( + 200, + headers={"Content-Type": "text/event-stream"}, + content=b"event: end\ndata: null\n\n", + ) + + transport = httpx.MockTransport(handler) + async with httpx.AsyncClient( + transport=transport, base_url="https://example.com" + ) as client: + runs_client = RunsClient(HttpClient(client)) + parts = [ + part + async for part in runs_client.stream( + thread_id=None, + assistant_id="agent", + input={"messages": []}, + stream_protocol_version="v2", + ) + ] + + assert len(parts) == 1 + assert parts[0].event == "end" + assert parts[0].data is None + + +def test_sync_runs_stream_includes_stream_protocol_version(): + def handler(request: httpx.Request) -> httpx.Response: + assert request.method == "POST" + assert request.url.path == "/runs/stream" + body = json.loads(request.content) + assert body["stream_protocol_version"] == "v2" + return httpx.Response( + 200, + headers={"Content-Type": "text/event-stream"}, + content=b"event: end\ndata: null\n\n", + ) + + transport = httpx.MockTransport(handler) + with httpx.Client(transport=transport, base_url="https://example.com") as client: + runs_client = SyncRunsClient(SyncHttpClient(client)) + parts = list( + runs_client.stream( + thread_id=None, + assistant_id="agent", + input={"messages": []}, + stream_protocol_version="v2", + ) + ) + + assert len(parts) == 1 + assert parts[0].event == "end" + assert parts[0].data is None + + # --- client-side v2 stream wrapping --- @@ -448,6 +527,8 @@ def _check_v2_type_narrowing(part: StreamPartV2) -> None: if part["type"] == "values": assert_type(part, ValuesStreamPart) assert_type(part["data"], dict[str, Any]) + elif part["type"] == "values-patch": + assert_type(part, ValuesPatchStreamPart) elif part["type"] == "updates": assert_type(part, UpdatesStreamPart) assert_type(part["data"], dict[str, Any])