From f4a4c4b7b1e18403c41b073639ddf8ee0f916487 Mon Sep 17 00:00:00 2001 From: Sydney Runkle <54324534+sydney-runkle@users.noreply.github.com> Date: Tue, 3 Mar 2026 12:46:06 -0500 Subject: [PATCH] feat(langgraph): more robust pydantic + dataclass support for `StateGraph` (#6963) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## More robust Pydantic support for v2 streaming When using `stream_version="v2"`, stream data and invoke results now respect the graph's output/state schema types (Pydantic models, dataclasses, etc.) instead of always returning raw dicts. This makes working with typed state much more natural — no more manual `Model(**chunk)` calls scattered through your code. ### Values stream coercion `values` stream parts coerce data through the graph's output schema mapper, so you get Pydantic models (or dataclasses) back directly: ```python class MyState(BaseModel): value: str items: Annotated[list[str], operator.add] graph = StateGraph(MyState).compile() # v1: you get raw dicts back, have to reconstruct manually for chunk in graph.stream(inputs, stream_mode="values"): state = MyState(**chunk) # manual, error-prone # v2: data is already a MyState instance for part in graph.stream(inputs, stream_mode="values", stream_version="v2"): assert isinstance(part["data"], MyState) # just works print(part["data"].value) # attribute access, IDE autocomplete ``` This also works for dataclass-based state schemas. TypedDict state stays as plain dicts (no change needed). ### Interrupts on stream parts `values` stream parts now carry an `interrupts` field directly, removing the need to cross-reference the `updates` stream: ```python for part in graph.stream(inputs, config, stream_mode="values", stream_version="v2"): if part["interrupts"]: # handle interrupts inline — no need to check updates stream for intr in part["interrupts"]: print(intr.value) ``` ### Checkpoint/debug coercion Checkpoint and debug stream payloads also coerce their `values` through the state schema mapper, so `stream_mode="checkpoints"` and `stream_mode="debug"` return typed state too. ### Generic stream types `StreamPart`, `ValuesStreamPart`, `CheckpointPayload`, etc. are now generic over `StateT`/`OutputT`, enabling better static type checking across the board. ### `GraphOutput` wrapper This adds a new return type to `invoke()` which is a meaningful API surface change. `invoke(stream_version="v2")` returns a `GraphOutput[OutputT]` dataclass with `.value` and `.interrupts` fields: ```python result = graph.invoke({"value": "x", "items": []}, stream_version="v2") # typed access assert isinstance(result, GraphOutput) assert isinstance(result.value, MyState) # coerced to schema type assert result.interrupts == () # always available # backward compat dict access still works assert result["value"] == "x_a" ``` The concern: this changes the return type of `invoke()` in a way that existing code patterns like `result["key"]` still work (via `__getitem__`), but `isinstance(result, dict)` checks would break. Worth discussing whether the ergonomic benefit justifies the migration cost. --- libs/langgraph/langgraph/graph/state.py | 29 +- libs/langgraph/langgraph/pregel/main.py | 230 +++++--- libs/langgraph/langgraph/pregel/protocol.py | 5 +- libs/langgraph/langgraph/pregel/remote.py | 15 +- libs/langgraph/langgraph/types.py | 74 ++- libs/langgraph/tests/test_stream_v2.py | 547 +++++++++++++++++++- 6 files changed, 811 insertions(+), 89 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index a87c17527..b1c24de2b 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -6,6 +6,7 @@ import typing import warnings from collections import defaultdict from collections.abc import Awaitable, Callable, Hashable, Sequence +from dataclasses import is_dataclass from functools import partial from inspect import isclass, isfunction, ismethod, signature from types import FunctionType @@ -14,6 +15,7 @@ from typing import ( Any, Generic, Literal, + TypeVar, Union, cast, get_args, @@ -1164,6 +1166,20 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): for key, node in self.nodes.items(): compiled.attach_node(key, node) + # Record output/state mappers for v2 stream coercion (pydantic/dataclass only) + compiled._output_mapper = _pick_mapper( + list(output_channels) + if isinstance(output_channels, list) + else [output_channels], + self.output_schema, + ) + compiled._state_mapper = _pick_mapper( + list(stream_channels) + if isinstance(stream_channels, list) + else [stream_channels], + self.state_schema, + ) + for start, end in self.edges: compiled.attach_edge(start, end) @@ -1183,6 +1199,8 @@ class CompiledStateGraph( ): builder: StateGraph[StateT, ContextT, InputT, OutputT] schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None] + _output_mapper: Callable[[Any], Any] | None + _state_mapper: Callable[[Any], Any] | None def __init__( self, @@ -1504,12 +1522,15 @@ def _pick_mapper( ) -> Callable[[Any], Any] | None: if state_keys == ["__root__"]: return None - if isclass(schema) and issubclass(schema, dict): - return None - return partial(_coerce_state, schema) + if isclass(schema) and (issubclass(schema, BaseModel) or is_dataclass(schema)): + return partial(_coerce_state, schema) + return None -def _coerce_state(schema: type[Any], input: dict[str, Any]) -> dict[str, Any]: +_S = TypeVar("_S") + + +def _coerce_state(schema: type[_S], input: dict[str, Any]) -> _S: return schema(**input) diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 597a60922..b3bfb2d9d 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -123,7 +123,10 @@ from langgraph.pregel._checkpoint import ( ) from langgraph.pregel._draw import draw_graph from langgraph.pregel._io import map_input, read_channels -from langgraph.pregel._loop import AsyncPregelLoop, SyncPregelLoop +from langgraph.pregel._loop import ( + AsyncPregelLoop, + SyncPregelLoop, +) from langgraph.pregel._messages import StreamMessagesHandler from langgraph.pregel._read import DEFAULT_BOUND, PregelNode from langgraph.pregel._retry import RetryPolicy @@ -140,6 +143,7 @@ from langgraph.types import ( Checkpointer, Command, Durability, + GraphOutput, Interrupt, Send, StateSnapshot, @@ -996,6 +1000,11 @@ class Pregel( for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse): yield name, node + # Mappers for v2 stream coercion (pydantic/dataclass). + # Set by CompiledStateGraph; None for base Pregel. + _output_mapper: Callable[[Any], Any] | None = None + _state_mapper: Callable[[Any], Any] | None = None + def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None: """Migrate a saved checkpoint to new channel layout.""" if checkpoint["v"] < 4 and checkpoint.get("pending_sends"): @@ -2644,6 +2653,10 @@ class Pregel( runtime = parent_runtime.merge(runtime) config[CONF][CONFIG_KEY_RUNTIME] = runtime + # resolve mappers for v2 stream coercion + _output_mapper = self._output_mapper if stream_version == "v2" else None + _state_mapper = self._state_mapper if stream_version == "v2" else None + with SyncPregelLoop( input, stream=StreamProtocol(stream.put, stream_modes), @@ -2722,6 +2735,8 @@ class Pregel( stream.get, queue.Empty, stream_version, + _output_mapper, + _state_mapper, ) loop.after_tick() # wait for checkpoint @@ -2735,6 +2750,8 @@ class Pregel( stream.get, queue.Empty, stream_version, + _output_mapper, + _state_mapper, ) # handle exit if loop.status == "out_of_steps": @@ -3002,6 +3019,10 @@ class Pregel( runtime = parent_runtime.merge(runtime) config[CONF][CONFIG_KEY_RUNTIME] = runtime + # resolve mappers for v2 stream coercion + _output_mapper = self._output_mapper if stream_version == "v2" else None + _state_mapper = self._state_mapper if stream_version == "v2" else None + async with AsyncPregelLoop( input, stream=StreamProtocol(stream.put_nowait, stream_modes), @@ -3099,6 +3120,8 @@ class Pregel( stream.get_nowait, asyncio.QueueEmpty, stream_version, + _output_mapper, + _state_mapper, ): yield o loop.after_tick() @@ -3118,6 +3141,8 @@ class Pregel( stream.get_nowait, asyncio.QueueEmpty, stream_version, + _output_mapper, + _state_mapper, ): yield o # handle exit @@ -3152,7 +3177,7 @@ class Pregel( durability: Durability | None = None, stream_version: Literal["v2"], **kwargs: Any, - ) -> dict[str, Any]: ... + ) -> GraphOutput[OutputT]: ... @overload def invoke( @@ -3239,44 +3264,64 @@ class Pregel( chunks: list[dict[str, Any] | Any] = [] interrupts: list[Interrupt] = [] - for chunk in self.stream( # type: ignore[misc] - input, - config, - context=context, - stream_mode=( - ["updates", "values"] if stream_mode == "values" else stream_mode - ), - print_mode=print_mode, - output_keys=output_keys, - interrupt_before=interrupt_before, - interrupt_after=interrupt_after, - durability=durability, - stream_version=stream_version, # type: ignore[arg-type] - **kwargs, - ): - if stream_mode == "values": - if stream_version == "v2": - mode = chunk["type"] - payload = chunk["data"] + if stream_version == "v2": + # v2: values stream parts carry interrupts directly + for chunk in self.stream( + input, + config, + context=context, + stream_mode="values" if stream_mode == "values" else stream_mode, + print_mode=print_mode, + output_keys=output_keys, + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, + durability=durability, + stream_version=stream_version, + **kwargs, + ): + if stream_mode == "values": + latest = chunk["data"] + if chunk_ints := chunk.get("interrupts", ()): + interrupts.extend(chunk_ints) # type: ignore[arg-type] else: + chunks.append(chunk) + else: + # v1: collect interrupts from updates stream + for chunk in self.stream( + input, + config, + context=context, + stream_mode=( + ["updates", "values"] if stream_mode == "values" else stream_mode + ), + print_mode=print_mode, + output_keys=output_keys, + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, + durability=durability, + **kwargs, + ): + if stream_mode == "values": if len(chunk) == 2: mode, payload = cast(tuple[StreamMode, Any], chunk) else: _, mode, payload = cast( tuple[tuple[str, ...], StreamMode, Any], chunk ) - if ( - mode == "updates" - and isinstance(payload, dict) - and (ints := payload.get(INTERRUPT)) is not None - ): - interrupts.extend(ints) # type: ignore[arg-type] - elif mode == "values": - latest = payload - else: - chunks.append(chunk) + if ( + mode == "updates" + and isinstance(payload, dict) + and (ints := payload.get(INTERRUPT)) is not None + ): + interrupts.extend(ints) + elif mode == "values": + latest = payload + else: + chunks.append(chunk) if stream_mode == "values": + if stream_version == "v2": + return GraphOutput(value=latest, interrupts=tuple(interrupts)) if interrupts: return ( {**latest, INTERRUPT: interrupts} @@ -3302,7 +3347,7 @@ class Pregel( durability: Durability | None = None, stream_version: Literal["v2"], **kwargs: Any, - ) -> dict[str, Any]: ... + ) -> GraphOutput[OutputT]: ... @overload async def ainvoke( @@ -3389,44 +3434,64 @@ class Pregel( chunks: list[dict[str, Any] | Any] = [] interrupts: list[Interrupt] = [] - async for chunk in self.astream( # type: ignore[misc] - input, - config, - context=context, - stream_mode=( - ["updates", "values"] if stream_mode == "values" else stream_mode - ), - print_mode=print_mode, - output_keys=output_keys, - interrupt_before=interrupt_before, - interrupt_after=interrupt_after, - durability=durability, - stream_version=stream_version, # type: ignore[arg-type] - **kwargs, - ): - if stream_mode == "values": - if stream_version == "v2": - mode = chunk["type"] - payload = chunk["data"] + if stream_version == "v2": + # v2: values stream parts carry interrupts directly + async for chunk in self.astream( + input, + config, + context=context, + stream_mode="values" if stream_mode == "values" else stream_mode, + print_mode=print_mode, + output_keys=output_keys, + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, + durability=durability, + stream_version=stream_version, + **kwargs, + ): + if stream_mode == "values": + latest = chunk["data"] + if chunk_ints := chunk.get("interrupts", ()): + interrupts.extend(chunk_ints) # type: ignore[arg-type] else: + chunks.append(chunk) + else: + # v1: collect interrupts from updates stream + async for chunk in self.astream( + input, + config, + context=context, + stream_mode=( + ["updates", "values"] if stream_mode == "values" else stream_mode + ), + print_mode=print_mode, + output_keys=output_keys, + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, + durability=durability, + **kwargs, + ): + if stream_mode == "values": if len(chunk) == 2: mode, payload = cast(tuple[StreamMode, Any], chunk) else: _, mode, payload = cast( tuple[tuple[str, ...], StreamMode, Any], chunk ) - if ( - mode == "updates" - and isinstance(payload, dict) - and (ints := payload.get(INTERRUPT)) is not None - ): - interrupts.extend(ints) # type: ignore[arg-type] - elif mode == "values": - latest = payload - else: - chunks.append(chunk) + if ( + mode == "updates" + and isinstance(payload, dict) + and (ints := payload.get(INTERRUPT)) is not None + ): + interrupts.extend(ints) + elif mode == "values": + latest = payload + else: + chunks.append(chunk) if stream_mode == "values": + if stream_version == "v2": + return GraphOutput(value=latest, interrupts=tuple(interrupts)) if interrupts: return ( {**latest, INTERRUPT: interrupts} @@ -3492,6 +3557,8 @@ def _output( getter: Callable[[], tuple[tuple[str, ...], str, Any]], empty_exc: type[Exception], stream_version: Literal["v1", "v2"] = "v1", + output_mapper: Callable[[Any], Any] | None = None, + state_mapper: Callable[[Any], Any] | None = None, ) -> Iterator: while True: try: @@ -3520,7 +3587,21 @@ def _output( ) if mode in stream_mode: if stream_version == "v2": - yield {"type": mode, "ns": ns, "data": payload} + if mode == "values": + # pop __interrupt__ into typed field, coerce data + ints: tuple[Interrupt, ...] = () + if isinstance(payload, dict): + ints = payload.pop(INTERRUPT, ()) + if output_mapper: + payload = output_mapper(payload) + yield {"type": mode, "ns": ns, "data": payload, "interrupts": ints} + elif mode in ("checkpoints", "debug"): + # coerce state values in checkpoint/debug payloads + if state_mapper: + _coerce_checkpoint_values(payload, state_mapper) + yield {"type": mode, "ns": ns, "data": payload} + else: + yield {"type": mode, "ns": ns, "data": payload} elif stream_subgraphs and isinstance(stream_mode, list): yield (ns, mode, payload) elif isinstance(stream_mode, list): @@ -3531,6 +3612,31 @@ def _output( yield payload +def _coerce_checkpoint_values(payload: Any, mapper: Callable[[Any], Any]) -> None: + """Coerce `values` dicts inside checkpoint or debug payloads in-place. + + Skips the initial checkpoint (where next contains ``__start__``) because + not all channels are populated yet and coercion would fail. + """ + _START = "__start__" + # debug wrapper: {"type": "checkpoint", "payload": {"values": dict, ...}} + if ( + isinstance(payload, dict) + and payload.get("type") == "checkpoint" + and isinstance(payload.get("payload"), dict) + and isinstance(payload["payload"].get("values"), dict) + and _START not in payload["payload"].get("next", ()) + ): + payload["payload"]["values"] = mapper(payload["payload"]["values"]) + # direct checkpoint payload: {"values": dict, ...} + elif ( + isinstance(payload, dict) + and isinstance(payload.get("values"), dict) + and _START not in payload.get("next", ()) + ): + payload["values"] = mapper(payload["values"]) + + def _coerce_context( context_schema: type[ContextT] | None, context: Any ) -> ContextT | None: diff --git a/libs/langgraph/langgraph/pregel/protocol.py b/libs/langgraph/langgraph/pregel/protocol.py index fb39c76c2..f5880455c 100644 --- a/libs/langgraph/langgraph/pregel/protocol.py +++ b/libs/langgraph/langgraph/pregel/protocol.py @@ -11,6 +11,7 @@ from typing_extensions import Self from langgraph.types import ( All, Command, + GraphOutput, StateSnapshot, StateUpdate, StreamMode, @@ -202,7 +203,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, stream_version: Literal["v2"], - ) -> dict[str, Any]: ... + ) -> GraphOutput[OutputT]: ... @overload @abstractmethod @@ -240,7 +241,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou interrupt_before: All | Sequence[str] | None = None, interrupt_after: All | Sequence[str] | None = None, stream_version: Literal["v2"], - ) -> dict[str, Any]: ... + ) -> GraphOutput[OutputT]: ... @overload @abstractmethod diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index 691c9646d..54d119485 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -58,6 +58,7 @@ from langgraph.pregel.protocol import PregelProtocol, StreamProtocol from langgraph.types import ( All, Command, + GraphOutput, Interrupt, PregelTask, StateSnapshot, @@ -1002,7 +1003,7 @@ class RemoteGraph(PregelProtocol): params: QueryParamTypes | None = None, stream_version: Literal["v2"], **kwargs: Any, - ) -> dict[str, Any]: ... + ) -> GraphOutput[dict[str, Any]]: ... @overload def invoke( @@ -1059,7 +1060,10 @@ class RemoteGraph(PregelProtocol): pass try: if stream_version == "v2": - return chunk["data"] + return GraphOutput( + value=chunk["data"], + interrupts=tuple(chunk.get("interrupts", ())), + ) return chunk except UnboundLocalError: logger.warning("No events received from remote graph") @@ -1077,7 +1081,7 @@ class RemoteGraph(PregelProtocol): params: QueryParamTypes | None = None, stream_version: Literal["v2"], **kwargs: Any, - ) -> dict[str, Any]: ... + ) -> GraphOutput[dict[str, Any]]: ... @overload async def ainvoke( @@ -1134,7 +1138,10 @@ class RemoteGraph(PregelProtocol): pass try: if stream_version == "v2": - return chunk["data"] + return GraphOutput( + value=chunk["data"], + interrupts=tuple(chunk.get("interrupts", ())), + ) return chunk except UnboundLocalError: logger.warning("No events received from remote graph") diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 1ed24feaa..d00eb3c61 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -23,11 +23,19 @@ from typing_extensions import NotRequired, TypedDict, Unpack, deprecated from xxhash import xxh3_128_hexdigest from langgraph._internal._cache import default_cache_key +from langgraph._internal._constants import INTERRUPT as _INTERRUPT_KEY from langgraph._internal._fields import get_cached_annotated_keys, get_update_as_tuples from langgraph._internal._retry import default_retry_on from langgraph._internal._typing import MISSING, DeprecatedKwargs from langgraph.warnings import LangGraphDeprecatedSinceV10 +# Local TypeVars for generic stream TypedDicts. +# We use separate TypeVars here (rather than importing from langgraph.typing) +# because the typing module TypeVars have defaults that cause mypy issues +# when used in standalone type aliases. +StateT = TypeVar("StateT") +OutputT = TypeVar("OutputT") + if TYPE_CHECKING: from langgraph.pregel.protocol import PregelProtocol @@ -70,6 +78,7 @@ __all__ = ( "Durability", "interrupt", "Overwrite", + "GraphOutput", "ensure_valid_checkpointer", ) @@ -165,22 +174,22 @@ class CheckpointTask(TypedDict): state: StateSnapshot | RunnableConfig | None -class CheckpointPayload(TypedDict): +class CheckpointPayload(TypedDict, Generic[StateT]): """Payload for a checkpoint event.""" config: RunnableConfig | None metadata: CheckpointMetadata - values: dict[str, Any] + values: StateT next: list[str] parent_config: RunnableConfig | None tasks: list[CheckpointTask] -class _DebugCheckpointPayload(TypedDict): +class _DebugCheckpointPayload(TypedDict, Generic[StateT]): step: int timestamp: str type: Literal["checkpoint"] - payload: CheckpointPayload + payload: CheckpointPayload[StateT] class _DebugTaskPayload(TypedDict): @@ -197,11 +206,13 @@ class _DebugTaskResultPayload(TypedDict): payload: TaskResultPayload -DebugPayload = _DebugCheckpointPayload | _DebugTaskPayload | _DebugTaskResultPayload +DebugPayload = ( + _DebugCheckpointPayload[StateT] | _DebugTaskPayload | _DebugTaskResultPayload +) """Wrapper payload for debug events. Discriminate on `type`.""" -class ValuesStreamPart(TypedDict): +class ValuesStreamPart(TypedDict, Generic[OutputT]): """Stream part emitted for `stream_mode="values"`. `data` contains the full state after each step, as returned by `read_channels()`. @@ -209,7 +220,8 @@ class ValuesStreamPart(TypedDict): type: Literal["values"] ns: tuple[str, ...] - data: dict[str, Any] + data: OutputT + interrupts: tuple[Interrupt, ...] class UpdatesStreamPart(TypedDict): @@ -248,12 +260,12 @@ class CustomStreamPart(TypedDict): data: Any -class CheckpointStreamPart(TypedDict): +class CheckpointStreamPart(TypedDict, Generic[StateT]): """Stream part emitted for `stream_mode="checkpoints"`.""" type: Literal["checkpoints"] ns: tuple[str, ...] - data: CheckpointPayload + data: CheckpointPayload[StateT] class TasksStreamPart(TypedDict): @@ -271,7 +283,7 @@ class TasksStreamPart(TypedDict): data: TaskPayload | TaskResultPayload -class DebugStreamPart(TypedDict): +class DebugStreamPart(TypedDict, Generic[StateT]): """Stream part emitted for `stream_mode="debug"`.""" type: Literal["debug"] @@ -280,22 +292,22 @@ class DebugStreamPart(TypedDict): StreamPart = ( - ValuesStreamPart + ValuesStreamPart[OutputT] | UpdatesStreamPart | MessagesStreamPart | CustomStreamPart - | CheckpointStreamPart + | CheckpointStreamPart[StateT] | TasksStreamPart - | DebugStreamPart + | DebugStreamPart[StateT] ) """A discriminated union of all v2 stream part types. Use `part["type"]` to narrow the type: ```python -async for part in graph.astream(input, version="v2"): +async for part in graph.astream(input, stream_version="v2"): if part["type"] == "values": - part["data"] # dict[str, Any] — full state + part["data"] # OutputT — full state (pydantic/dataclass/dict) elif part["type"] == "messages": part["data"] # tuple[BaseMessage, dict] — (message, metadata) elif part["type"] == "custom": @@ -303,6 +315,38 @@ async for part in graph.astream(input, version="v2"): ``` """ + +@dataclass(frozen=True) +class GraphOutput(Generic[OutputT]): + """Typed container returned by `invoke()` / `ainvoke()` with `stream_version="v2"`. + + Attributes: + value: The final output of the graph (dict, Pydantic model, dataclass, etc.). + interrupts: Any interrupts that occurred during execution. + """ + + value: OutputT + interrupts: tuple[Interrupt, ...] = () + + def __getitem__(self, key: str) -> Any: + """Backward compat: `result['__interrupt__']` and dict-key access.""" + if key == _INTERRUPT_KEY: + return self.interrupts + if isinstance(self.value, dict): + return self.value[key] + try: + return getattr(self.value, key) + except AttributeError: + raise KeyError(key) + + def __contains__(self, key: object) -> bool: + if key == _INTERRUPT_KEY: + return bool(self.interrupts) + if isinstance(self.value, dict): + return key in self.value + return isinstance(key, str) and hasattr(self.value, key) + + _DC_KWARGS = {"kw_only": True, "slots": True, "frozen": True} diff --git a/libs/langgraph/tests/test_stream_v2.py b/libs/langgraph/tests/test_stream_v2.py index 07b1974f1..772c00223 100644 --- a/libs/langgraph/tests/test_stream_v2.py +++ b/libs/langgraph/tests/test_stream_v2.py @@ -8,13 +8,16 @@ from __future__ import annotations import operator import sys +from dataclasses import dataclass from typing import Annotated, Any import pytest from langchain_core.messages import AIMessage, BaseMessage from langgraph.checkpoint.memory import InMemorySaver +from pydantic import BaseModel, ValidationError from typing_extensions import TypedDict, assert_type +from langgraph._internal._constants import INTERRUPT from langgraph.constants import END, START from langgraph.func import entrypoint from langgraph.graph import StateGraph @@ -25,6 +28,8 @@ from langgraph.types import ( CustomStreamPart, DebugPayload, DebugStreamPart, + GraphOutput, + Interrupt, MessagesStreamPart, StreamPart, StreamWriter, @@ -33,6 +38,7 @@ from langgraph.types import ( TasksStreamPart, UpdatesStreamPart, ValuesStreamPart, + interrupt, ) from tests.fake_chat import FakeChatModel @@ -118,6 +124,9 @@ def _assert_stream_part_shape(part: StreamPart) -> None: assert isinstance(part["ns"], tuple) for elem in part["ns"]: assert isinstance(elem, str) + if part["type"] == "values": + assert "interrupts" in part, "values stream part missing 'interrupts' field" + assert isinstance(part["interrupts"], tuple) # --- v1 backwards compatibility --- @@ -336,10 +345,119 @@ class TestV2Invoke: def test_values_default(self) -> None: graph = _make_simple_graph().compile() result = graph.invoke(_SIMPLE_INPUT, stream_version="v2") - assert isinstance(result, dict) + assert isinstance(result, GraphOutput) + assert result.value == {"value": "x_a_b", "items": ["a", "b"]} + assert result.interrupts == () + # backward compat dict access assert result["value"] == "x_a_b" assert result["items"] == ["a", "b"] + def test_invoke_v2_graph_output_with_interrupts(self) -> None: + def my_node(state: SimpleState) -> dict[str, Any]: + answer = interrupt("what is your name?") + return {"value": answer, "items": ["done"]} + + builder: StateGraph = StateGraph(SimpleState) + builder.add_node("my_node", my_node) + builder.add_edge(START, "my_node") + builder.add_edge("my_node", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-invoke-v2-interrupts"}} + result = graph.invoke({"value": "x", "items": []}, config, stream_version="v2") + assert isinstance(result, GraphOutput) + assert len(result.interrupts) > 0 + for intr in result.interrupts: + assert isinstance(intr, Interrupt) + # value should still be the state (not None or empty) + assert isinstance(result.value, dict) + + def test_invoke_v2_graph_output_interrupt_compat(self) -> None: + """result['__interrupt__'] works via __getitem__.""" + + def my_node(state: SimpleState) -> dict[str, Any]: + answer = interrupt("what is your name?") + return {"value": answer, "items": ["done"]} + + builder: StateGraph = StateGraph(SimpleState) + builder.add_node("my_node", my_node) + builder.add_edge(START, "my_node") + builder.add_edge("my_node", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-invoke-v2-compat"}} + result = graph.invoke({"value": "x", "items": []}, config, stream_version="v2") + assert isinstance(result, GraphOutput) + assert INTERRUPT in result + assert result[INTERRUPT] == result.interrupts + assert len(result[INTERRUPT]) > 0 + + def test_invoke_v2_graph_output_no_interrupts(self) -> None: + graph = _make_simple_graph().compile() + result = graph.invoke(_SIMPLE_INPUT, stream_version="v2") + assert isinstance(result, GraphOutput) + assert result.interrupts == () + assert INTERRUPT not in result + + def test_invoke_v2_pydantic_state(self) -> None: + """invoke with v2 and pydantic state returns GraphOutput with pydantic value.""" + + def node_a(state: PydanticState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile() + + result = graph.invoke({"value": "x", "items": []}, stream_version="v2") + assert isinstance(result, GraphOutput) + assert isinstance(result.value, PydanticState) + assert result.value.value == "x_a" + assert result.interrupts == () + + def test_invoke_v2_dataclass_state(self) -> None: + """invoke with v2 and dataclass state returns GraphOutput with dataclass value.""" + + def node_a(state: DataclassState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(DataclassState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile() + + result = graph.invoke({"value": "x", "items": []}, stream_version="v2") + assert isinstance(result, GraphOutput) + assert isinstance(result.value, DataclassState) + assert result.value.value == "x_a" + assert result.value.items == ["a"] + assert result.interrupts == () + + def test_invoke_v2_non_values_mode_pydantic(self) -> None: + """invoke with v2 + non-values mode + pydantic state returns list[StreamPart].""" + + def node_a(state: PydanticState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile() + + result = graph.invoke( + {"value": "x", "items": []}, stream_mode="updates", stream_version="v2" + ) + assert isinstance(result, list) + for chunk in result: + _assert_stream_part_shape(chunk) + assert chunk["type"] == "updates" + # updates data should be plain dicts, not coerced to pydantic + assert isinstance(chunk["data"], dict) + def test_updates_mode(self) -> None: graph = _make_simple_graph().compile() result = graph.invoke(_SIMPLE_INPUT, stream_mode="updates", stream_version="v2") @@ -562,10 +680,84 @@ class TestV2InvokeAsync: async def test_values_default(self) -> None: graph = _make_simple_graph().compile() result = await graph.ainvoke(_SIMPLE_INPUT, stream_version="v2") - assert isinstance(result, dict) + assert isinstance(result, GraphOutput) + assert result.value == {"value": "x_a_b", "items": ["a", "b"]} + assert result.interrupts == () + # backward compat dict access assert result["value"] == "x_a_b" assert result["items"] == ["a", "b"] + @NEEDS_CONTEXTVARS + @pytest.mark.anyio + async def test_ainvoke_v2_graph_output_with_interrupts(self) -> None: + def my_node(state: SimpleState) -> dict[str, Any]: + answer = interrupt("what is your name?") + return {"value": answer, "items": ["done"]} + + builder: StateGraph = StateGraph(SimpleState) + builder.add_node("my_node", my_node) + builder.add_edge(START, "my_node") + builder.add_edge("my_node", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-ainvoke-v2-interrupts"}} + result = await graph.ainvoke( + {"value": "x", "items": []}, config, stream_version="v2" + ) + assert isinstance(result, GraphOutput) + assert len(result.interrupts) > 0 + for intr in result.interrupts: + assert isinstance(intr, Interrupt) + assert isinstance(result.value, dict) + + @pytest.mark.anyio + async def test_ainvoke_v2_pydantic_state(self) -> None: + """ainvoke with v2 and pydantic state returns GraphOutput with pydantic value.""" + + def node_a(state: PydanticState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile() + + result = await graph.ainvoke({"value": "x", "items": []}, stream_version="v2") + assert isinstance(result, GraphOutput) + assert isinstance(result.value, PydanticState) + assert result.value.value == "x_a" + assert result.value.items == ["a"] + assert result.interrupts == () + + @pytest.mark.anyio + async def test_ainvoke_v2_dataclass_state(self) -> None: + """ainvoke with v2 and dataclass state returns GraphOutput with dataclass value.""" + + def node_a(state: DataclassState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(DataclassState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile() + + result = await graph.ainvoke({"value": "x", "items": []}, stream_version="v2") + assert isinstance(result, GraphOutput) + assert isinstance(result.value, DataclassState) + assert result.value.value == "x_a" + assert result.value.items == ["a"] + assert result.interrupts == () + + @pytest.mark.anyio + async def test_ainvoke_v2_graph_output_no_interrupts(self) -> None: + graph = _make_simple_graph().compile() + result = await graph.ainvoke(_SIMPLE_INPUT, stream_version="v2") + assert isinstance(result, GraphOutput) + assert result.interrupts == () + assert INTERRUPT not in result + @pytest.mark.anyio async def test_updates_mode(self) -> None: graph = _make_simple_graph().compile() @@ -594,6 +786,357 @@ class TestV2InvokeAsync: _assert_stream_part_shape(c) +# --- type-safe streaming: coercion + interrupt separation --- + + +class PydanticState(BaseModel): + value: str + items: Annotated[list[str], operator.add] + + +@dataclass +class DataclassState: + value: str + items: Annotated[list[str], operator.add] + + +class TestV2TypeSafeStreaming: + """Test that v2 streaming coerces values to pydantic/dataclass instances + and separates interrupts into a dedicated field.""" + + def test_values_pydantic_state(self) -> None: + """v2 values + pydantic state -> data is pydantic model instance.""" + + def node_a(state: PydanticState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile() + + chunks = list( + graph.stream( + {"value": "x", "items": []}, + stream_mode="values", + stream_version="v2", + ) + ) + assert len(chunks) >= 1 + for c in chunks: + _assert_stream_part_shape(c) + assert c["type"] == "values" + assert isinstance(c["data"], PydanticState), ( + f"Expected PydanticState, got {type(c['data'])}" + ) + assert c["interrupts"] == () + + def test_values_dataclass_state(self) -> None: + """v2 values + dataclass state -> data is dataclass instance.""" + + def node_a(state: DataclassState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(DataclassState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile() + + chunks = list( + graph.stream( + {"value": "x", "items": []}, + stream_mode="values", + stream_version="v2", + ) + ) + assert len(chunks) >= 1 + for c in chunks: + _assert_stream_part_shape(c) + assert c["type"] == "values" + assert isinstance(c["data"], DataclassState), ( + f"Expected DataclassState, got {type(c['data'])}" + ) + + def test_values_typeddict_state(self) -> None: + """v2 values + TypedDict state -> data stays plain dict (no coercion).""" + graph = _make_simple_graph().compile() + chunks = list( + graph.stream(_SIMPLE_INPUT, stream_mode="values", stream_version="v2") + ) + assert len(chunks) >= 1 + for c in chunks: + _assert_stream_part_shape(c) + assert c["type"] == "values" + # TypedDict state should remain a plain dict + assert isinstance(c["data"], dict) + assert type(c["data"]) is dict + + def test_values_interrupt_v2(self) -> None: + """v2 values + interrupt -> interrupts in typed field, not in data.""" + + def my_node(state: SimpleState) -> dict[str, Any]: + answer = interrupt("what is your name?") + return {"value": answer, "items": ["done"]} + + builder: StateGraph = StateGraph(SimpleState) + builder.add_node("my_node", my_node) + builder.add_edge(START, "my_node") + builder.add_edge("my_node", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-v2-interrupt"}} + chunks = list( + graph.stream( + {"value": "x", "items": []}, + config, + stream_mode="values", + stream_version="v2", + ) + ) + # should have at least one values chunk with interrupts + interrupt_chunks = [c for c in chunks if c.get("interrupts", ())] + assert len(interrupt_chunks) >= 1, f"Expected interrupt chunks, got {chunks}" + for c in interrupt_chunks: + assert c["type"] == "values" + assert isinstance(c["interrupts"], tuple) + assert len(c["interrupts"]) > 0 + for intr in c["interrupts"]: + assert isinstance(intr, Interrupt) + # __interrupt__ should NOT be in data + if isinstance(c["data"], dict): + assert INTERRUPT not in c["data"] + + def test_values_interrupt_v1_compat(self) -> None: + """v1 values + interrupt -> __interrupt__ still in dict (v1 compat).""" + + def my_node(state: SimpleState) -> dict[str, Any]: + answer = interrupt("what is your name?") + return {"value": answer, "items": ["done"]} + + builder: StateGraph = StateGraph(SimpleState) + builder.add_node("my_node", my_node) + builder.add_edge(START, "my_node") + builder.add_edge("my_node", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-v1-interrupt-compat"}} + chunks = list( + graph.stream( + {"value": "x", "items": []}, + config, + stream_mode="values", + ) + ) + # v1 format: should have __interrupt__ in dict + interrupt_chunks = [c for c in chunks if isinstance(c, dict) and INTERRUPT in c] + assert len(interrupt_chunks) >= 1, ( + f"Expected v1 interrupt chunks with {INTERRUPT}, got {chunks}" + ) + + def test_checkpoints_pydantic_state(self) -> None: + """v2 checkpoints + pydantic state -> values is pydantic model instance + (at least for checkpoints emitted after all channels are populated).""" + + def node_a(state: PydanticState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-v2-ckpt-pydantic"}} + chunks = list( + graph.stream( + {"value": "x", "items": []}, + config, + stream_mode="checkpoints", + stream_version="v2", + ) + ) + ckpt_chunks = [c for c in chunks if c["type"] == "checkpoints"] + assert len(ckpt_chunks) >= 1 + # At least one checkpoint (after first node runs) should have coerced values + coerced_ckpts = [ + c for c in ckpt_chunks if isinstance(c["data"]["values"], PydanticState) + ] + assert len(coerced_ckpts) >= 1, ( + f"Expected at least one checkpoint with PydanticState values, got types: " + f"{[type(c['data']['values']) for c in ckpt_chunks]}" + ) + + def test_debug_pydantic_state(self) -> None: + """v2 debug + pydantic state -> inner checkpoint payload has coerced values + (at least for checkpoints emitted after all channels are populated).""" + + def node_a(state: PydanticState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-v2-debug-pydantic"}} + chunks = list( + graph.stream( + {"value": "x", "items": []}, + config, + stream_mode="debug", + stream_version="v2", + ) + ) + debug_chunks = [c for c in chunks if c["type"] == "debug"] + checkpoint_debug = [ + c for c in debug_chunks if c["data"]["type"] == "checkpoint" + ] + assert len(checkpoint_debug) >= 1 + # At least one debug checkpoint should have coerced values + coerced_debug = [ + c + for c in checkpoint_debug + if isinstance(c["data"]["payload"]["values"], PydanticState) + ] + assert len(coerced_debug) >= 1, ( + f"Expected at least one debug checkpoint with PydanticState values, got types: " + f"{[type(c['data']['payload']['values']) for c in checkpoint_debug]}" + ) + + def test_values_pydantic_interrupt(self) -> None: + """v2 values + pydantic state + interrupt -> data is model, interrupts separated.""" + + def my_node(state: PydanticState) -> dict[str, Any]: + answer = interrupt("what is your name?") + return {"value": answer, "items": ["done"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("my_node", my_node) + builder.add_edge(START, "my_node") + builder.add_edge("my_node", END) + graph = builder.compile(checkpointer=InMemorySaver()) + + config: Any = {"configurable": {"thread_id": "test-v2-pydantic-interrupt"}} + chunks = list( + graph.stream( + {"value": "x", "items": []}, + config, + stream_mode="values", + stream_version="v2", + ) + ) + interrupt_chunks = [c for c in chunks if c.get("interrupts", ())] + assert len(interrupt_chunks) >= 1 + for c in interrupt_chunks: + assert isinstance(c["data"], PydanticState), ( + f"Expected PydanticState, got {type(c['data'])}" + ) + assert isinstance(c["interrupts"], tuple) + assert len(c["interrupts"]) > 0 + + def test_subgraph_different_pydantic_schema(self) -> None: + """Subgraph with different pydantic schema -> subgraph data coerced with subgraph's schema.""" + + class InnerState(BaseModel): + value: str + + class OuterState(BaseModel): + value: str + + def inner_node(state: InnerState) -> dict[str, Any]: + return {"value": state.value + "_inner"} + + def outer_node(state: OuterState) -> dict[str, Any]: + return {"value": state.value + "_outer"} + + inner_builder: StateGraph = StateGraph(InnerState) + inner_builder.add_node("inner_node", inner_node) + inner_builder.add_edge(START, "inner_node") + inner_builder.add_edge("inner_node", END) + inner_graph = inner_builder.compile() + + outer_builder: StateGraph = StateGraph(OuterState) + outer_builder.add_node("outer_node", outer_node) + outer_builder.add_node("inner", inner_graph) + outer_builder.add_edge(START, "outer_node") + outer_builder.add_edge("outer_node", "inner") + outer_builder.add_edge("inner", END) + outer = outer_builder.compile() + + chunks = list( + outer.stream( + {"value": "x"}, + stream_mode="values", + subgraphs=True, + stream_version="v2", + ) + ) + # Root-level values should be OuterState instances + root_values = [c for c in chunks if c["type"] == "values" and c["ns"] == ()] + assert len(root_values) >= 1 + for c in root_values: + assert isinstance(c["data"], OuterState), ( + f"Expected OuterState, got {type(c['data'])}" + ) + # Subgraph values are streamed from the subgraph's own stream() + # which runs with default stream_version="v1", so no coercion + sub_values = [c for c in chunks if c["type"] == "values" and c["ns"] != ()] + assert len(sub_values) >= 1 + + +# --- v2 validation errors --- + + +def _make_pydantic_graph() -> Any: + """Build a simple graph with PydanticState for validation error tests.""" + + def node_a(state: PydanticState) -> dict[str, Any]: + return {"value": state.value + "_a", "items": ["a"]} + + builder: StateGraph = StateGraph(PydanticState) + builder.add_node("node_a", node_a) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", END) + return builder.compile() + + +class TestV2ValidationErrors: + """Validation errors propagate for pydantic state in both v1 and v2. + + Uses `value=[1, 2, 3]` which channels accept (LastValue stores anything) + but pydantic rejects (list is not coercible to str even in lax mode). + """ + + _INVALID_INPUT: dict[str, Any] = {"value": [1, 2, 3], "items": []} + + def test_stream_v2_pydantic_validation_error(self) -> None: + """Invalid input to stream with v2 + pydantic state raises ValidationError.""" + graph = _make_pydantic_graph() + with pytest.raises(ValidationError): + list( + graph.stream( + self._INVALID_INPUT, + stream_mode="values", + stream_version="v2", + ) + ) + + def test_invoke_v2_pydantic_validation_error(self) -> None: + """Invalid input to invoke with v2 + pydantic state raises ValidationError.""" + graph = _make_pydantic_graph() + with pytest.raises(ValidationError): + graph.invoke(self._INVALID_INPUT, stream_version="v2") + + def test_invoke_v1_pydantic_validation_error(self) -> None: + """Regression: invalid input to invoke without stream_version raises ValidationError.""" + graph = _make_pydantic_graph() + with pytest.raises(ValidationError): + graph.invoke(self._INVALID_INPUT) + + # --- type narrowing compile-time checks --- # These assert_type calls verify that mypy narrows the union correctly.