feat(langgraph): add protocol improvements for better streaming

This commit is contained in:
Christian Bromann
2026-03-12 15:08:25 -07:00
parent 93a0dfec08
commit f6286dce38
10 changed files with 435 additions and 21 deletions
+10 -1
View File
@@ -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(
+2
View File
@@ -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",
)
)
+85 -12
View File
@@ -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:
+5 -4
View File
@@ -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
```
+141
View File
@@ -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"
)
+64 -4
View File
@@ -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
+12
View File
@@ -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,
+12
View File
@@ -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,
+23
View File
@@ -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
+81
View File
@@ -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])