diff --git a/libs/langgraph/langgraph/_internal/_constants.py b/libs/langgraph/langgraph/_internal/_constants.py index 68cb48fe8..d28289053 100644 --- a/libs/langgraph/langgraph/_internal/_constants.py +++ b/libs/langgraph/langgraph/_internal/_constants.py @@ -66,6 +66,9 @@ CONFIG_KEY_RUNTIME = sys.intern("__pregel_runtime") # holds a `Runtime` instance with context, store, stream writer, etc. CONFIG_KEY_RESUME_MAP = sys.intern("__pregel_resume_map") # holds a mapping of task ns -> resume value for resuming tasks +CONFIG_KEY_STREAM_MESSAGES_V2 = sys.intern("__pregel_stream_messages_v2") +# when True, attach StreamMessagesHandlerV2 so content-block (v2) events +# flow through stream_mode="messages"; set by StreamingHandler only. # --- Other constants --- PUSH = sys.intern("__pregel_push") @@ -107,6 +110,7 @@ RESERVED = { CONFIG_KEY_CHECKPOINT_ID, CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_RESUME_MAP, + CONFIG_KEY_STREAM_MESSAGES_V2, # other constants PUSH, PULL, diff --git a/libs/langgraph/langgraph/pregel/_messages.py b/libs/langgraph/langgraph/pregel/_messages.py index 6cda6b0a1..07f26f446 100644 --- a/libs/langgraph/langgraph/pregel/_messages.py +++ b/libs/langgraph/langgraph/pregel/_messages.py @@ -24,6 +24,11 @@ try: except ImportError: _StreamingCallbackHandler = object # type: ignore +try: + from langchain_core.tracers._streaming import _V2StreamingCallbackHandler +except ImportError: + _V2StreamingCallbackHandler = object # type: ignore + T = TypeVar("T") Meta = tuple[tuple[str, ...], dict[str, Any]] @@ -256,3 +261,126 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): **kwargs: Any, ) -> Any: self.metadata.pop(run_id, None) + + +class StreamMessagesHandlerV2(StreamMessagesHandler, _V2StreamingCallbackHandler): + """v2 variant of `StreamMessagesHandler`. + + Declaring `_V2StreamingCallbackHandler` as a base flips + `BaseChatModel.invoke` to route through `_stream_chat_model_events` + (firing `on_stream_event`) instead of `_stream` (firing + `on_llm_new_token`). Inherits `on_stream_event` from the parent, + which forwards protocol events onto the messages stream channel. + + Pregel attaches this class instead of the v1 handler only when + `StreamingHandler` opts in via the internal + `CONFIG_KEY_STREAM_MESSAGES_V2` config key; direct + `graph.stream(stream_mode="messages")` callers keep the v1 + AIMessageChunk shape. + """ + + def on_llm_new_token( + self, + token: str, + *, + chunk: ChatGenerationChunk | None = None, + run_id: UUID, + parent_run_id: UUID | None = None, + tags: list[str] | None = None, + **kwargs: Any, + ) -> Any: + """Intentional no-op — v1 chunks are not used on v2-flagged runs. + + The v2 marker already steers `invoke` to the event generator, so + `on_llm_new_token` should not fire under normal routing. This + override stays a pass-through (no call to `super()`) to make + the intent explicit and to guard against any caller (e.g. a + node that calls `model.stream()` directly, which still fires + the v1 callback) leaking AIMessageChunks onto a v2-flagged + messages stream. + """ + # Intentionally empty: v2 handler does not forward v1 chunks. + + def __init__( + self, + stream: Callable[[StreamChunk], None], + subgraphs: bool, + *, + parent_ns: tuple[str, ...] | None = None, + ) -> None: + super().__init__(stream, subgraphs, parent_ns=parent_ns) + self._streamed_run_ids: set[UUID] = set() + + def on_llm_end( + self, + response: LLMResult, + *, + run_id: UUID, + parent_run_id: UUID | None = None, + **kwargs: Any, + ) -> Any: + if meta := self.metadata.get(run_id): + if response.generations and response.generations[0]: + gen = response.generations[0][0] + if isinstance(gen, ChatGeneration): + if run_id in self._streamed_run_ids: + if gen.message.id is None: + gen.message.id = str(uuid4()) + self.seen.add(gen.message.id) + else: + self._emit(meta, gen.message, dedupe=True) + self._streamed_run_ids.discard(run_id) + self.metadata.pop(run_id, None) + + def on_llm_error( + self, + error: BaseException, + *, + run_id: UUID, + parent_run_id: UUID | None = None, + **kwargs: Any, + ) -> Any: + self._streamed_run_ids.discard(run_id) + super().on_llm_error( + error, + run_id=run_id, + parent_run_id=parent_run_id, + **kwargs, + ) + + def on_stream_event( + self, + event: dict[str, Any], + *, + run_id: UUID, + parent_run_id: UUID | None = None, + tags: list[str] | None = None, + **kwargs: Any, + ) -> Any: + """Forward a protocol event from `stream_v2` as a messages stream part. + + Fires once per `MessagesData` event (`message-start`, per-block + `content-block-*`, `message-finish`). The transformer layer + correlates events back to a single `ChatModelStream` via + `metadata["run_id"]` — attached here so the v1 + `stream_mode="messages"` output (which emits + `(AIMessageChunk, metadata)` via `on_llm_new_token`) keeps its + original metadata shape. + + Lives on the v2 handler rather than the v1 base: content-block + events are a v2-only concept, and forwarding them only when the + v2 handler is attached keeps the message channel's shape + predictable for v1 callers. + """ + if meta := self.metadata.get(run_id): + # Record message_id on message-start so on_chain_end's + # dedupe skips the finalized AIMessage the node returns + # (otherwise the messages projection double-counts: once + # from streaming, once from the chain output). + if event.get("event") == "message-start": + self._streamed_run_ids.add(run_id) + msg_id = event.get("message_id") + if msg_id: + self.seen.add(msg_id) + v2_meta = {**meta[1], "run_id": str(run_id)} + self.stream((meta[0], "messages", (event, v2_meta))) diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index ce38ae1bf..8e9719d3f 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -73,6 +73,7 @@ from langgraph._internal._constants import ( CONFIG_KEY_RUNTIME, CONFIG_KEY_SEND, CONFIG_KEY_STREAM, + CONFIG_KEY_STREAM_MESSAGES_V2, CONFIG_KEY_TASK_ID, CONFIG_KEY_THREAD_ID, ERROR, @@ -133,7 +134,10 @@ from langgraph.pregel._loop import ( AsyncPregelLoop, SyncPregelLoop, ) -from langgraph.pregel._messages import StreamMessagesHandler +from langgraph.pregel._messages import ( + StreamMessagesHandler, + StreamMessagesHandlerV2, +) from langgraph.pregel._read import DEFAULT_BOUND, PregelNode from langgraph.pregel._retry import RetryPolicy from langgraph.pregel._runner import PregelRunner @@ -340,6 +344,17 @@ class NodeBuilder: ) +_STREAM_V2_MODES: list[StreamMode] = [ + "values", + "updates", + "messages", + "custom", + "checkpoints", + "tasks", + "debug", +] + + class Pregel( PregelProtocol[StateT, ContextT, InputT, OutputT], Generic[StateT, ContextT, InputT, OutputT], @@ -2586,19 +2601,7 @@ class Pregel( stream = SyncQueue() config = ensure_config(self.config, config) - callback_manager = get_callback_manager_for_config(config) - if "ls_integration" not in callback_manager.metadata: - callback_manager.add_metadata({"ls_integration": "langgraph"}) - run_manager = callback_manager.on_chain_start( - None, - input, - name=config.get("run_name", self.get_name()), - run_id=config.get("run_id"), - ) - graph_callback_manager = get_sync_graph_callback_manager_for_config( - config, - run_id=run_manager.run_id, - ) + run_manager = None try: # assign defaults ( @@ -2619,6 +2622,36 @@ class Pregel( interrupt_after=interrupt_after, durability=durability, ) + callback_manager = get_callback_manager_for_config(config) + if "messages" in stream_modes and version != "v2": + # Strip any inherited v2 messages handler so a v1 stream + # does not get routed through the content-block event + # protocol. Leave v1 handlers in place — an outer + # stream(stream_mode="messages", subgraphs=True) relies + # on its inheritable handler to observe events emitted + # by inner stream(stream_mode="messages") calls. + callback_manager.handlers = [ + h + for h in callback_manager.handlers + if not isinstance(h, StreamMessagesHandlerV2) + ] + callback_manager.inheritable_handlers = [ + h + for h in callback_manager.inheritable_handlers + if not isinstance(h, StreamMessagesHandlerV2) + ] + if "ls_integration" not in callback_manager.metadata: + callback_manager.add_metadata({"ls_integration": "langgraph"}) + run_manager = callback_manager.on_chain_start( + None, + input, + name=config.get("run_name", self.get_name()), + run_id=config.get("run_id"), + ) + graph_callback_manager = get_sync_graph_callback_manager_for_config( + config, + run_id=run_manager.run_id, + ) if checkpointer is None and durability is not None: warnings.warn( "`durability` has no effect when no checkpointer is present.", @@ -2630,8 +2663,16 @@ class Pregel( # set up messages stream mode if "messages" in stream_modes: ns_ = cast(str | None, config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)) + use_stream_messages_v2 = bool( + version == "v2" and config[CONF].get(CONFIG_KEY_STREAM_MESSAGES_V2) + ) + messages_handler_cls = ( + StreamMessagesHandlerV2 + if use_stream_messages_v2 + else StreamMessagesHandler + ) run_manager.inheritable_handlers.append( - StreamMessagesHandler( + messages_handler_cls( stream.put, subgraphs, parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None, @@ -2808,7 +2849,8 @@ class Pregel( # set final channel values as run output run_manager.on_chain_end(loop.output) except BaseException as e: - run_manager.on_chain_error(e) + if run_manager is not None: + run_manager.on_chain_error(e) raise @overload @@ -2948,33 +2990,7 @@ class Pregel( ) config = ensure_config(self.config, config) - callback_manager = get_async_callback_manager_for_config(config) - if "ls_integration" not in callback_manager.metadata: - callback_manager.add_metadata({"ls_integration": "langgraph"}) - run_manager = await callback_manager.on_chain_start( - None, - input, - name=config.get("run_name", self.get_name()), - run_id=config.get("run_id"), - ) - graph_callback_manager = get_async_graph_callback_manager_for_config( - config, - run_id=run_manager.run_id, - ) - # if running from astream_log() run each proc with streaming - do_stream = ( - next( - ( - True - for h in run_manager.handlers - if isinstance(h, _StreamingCallbackHandler) - and not isinstance(h, StreamMessagesHandler) - ), - False, - ) - if _StreamingCallbackHandler is not None - else False - ) + run_manager = None try: # assign defaults ( @@ -2995,6 +3011,50 @@ class Pregel( interrupt_after=interrupt_after, durability=durability, ) + callback_manager = get_async_callback_manager_for_config(config) + if "messages" in stream_modes and version != "v2": + # Strip any inherited v2 messages handler so a v1 stream + # does not get routed through the content-block event + # protocol. Leave v1 handlers in place — an outer + # astream(stream_mode="messages", subgraphs=True) relies + # on its inheritable handler to observe events emitted + # by inner astream(stream_mode="messages") calls. + callback_manager.handlers = [ + h + for h in callback_manager.handlers + if not isinstance(h, StreamMessagesHandlerV2) + ] + callback_manager.inheritable_handlers = [ + h + for h in callback_manager.inheritable_handlers + if not isinstance(h, StreamMessagesHandlerV2) + ] + if "ls_integration" not in callback_manager.metadata: + callback_manager.add_metadata({"ls_integration": "langgraph"}) + run_manager = await callback_manager.on_chain_start( + None, + input, + name=config.get("run_name", self.get_name()), + run_id=config.get("run_id"), + ) + graph_callback_manager = get_async_graph_callback_manager_for_config( + config, + run_id=run_manager.run_id, + ) + # if running from astream_log() run each proc with streaming + do_stream = ( + next( + ( + True + for h in run_manager.handlers + if isinstance(h, _StreamingCallbackHandler) + and not isinstance(h, StreamMessagesHandler) + ), + False, + ) + if _StreamingCallbackHandler is not None + else False + ) if checkpointer is None and durability is not None: warnings.warn( "`durability` has no effect when no checkpointer is present.", @@ -3007,8 +3067,16 @@ class Pregel( if "messages" in stream_modes: # namespace can be None in a root level graph? ns_ = cast(str | None, config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)) + use_stream_messages_v2 = bool( + version == "v2" and config[CONF].get(CONFIG_KEY_STREAM_MESSAGES_V2) + ) + messages_handler_cls = ( + StreamMessagesHandlerV2 + if use_stream_messages_v2 + else StreamMessagesHandler + ) run_manager.inheritable_handlers.append( - StreamMessagesHandler( + messages_handler_cls( stream_put, subgraphs, parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None, @@ -3238,7 +3306,8 @@ class Pregel( # set final channel values as run output await run_manager.on_chain_end(loop.output) except BaseException as e: - await asyncio.shield(run_manager.on_chain_error(e)) + if run_manager is not None: + await asyncio.shield(run_manager.on_chain_error(e)) raise def stream_v2( @@ -3277,12 +3346,13 @@ class Pregel( ValuesTransformer, ) - values_t = ValuesTransformer() + parent_ns = _resolve_parent_ns(self.config, config) + values_t = ValuesTransformer(parent_ns=parent_ns) compiled_instances = [f() for f in self._stream_transformers] mux = StreamMux( [ values_t, - MessagesTransformer(), + MessagesTransformer(parent_ns=parent_ns), *compiled_instances, *(transformers or ()), ], @@ -3291,16 +3361,8 @@ class Pregel( graph_iter = iter( self.stream( input, - config, - stream_mode=[ - "values", - "updates", - "messages", - "custom", - "checkpoints", - "tasks", - "debug", - ], + patch_configurable(config, {CONFIG_KEY_STREAM_MESSAGES_V2: True}), + stream_mode=_STREAM_V2_MODES, subgraphs=True, version="v2", interrupt_before=interrupt_before, @@ -3339,12 +3401,13 @@ class Pregel( ValuesTransformer, ) - values_t = ValuesTransformer() + parent_ns = _resolve_parent_ns(self.config, config) + values_t = ValuesTransformer(parent_ns=parent_ns) compiled_instances = [f() for f in self._stream_transformers] mux = StreamMux( [ values_t, - MessagesTransformer(), + MessagesTransformer(parent_ns=parent_ns), *compiled_instances, *(transformers or ()), ], @@ -3352,16 +3415,8 @@ class Pregel( ) graph_aiter = self.astream( input, - config, - stream_mode=[ - "values", - "updates", - "messages", - "custom", - "checkpoints", - "tasks", - "debug", - ], + patch_configurable(config, {CONFIG_KEY_STREAM_MESSAGES_V2: True}), + stream_mode=_STREAM_V2_MODES, subgraphs=True, version="v2", interrupt_before=interrupt_before, @@ -3844,6 +3899,24 @@ def _coerce_checkpoint_values(payload: Any, mapper: Callable[[Any], Any]) -> Non payload["values"] = mapper(payload["values"]) +def _resolve_parent_ns( + graph_config: RunnableConfig | None, call_config: RunnableConfig | None +) -> tuple[str, ...]: + """Return the checkpoint namespace the caller is running under. + + `stream_v2` uses this to scope its native projections + (`ValuesTransformer`, `MessagesTransformer`) to events emitted at + the run's own level. A root call resolves to `()`; a call made + from inside a node carries the outer graph's task namespace so the + projection still matches its own root-level events. + """ + merged = ensure_config(graph_config, call_config) + ns = merged.get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS) + if not ns: + return () + return tuple(ns.split(NS_SEP)) + + def _build_server_info( config: RunnableConfig, parent_runtime: Runtime[Any] ) -> ServerInfo | None: diff --git a/libs/langgraph/langgraph/stream/run_stream.py b/libs/langgraph/langgraph/stream/run_stream.py index 98ab08586..7be3f7534 100644 --- a/libs/langgraph/langgraph/stream/run_stream.py +++ b/libs/langgraph/langgraph/stream/run_stream.py @@ -64,14 +64,23 @@ class GraphRunStream: self._wire_request_more(mux) def _wire_request_more(self, mux: StreamMux) -> None: - """Install `_request_more` on every sync EventLog so cursors - can drive the pump when their buffer catches up.""" + """Install `_request_more` on every sync EventLog so cursors can + drive the pump when their buffer catches up. + + Also calls `_bind_pump` on any transformer that exposes it, so + transformers producing ChatModelStream objects (e.g. + MessagesTransformer) can wire the pull callback on each stream + as it's created. + """ mux._events._request_more = self._pump_next for value in mux.extensions.values(): if isinstance(value, EventLog): value._request_more = self._pump_next elif isinstance(value, StreamChannel): value._log._request_more = self._pump_next + for transformer in mux._transformers: + if hasattr(transformer, "_bind_pump"): + transformer._bind_pump(self._pump_next) def _pump_next(self) -> bool: """Pull one event from the graph and push it through the mux. @@ -125,8 +134,7 @@ class GraphRunStream: def output(self) -> dict[str, Any] | None: """Drive the run to completion and return the final state.""" _drive_until_done(self._pump_next) - err = self._values_transformer.error - if err is not None: + if (err := self._values_transformer.error) is not None: raise err return self._values_transformer._latest @@ -139,8 +147,7 @@ class GraphRunStream: BaseException: If the run ended with an error. """ _drive_until_done(self._pump_next) - err = self._values_transformer.error - if err is not None: + if (err := self._values_transformer.error) is not None: raise err return self._values_transformer._interrupted @@ -152,8 +159,7 @@ class GraphRunStream: BaseException: If the run ended with an error. """ _drive_until_done(self._pump_next) - err = self._values_transformer.error - if err is not None: + if (err := self._values_transformer.error) is not None: raise err return self._values_transformer._interrupts @@ -262,13 +268,23 @@ class AsyncGraphRunStream: def _wire_arequest_more(self, mux: StreamMux) -> None: """Install `_arequest_more` on every async EventLog so cursors - can drive the pump when their buffer catches up.""" + can drive the pump when their buffer catches up. + + Also calls `_bind_apump` on any transformer that exposes it, + so transformers producing `AsyncChatModelStream` objects (e.g. + `MessagesTransformer`) can fan the pull callback out to each + stream's projections. Mirrors the sync `_wire_request_more` + plumbing. + """ mux._events._arequest_more = self._apump_next for value in mux.extensions.values(): if isinstance(value, EventLog): value._arequest_more = self._apump_next elif isinstance(value, StreamChannel): value._log._arequest_more = self._apump_next + for transformer in mux._transformers: + if hasattr(transformer, "_bind_apump"): + transformer._bind_apump(self._apump_next) async def _apump_next(self) -> bool: """Drive one pump step, or wait for the active pumper to drive one. diff --git a/libs/langgraph/langgraph/stream/transformers.py b/libs/langgraph/langgraph/stream/transformers.py index dec52515f..c548f4aaf 100644 --- a/libs/langgraph/langgraph/stream/transformers.py +++ b/libs/langgraph/langgraph/stream/transformers.py @@ -1,10 +1,21 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any, cast + +from langchain_core.language_models._compat_bridge import message_to_events +from langchain_core.language_models.chat_model_stream import ( + AsyncChatModelStream, + ChatModelStream, +) +from langchain_core.messages import AIMessageChunk, BaseMessage +from langchain_protocol.protocol import MessagesData from langgraph.stream._event_log import EventLog from langgraph.stream._types import ProtocolEvent, StreamTransformer +if TYPE_CHECKING: + from collections.abc import Awaitable, Callable + class ValuesTransformer(StreamTransformer): """Capture values events as a drainable stream of state snapshots. @@ -17,17 +28,22 @@ class ValuesTransformer(StreamTransformer): Native transformer — projection keys are exposed as direct attributes on the run stream (e.g. `run.values`). - Only root-namespace values events are captured; subgraph state - snapshots are ignored. + Only values events at the run's own level are captured; snapshots + from deeper subgraphs are left in the main event log but excluded + from the projection. "Own level" is defined by `parent_ns`, which + `stream_v2` / `astream_v2` populate from the caller's checkpoint + namespace so that a nested `stream_v2` call still sees its own + root snapshots. """ _native = True - def __init__(self) -> None: + def __init__(self, *, parent_ns: tuple[str, ...] = ()) -> None: self._log: EventLog[dict[str, Any]] = EventLog() self._latest: dict[str, Any] | None = None self._interrupted = False self._interrupts: list[Any] = [] + self._parent_ns: list[str] = list(parent_ns) def init(self) -> dict[str, Any]: return {"values": self._log} @@ -44,7 +60,7 @@ class ValuesTransformer(StreamTransformer): if event["method"] != "values": return True params = event["params"] - if params["namespace"]: + if params["namespace"] != self._parent_ns: return True self._latest = params["data"] interrupts = params.get("interrupts", ()) @@ -56,34 +72,175 @@ class ValuesTransformer(StreamTransformer): class MessagesTransformer(StreamTransformer): - """Pass through raw (chunk, metadata) tuples from messages events. + """Capture messages events as ChatModelStream objects. - This is the same shape as today's `stream_mode="messages"` output. - A follow-on PR will replace this with a richer transformer that - produces ChatModelStream objects using the protocol handler. + The messages projection yields one `ChatModelStream` (or + `AsyncChatModelStream`) per LLM call. Consumers iterate + `run.messages` to get stream handles, then use each handle's typed + projections (`.text`, `.reasoning`, `.tool_calls`, `.usage`, + `.output`) for per-message content. - Only root-namespace messages events are captured; tokens emitted - from subgraphs are dropped from the `messages` projection. - Consumers that need subgraph tokens should iterate the raw event - stream or register a custom transformer. + Two input shapes are handled (via `params["data"] = (payload, + metadata)` from `StreamMessagesHandler`): - Native transformer — projection keys are exposed as direct - attributes on the run stream (e.g. `run.messages`). + 1. Protocol event (dict with `"event"` key) — emitted by + `stream_v2()` / `astream_v2()` via the `on_stream_event` + callback. Routed to an existing `ChatModelStream` by + `metadata["run_id"]`. A `message-start` event creates a new + stream; `message-finish` closes it. + 2. Whole `AIMessage` — emitted from `on_chain_end` when a node + returns a finalized message. Replayed as a synthetic protocol + event lifecycle via `message_to_events`, then the + already-complete stream is pushed to the log. + + V1 `AIMessageChunk` tuples (from `on_llm_new_token`) are not + streamed into this projection: chat models that want to populate + `run.messages` with content-block streaming must use + `stream_v2()` / `astream_v2()`. Models called via the legacy + `stream()` method still surface their final `AIMessage` via + `on_chain_end` when a node returns it as state. + + Only events at the run's own level are projected; tokens from + deeper subgraphs are left in the main event log but excluded from + `.messages`. "Own level" is defined by `parent_ns`, which + `stream_v2` / `astream_v2` populate from the caller's checkpoint + namespace so that a `stream_v2` call inside a node still sees its + own root chat model streams on `.messages`. Consumers that need + subgraph tokens should iterate the raw event stream or register a + custom transformer. + + Native transformer — the `messages` projection is exposed as a + direct attribute on the run stream. """ _native = True - def __init__(self) -> None: - self._log: EventLog[tuple[Any, dict[str, Any]]] = EventLog() + def __init__(self, *, parent_ns: tuple[str, ...] = ()) -> None: + self._log: EventLog[ChatModelStream] = EventLog() + # Correlate protocol events back to a ChatModelStream by run_id + # (attached to the event's metadata by StreamMessagesHandler). + self._by_run: dict[str, ChatModelStream] = {} + self._pump_fn: Callable[[], bool] | None = None + self._apump_fn: Callable[[], Awaitable[bool]] | None = None + # Root scope for this projection. Only chat model streams whose + # emitted namespace matches `parent_ns` are surfaced on + # `.messages`; events from deeper subgraphs stay in the main + # event log for other consumers but are not projected here. + self._parent_ns: list[str] = list(parent_ns) def init(self) -> dict[str, Any]: return {"messages": self._log} + def _bind_pump(self, fn: Callable[[], bool]) -> None: + """Wire the sync pull callback. Called by GraphRunStream._wire_request_more.""" + self._pump_fn = fn + + def _bind_apump(self, fn: Callable[[], Awaitable[bool]]) -> None: + """Wire the async pull callback. + + Called by `AsyncGraphRunStream._wire_arequest_more` so each + `AsyncChatModelStream` this transformer creates can drive the + shared graph pump from its projection cursors. + """ + self._apump_fn = fn + + def _make_stream( + self, + *, + namespace: list[str], + node: str | None, + message_id: str | None, + ) -> ChatModelStream: + """Create a ChatModelStream (sync) or AsyncChatModelStream (async). + + Wires whichever pump is bound. Prefers the async pump so nested + iteration under `AsyncGraphRunStream` drives the graph forward + without a background task. The unwired fallback (no pump bound) + is used by unit tests that dispatch events manually. + """ + if self._apump_fn is not None: + astream = AsyncChatModelStream( + namespace=namespace, + node=node, + message_id=message_id, + ) + astream.set_arequest_more(self._apump_fn) + return astream + if self._pump_fn is not None: + stream: ChatModelStream = ChatModelStream( + namespace=namespace, + node=node, + message_id=message_id, + ) + stream.set_request_more(self._pump_fn) + return stream + return AsyncChatModelStream( + namespace=namespace, + node=node, + message_id=message_id, + ) + def process(self, event: ProtocolEvent) -> bool: if event["method"] != "messages": return True params = event["params"] - if params["namespace"]: + if params["namespace"] != self._parent_ns: return True - self._log.push(params["data"]) + + payload, metadata = params["data"] + node: str | None = metadata.get("langgraph_node") + run_id = str(metadata.get("run_id", "")) if metadata else "" + + if isinstance(payload, dict) and "event" in payload: + self._route_protocol_event( + cast("MessagesData", payload), run_id=run_id, node=node + ) + elif isinstance(payload, BaseMessage) and not isinstance( + payload, AIMessageChunk + ): + self._route_whole_message(payload, node=node) + # Legacy AIMessageChunk tuples (from on_llm_new_token) are ignored; + # v1 streaming callers must switch to stream_v2() to populate this + # projection. + return True + + def _route_protocol_event( + self, + event: MessagesData, + *, + run_id: str, + node: str | None, + ) -> None: + event_type = event.get("event") + if event_type == "message-start": + message_id = event.get("message_id") + stream = self._make_stream( + namespace=[], + node=node, + message_id=str(message_id) if message_id is not None else None, + ) + self._by_run[run_id] = stream + self._log.push(stream) + stream.dispatch(event) + elif run_id in self._by_run: + stream = self._by_run[run_id] + stream.dispatch(event) + if event_type == "message-finish": + del self._by_run[run_id] + + def _route_whole_message(self, message: BaseMessage, *, node: str | None) -> None: + stream = self._make_stream(namespace=[], node=node, message_id=message.id) + for evt in message_to_events(message, message_id=message.id): + stream.dispatch(evt) + self._log.push(stream) + + def finalize(self) -> None: + """Clear any routing state — streams close themselves via `message-finish`.""" + self._by_run.clear() + + def fail(self, err: BaseException) -> None: + """Propagate run error to any streams still open when the graph fails.""" + for stream in list(self._by_run.values()): + stream.fail(err) + self._by_run.clear() diff --git a/libs/langgraph/pyproject.toml b/libs/langgraph/pyproject.toml index 2573395fb..b5716f639 100644 --- a/libs/langgraph/pyproject.toml +++ b/libs/langgraph/pyproject.toml @@ -24,7 +24,7 @@ classifiers = [ 'Programming Language :: Python :: 3.13', ] dependencies = [ - "langchain-core==1.3.0a2", + "langchain-core==1.3.2", "langgraph-checkpoint>=2.1.0,<5.0.0", "langgraph-sdk>=0.3.0,<0.4.0", "langgraph-prebuilt>=1.0.9,<1.1.0", diff --git a/libs/langgraph/tests/test_pregel_stream_v2.py b/libs/langgraph/tests/test_pregel_stream_v2.py index e0d47126a..4e890d510 100644 --- a/libs/langgraph/tests/test_pregel_stream_v2.py +++ b/libs/langgraph/tests/test_pregel_stream_v2.py @@ -1,4 +1,4 @@ -"""Tests for `Pregel.stream_v2` / `astream_v2` and the transformer pipeline.""" +"""Tests for Pregel.stream_v2 / astream_v2 and the transformer pipeline.""" from __future__ import annotations @@ -40,7 +40,6 @@ def _event( namespace: list[str] | None = None, interrupts: tuple[Any, ...] | None = None, ) -> ProtocolEvent: - """Build a test ProtocolEvent with sensible defaults.""" params: dict[str, Any] = { "namespace": namespace or [], "timestamp": TS, @@ -52,7 +51,7 @@ def _event( # --------------------------------------------------------------------------- -# Shared state and graph builders +# Shared graph builders # --------------------------------------------------------------------------- @@ -62,8 +61,6 @@ class SimpleState(TypedDict): def _build_simple_graph(): - """Two-node graph: node_a appends 'a', node_b appends 'b'.""" - def node_a(state: SimpleState) -> dict: return {"value": state["value"] + "A", "items": ["a"]} @@ -80,8 +77,6 @@ def _build_simple_graph(): def _build_interrupt_graph(): - """Graph that interrupts before node_b.""" - def node_a(state: SimpleState) -> dict: return {"value": state["value"] + "A", "items": ["a"]} @@ -99,8 +94,6 @@ def _build_interrupt_graph(): def _build_error_graph(): - """Graph where node_b raises.""" - def node_a(state: SimpleState) -> dict: return {"value": state["value"] + "A", "items": ["a"]} @@ -117,8 +110,6 @@ def _build_error_graph(): def _build_custom_stream_graph(): - """Graph that emits custom stream events.""" - def node_a(state: SimpleState, *, writer: StreamWriter) -> dict: writer({"step": "start"}) writer({"step": "end"}) @@ -140,7 +131,7 @@ class TestEventLog: def test_sync_iteration(self) -> None: log: EventLog[int] = EventLog() log._bind(is_async=False) - it = iter(log) # subscribe before pushing + it = iter(log) log.push(1) log.push(2) log.push(3) @@ -148,7 +139,6 @@ class TestEventLog: assert list(it) == [1, 2, 3] def test_drain_on_consume(self) -> None: - """Items are popped as consumed — no retention across iterations.""" log: EventLog[str] = EventLog() log._bind(is_async=False) it = iter(log) @@ -156,11 +146,9 @@ class TestEventLog: log.push("b") log.close() assert list(it) == ["a", "b"] - # Buffer is drained. assert list(log._items) == [] def test_second_subscribe_raises(self) -> None: - """Only one subscriber allowed; tee() is the fan-out escape hatch.""" log: EventLog[str] = EventLog() log._bind(is_async=False) log.close() @@ -169,7 +157,7 @@ class TestEventLog: iter(log) def test_pre_subscription_push_is_noop(self) -> None: - """Lazy-subscribe: pushes before subscription are dropped silently.""" + # Lazy-subscribe: pushes before subscription are dropped silently. log: EventLog[int] = EventLog() log._bind(is_async=False) log.push(1) @@ -177,7 +165,6 @@ class TestEventLog: it = iter(log) log.push(3) log.close() - # Only the post-subscribe push survives. assert list(it) == [3] def test_fail_propagation(self) -> None: @@ -189,6 +176,72 @@ class TestEventLog: with pytest.raises(ValueError, match="test error"): list(it) + def test_sync_cursor_yields_items_before_error(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=False) + it = iter(log) + log.push(1) + log.push(2) + log.push(3) + log.fail(ValueError("late error")) + items: list[int] = [] + with pytest.raises(ValueError, match="late error"): + for item in it: + items.append(item) + assert items == [1, 2, 3] + + def test_push_after_close_raises(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=False) + it = iter(log) + log.push(1) + log.close() + with pytest.raises(RuntimeError, match="Cannot push to a closed EventLog"): + log.push(2) + _ = list(it) + + def test_push_after_fail_raises(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=False) + it = iter(log) + log.fail(ValueError("err")) + with pytest.raises(RuntimeError, match="Cannot push to a closed EventLog"): + log.push(1) + with pytest.raises(ValueError, match="err"): + list(it) + + def test_empty_log_sync(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=False) + log.close() + assert list(log) == [] + + def test_empty_log_fail_sync(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=False) + log.fail(ValueError("empty fail")) + with pytest.raises(ValueError, match="empty fail"): + list(log) + + def test_unbound_iter_raises(self) -> None: + log: EventLog[int] = EventLog() + log.close() + with pytest.raises(TypeError, match="has not been bound"): + list(log) + + def test_sync_bound_aiter_raises(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=False) + log.close() + with pytest.raises(TypeError, match="bound to sync mode"): + log.__aiter__() + + def test_double_bind_raises(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=False) + with pytest.raises(RuntimeError, match="already bound"): + log._bind(is_async=True) + @pytest.mark.anyio async def test_async_iteration(self) -> None: log: EventLog[int] = EventLog() @@ -197,8 +250,7 @@ class TestEventLog: for i in range(3): log.push(i) log.close() - items = [item async for item in cursor] - assert items == [0, 1, 2] + assert [item async for item in cursor] == [0, 1, 2] @pytest.mark.anyio async def test_async_second_subscribe_raises(self) -> None: @@ -220,24 +272,8 @@ class TestEventLog: async for _ in cursor: pass - def test_sync_cursor_yields_items_before_error(self) -> None: - """Sync cursor should yield all buffered items before raising.""" - log: EventLog[int] = EventLog() - log._bind(is_async=False) - it = iter(log) - log.push(1) - log.push(2) - log.push(3) - log.fail(ValueError("late error")) - items: list[int] = [] - with pytest.raises(ValueError, match="late error"): - for item in it: - items.append(item) - assert items == [1, 2, 3] - @pytest.mark.anyio async def test_async_cursor_yields_items_before_error(self) -> None: - """Async cursor should yield all buffered items before raising.""" log: EventLog[int] = EventLog() log._bind(is_async=True) cursor = aiter(log) @@ -251,54 +287,15 @@ class TestEventLog: items.append(item) assert items == [1, 2, 3] - def test_push_after_close_raises(self) -> None: - """Push after close should raise RuntimeError (when subscribed).""" - log: EventLog[int] = EventLog() - log._bind(is_async=False) - it = iter(log) - log.push(1) - log.close() - with pytest.raises(RuntimeError, match="Cannot push to a closed EventLog"): - log.push(2) - _ = list(it) - - def test_push_after_fail_raises(self) -> None: - """Fail closes the log, so push after fail should also raise (when subscribed).""" - log: EventLog[int] = EventLog() - log._bind(is_async=False) - it = iter(log) - log.fail(ValueError("err")) - with pytest.raises(RuntimeError, match="Cannot push to a closed EventLog"): - log.push(1) - with pytest.raises(ValueError, match="err"): - list(it) - - def test_empty_log_sync(self) -> None: - """Iterating a closed empty log should yield nothing.""" - log: EventLog[int] = EventLog() - log._bind(is_async=False) - log.close() - assert list(log) == [] - @pytest.mark.anyio async def test_empty_log_async(self) -> None: - """Async-iterating a closed empty log should yield nothing.""" log: EventLog[int] = EventLog() log._bind(is_async=True) log.close() assert [item async for item in log] == [] - def test_empty_log_fail_sync(self) -> None: - """Failing an empty log should raise immediately with no items.""" - log: EventLog[int] = EventLog() - log._bind(is_async=False) - log.fail(ValueError("empty fail")) - with pytest.raises(ValueError, match="empty fail"): - list(log) - @pytest.mark.anyio async def test_empty_log_fail_async(self) -> None: - """Failing an empty log should raise immediately with no items (async).""" log: EventLog[int] = EventLog() log._bind(is_async=True) log.fail(ValueError("empty fail")) @@ -306,37 +303,14 @@ class TestEventLog: async for _ in log: pass - def test_unbound_iter_raises(self) -> None: - """Iterating an unbound EventLog should raise TypeError.""" - log: EventLog[int] = EventLog() - log.close() - with pytest.raises(TypeError, match="has not been bound"): - list(log) - - def test_sync_bound_aiter_raises(self) -> None: - """Sync-bound EventLog should reject async iteration.""" - log: EventLog[int] = EventLog() - log._bind(is_async=False) - log.close() - with pytest.raises(TypeError, match="bound to sync mode"): - log.__aiter__() - @pytest.mark.anyio async def test_async_bound_iter_raises(self) -> None: - """Async-bound EventLog should reject sync iteration.""" log: EventLog[int] = EventLog() log._bind(is_async=True) log.close() with pytest.raises(TypeError, match="bound to async mode"): iter(log) - def test_double_bind_raises(self) -> None: - """Binding an already-bound EventLog should raise.""" - log: EventLog[int] = EventLog() - log._bind(is_async=False) - with pytest.raises(RuntimeError, match="already bound"): - log._bind(is_async=True) - # --------------------------------------------------------------------------- # StreamChannel unit tests @@ -362,12 +336,10 @@ class TestStreamChannel: ch.push("x") ch.push("y") ch._close() - # Wire callback fires on every push, regardless of subscription. assert forwarded == ["x", "y"] assert list(it) == ["x", "y"] def test_fail_propagation(self) -> None: - """_fail() should propagate the error through the underlying log.""" ch: StreamChannel[str] = StreamChannel("test") ch._bind(is_async=False) it = iter(ch) @@ -379,20 +351,7 @@ class TestStreamChannel: items.append(item) assert items == ["a"] - @pytest.mark.anyio - async def test_async_iteration(self) -> None: - """Async iteration should delegate to the inner event log.""" - ch: StreamChannel[str] = StreamChannel("test") - ch._bind(is_async=True) - cursor = ch.__aiter__() - ch._log.push("x") - ch._log.push("y") - ch._close() - items = [item async for item in cursor] - assert items == ["x", "y"] - def test_push_without_wire(self) -> None: - """Push without a wire callback should still append to the log.""" ch: StreamChannel[int] = StreamChannel("test") ch._bind(is_async=False) assert ch._wire_fn is None @@ -401,6 +360,16 @@ class TestStreamChannel: ch._close() assert list(it) == [42] + @pytest.mark.anyio + async def test_async_iteration(self) -> None: + ch: StreamChannel[str] = StreamChannel("test") + ch._bind(is_async=True) + cursor = ch.__aiter__() + ch._log.push("x") + ch._log.push("y") + ch._close() + assert [item async for item in cursor] == ["x", "y"] + # --------------------------------------------------------------------------- # stream_v2 sync tests @@ -409,30 +378,19 @@ class TestStreamChannel: class TestStreamV2Sync: def test_values_projection(self) -> None: - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) snapshots = list(run.values) - # Should have at least the initial + per-node snapshots. assert len(snapshots) >= 1 - # Last snapshot should have both nodes' effects. last = snapshots[-1] - assert "A" in last["value"] - assert "B" in last["value"] + assert "A" in last["value"] and "B" in last["value"] def test_output(self) -> None: - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) output = run.output - assert output is not None - assert output["value"] == "xAB" - assert output["items"] == ["a", "b"] + assert output == {"value": "xAB", "items": ["a", "b"]} def test_raw_event_iteration(self) -> None: - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) events = list(run) assert len(events) > 0 for event in events: @@ -442,124 +400,85 @@ class TestStreamV2Sync: assert isinstance(event["params"]["timestamp"], int) def test_extensions_has_native_keys(self) -> None: - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) - # Drain events so the run completes. + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) _ = run.output - assert "values" in run.extensions - assert "messages" in run.extensions - # Native keys should also be direct attributes. + assert "values" in run.extensions and "messages" in run.extensions assert run.values is run.extensions["values"] assert run.messages is run.extensions["messages"] def test_extensions_is_read_only(self) -> None: - """`run.extensions` must reject mutations so users can't corrupt mux state.""" - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) with pytest.raises(TypeError): run.extensions["new_key"] = object() # type: ignore[index] with pytest.raises(TypeError): del run.extensions["values"] # type: ignore[attr-defined] def test_custom_stream_events(self) -> None: - graph = _build_custom_stream_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_custom_stream_graph().stream_v2({"value": "x", "items": []}) custom_events = [e for e in run if e["method"] == "custom"] assert len(custom_events) == 2 assert custom_events[0]["params"]["data"] == {"step": "start"} assert custom_events[1]["params"]["data"] == {"step": "end"} def test_interleave_values_and_messages(self) -> None: - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) - + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) tagged = list(run.interleave("values", "messages")) names = [name for name, _ in tagged] assert set(names).issubset({"values", "messages"}) - # Values projection must have fired at least once. assert names.count("values") >= 1 - # Values have been drained by the interleave cursor — re-subscribing raises. with pytest.raises(RuntimeError, match="already has a subscriber"): list(run.values) def test_abort_marks_exhausted_and_closes_mux(self) -> None: - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) values_iter = iter(run.values) - # Consume one item so the pump advances. _ = next(values_iter) run.abort() - # Remaining iteration yields whatever was buffered, then stops. list(values_iter) assert run._exhausted is True - # Second abort is idempotent. - run.abort() + run.abort() # idempotent def test_context_manager_calls_abort_on_exit(self) -> None: - graph = _build_simple_graph() - handler = graph - with handler.stream_v2({"value": "x", "items": []}) as run: - values_iter = iter(run.values) - _ = next(values_iter) + with _build_simple_graph().stream_v2({"value": "x", "items": []}) as run: + _ = next(iter(run.values)) assert run._exhausted is True def test_interleave_unknown_projection(self) -> None: - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) with pytest.raises(KeyError): list(run.interleave("values", "does_not_exist")) class TestStreamV2SyncErrors: def test_error_propagation_output(self) -> None: - graph = _build_error_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_error_graph().stream_v2({"value": "x", "items": []}) with pytest.raises(ValueError, match="boom"): _ = run.output def test_error_propagation_values(self) -> None: - graph = _build_error_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_error_graph().stream_v2({"value": "x", "items": []}) with pytest.raises(ValueError, match="boom"): list(run.values) def test_error_propagation_raw_events(self) -> None: - graph = _build_error_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_error_graph().stream_v2({"value": "x", "items": []}) with pytest.raises(ValueError, match="boom"): list(run) def test_error_propagation_interrupted(self) -> None: - """`run.interrupted` should raise on a failed run, not silently return False.""" - graph = _build_error_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_error_graph().stream_v2({"value": "x", "items": []}) with pytest.raises(ValueError, match="boom"): _ = run.interrupted def test_error_propagation_interrupts(self) -> None: - """`run.interrupts` should raise on a failed run.""" - graph = _build_error_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_error_graph().stream_v2({"value": "x", "items": []}) with pytest.raises(ValueError, match="boom"): _ = run.interrupts class TestStreamV2SyncInterrupt: def test_interrupted(self) -> None: - graph = _build_interrupt_graph() - handler = graph - run = handler.stream_v2( + run = _build_interrupt_graph().stream_v2( {"value": "x", "items": []}, {"configurable": {"thread_id": "t1"}}, ) @@ -573,153 +492,53 @@ class TestStreamV2SyncInterrupt: # --------------------------------------------------------------------------- +@pytest.mark.anyio +@NEEDS_CONTEXTVARS class TestStreamV2Async: - @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_values_projection(self) -> None: - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) snapshots = [s async for s in run.values] assert len(snapshots) >= 1 last = snapshots[-1] - assert "A" in last["value"] - assert "B" in last["value"] + assert "A" in last["value"] and "B" in last["value"] - @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_output(self) -> None: - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) output = await run.output() - assert output is not None - assert output["value"] == "xAB" - assert output["items"] == ["a", "b"] + assert output == {"value": "xAB", "items": ["a", "b"]} - @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_raw_event_iteration(self) -> None: - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) events = [e async for e in run] assert len(events) > 0 for event in events: assert event["type"] == "event" - @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_abort_marks_exhausted_and_closes_mux(self) -> None: - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) values_iter = aiter(run.values) _ = await anext(values_iter) await run.abort() - # Drain the rest; should terminate promptly now that mux is closed. async for _item in values_iter: pass assert run._exhausted is True await run.abort() # idempotent - @pytest.mark.anyio - @NEEDS_CONTEXTVARS - async def test_async_context_manager_calls_abort_on_exit(self) -> None: - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + async def test_context_manager_calls_abort_on_exit(self) -> None: + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) async with run: - values_iter = aiter(run.values) - _ = await anext(values_iter) + _ = await anext(aiter(run.values)) assert run._exhausted is True - @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_extensions_has_native_keys(self) -> None: - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) _ = await run.output() - assert "values" in run.extensions - assert "messages" in run.extensions + assert "values" in run.extensions and "messages" in run.extensions assert run.values is run.extensions["values"] assert run.messages is run.extensions["messages"] - -class TestStreamV2AsyncErrors: - @pytest.mark.anyio - @NEEDS_CONTEXTVARS - async def test_error_propagation_output(self) -> None: - graph = _build_error_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) - with pytest.raises(ValueError, match="boom"): - await run.output() - - @pytest.mark.anyio - @NEEDS_CONTEXTVARS - async def test_error_propagation_values(self) -> None: - graph = _build_error_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) - with pytest.raises(ValueError, match="boom"): - async for _ in run.values: - pass - - @pytest.mark.anyio - @NEEDS_CONTEXTVARS - async def test_error_propagation_raw_events(self) -> None: - graph = _build_error_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) - with pytest.raises(ValueError, match="boom"): - async for _ in run: - pass - - @pytest.mark.anyio - @NEEDS_CONTEXTVARS - async def test_error_propagation_interrupted(self) -> None: - """`await run.interrupted()` should raise on a failed async run.""" - graph = _build_error_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) - with pytest.raises(ValueError, match="boom"): - await run.interrupted() - - @pytest.mark.anyio - @NEEDS_CONTEXTVARS - async def test_error_propagation_interrupts(self) -> None: - """`await run.interrupts()` should raise on a failed async run.""" - graph = _build_error_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) - with pytest.raises(ValueError, match="boom"): - await run.interrupts() - - -class TestStreamV2AsyncInterrupt: - @pytest.mark.anyio - @NEEDS_CONTEXTVARS - async def test_interrupted(self) -> None: - graph = _build_interrupt_graph() - handler = graph - run = await handler.astream_v2( - {"value": "x", "items": []}, - {"configurable": {"thread_id": "t2"}}, - ) - _ = await run.output() - assert await run.interrupted() is True - assert len(await run.interrupts()) > 0 - - -class TestStreamV2AsyncCustom: - @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_custom_stream_events(self) -> None: - graph = _build_custom_stream_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_custom_stream_graph().astream_v2({"value": "x", "items": []}) events = [e async for e in run] custom_events = [e for e in events if e["method"] == "custom"] assert len(custom_events) == 2 @@ -727,9 +546,48 @@ class TestStreamV2AsyncCustom: assert custom_events[1]["params"]["data"] == {"step": "end"} -# --------------------------------------------------------------------------- -# Custom transformer tests -# --------------------------------------------------------------------------- +@pytest.mark.anyio +@NEEDS_CONTEXTVARS +class TestStreamV2AsyncErrors: + async def test_error_propagation_output(self) -> None: + run = await _build_error_graph().astream_v2({"value": "x", "items": []}) + with pytest.raises(ValueError, match="boom"): + await run.output() + + async def test_error_propagation_values(self) -> None: + run = await _build_error_graph().astream_v2({"value": "x", "items": []}) + with pytest.raises(ValueError, match="boom"): + async for _ in run.values: + pass + + async def test_error_propagation_raw_events(self) -> None: + run = await _build_error_graph().astream_v2({"value": "x", "items": []}) + with pytest.raises(ValueError, match="boom"): + async for _ in run: + pass + + async def test_error_propagation_interrupted(self) -> None: + run = await _build_error_graph().astream_v2({"value": "x", "items": []}) + with pytest.raises(ValueError, match="boom"): + await run.interrupted() + + async def test_error_propagation_interrupts(self) -> None: + run = await _build_error_graph().astream_v2({"value": "x", "items": []}) + with pytest.raises(ValueError, match="boom"): + await run.interrupts() + + +@pytest.mark.anyio +@NEEDS_CONTEXTVARS +class TestStreamV2AsyncInterrupt: + async def test_interrupted(self) -> None: + run = await _build_interrupt_graph().astream_v2( + {"value": "x", "items": []}, + {"configurable": {"thread_id": "t2"}}, + ) + _ = await run.output() + assert await run.interrupted() is True + assert len(await run.interrupts()) > 0 # --------------------------------------------------------------------------- @@ -740,8 +598,9 @@ class TestStreamV2AsyncCustom: class TestConvertToProtocolEvent: def test_basic_conversion(self) -> None: before = int(time.time() * 1000) - part = {"type": "values", "ns": ("sub", "graph"), "data": {"key": "val"}} - event = convert_to_protocol_event(part) + event = convert_to_protocol_event( + {"type": "values", "ns": ("sub", "graph"), "data": {"key": "val"}} + ) after = int(time.time() * 1000) assert event["type"] == "event" assert event["method"] == "values" @@ -751,21 +610,20 @@ class TestConvertToProtocolEvent: assert before <= event["params"]["timestamp"] <= after def test_conversion_with_interrupts(self) -> None: - part = { - "type": "values", - "ns": (), - "data": {"k": 1}, - "interrupts": ({"value": "pause"},), - } - event = convert_to_protocol_event(part) + event = convert_to_protocol_event( + { + "type": "values", + "ns": (), + "data": {"k": 1}, + "interrupts": ({"value": "pause"},), + } + ) assert event["params"]["interrupts"] == ({"value": "pause"},) - assert isinstance(event["params"]["timestamp"], int) def test_namespace_tuple_becomes_list(self) -> None: - """ns tuple should be converted to a list.""" - part = {"type": "updates", "ns": ("a", "b", "c"), "data": {}} - event = convert_to_protocol_event(part) - assert isinstance(event["params"]["namespace"], list) + event = convert_to_protocol_event( + {"type": "updates", "ns": ("a", "b", "c"), "data": {}} + ) assert event["params"]["namespace"] == ["a", "b", "c"] @@ -776,8 +634,6 @@ class TestConvertToProtocolEvent: class TestStreamMux: def test_register_non_dict_raises(self) -> None: - """init() returning a non-dict should raise TypeError at construction.""" - class BadTransformer(StreamTransformer): def init(self) -> Any: return ["not", "a", "dict"] @@ -789,33 +645,24 @@ class TestStreamMux: StreamMux([BadTransformer()]) def test_event_suppression(self) -> None: - """When process() returns False, the event should not appear in the main log.""" - class FilterTransformer(StreamTransformer): def init(self) -> dict[str, Any]: return {} def process(self, event: ProtocolEvent) -> bool: - # Suppress "updates" events return event["method"] != "updates" mux = StreamMux([FilterTransformer()]) it = iter(mux._events) - mux.push(_event("values", {"a": 1})) mux.push(_event("updates", {"b": 2})) mux.push(_event("custom", {"c": 3})) mux.close() + assert [e["method"] for e in it] == ["values", "custom"] - events = list(it) - methods = [e["method"] for e in events] - assert "updates" not in methods - assert methods == ["values", "custom"] - - def test_suppression_partial_transformers(self) -> None: - """If any transformer returns False, the event is suppressed, - but all transformers still see it.""" - + def test_suppression_all_transformers_still_see_event(self) -> None: + """If any transformer returns False, the event is suppressed from the main + log, but all transformers still receive it.""" seen_by_second: list[str] = [] class PassTransformer(StreamTransformer): @@ -834,17 +681,12 @@ class TestStreamMux: return False mux = StreamMux([PassTransformer(), RejectTransformer()]) - mux.push(_event("values")) mux.close() - - # RejectTransformer saw the event even though it rejected it assert seen_by_second == ["values"] - # But nothing in the main log assert list(mux._events) == [] def test_empty_mux(self) -> None: - """Push/close/fail on a mux with no transformers should work.""" mux = StreamMux() it = iter(mux._events) mux.push(_event("values", {"x": 1})) @@ -854,7 +696,6 @@ class TestStreamMux: assert events[0]["method"] == "values" def test_empty_mux_fail(self) -> None: - """Fail on an empty mux should propagate to the event log.""" mux = StreamMux() mux.fail(ValueError("boom")) with pytest.raises(ValueError, match="boom"): @@ -868,37 +709,29 @@ class TestStreamMux: class TestValuesTransformer: def test_ignores_non_root_namespace(self) -> None: - """Values events from subgraphs (non-empty namespace) should be ignored.""" t = ValuesTransformer() t.init() t._log._bind(is_async=False) it = iter(t._log) - t.process(_event("values", {"val": "root"})) t.process(_event("values", {"val": "sub"}, namespace=["sub"])) - t._log.close() items = list(it) assert len(items) == 1 assert items[0]["val"] == "root" def test_ignores_non_values_methods(self) -> None: - """Non-values events should be passed through but not captured.""" t = ValuesTransformer() t.init() t._log._bind(is_async=False) it = iter(t._log) - - result = t.process(_event("updates", {"x": 1})) - assert result is True # passed through + assert t.process(_event("updates", {"x": 1})) is True t._log.close() - assert list(it) == [] # but not captured + assert list(it) == [] def test_tracks_interrupts(self) -> None: - """Interrupts should be accumulated across events.""" t = ValuesTransformer() t.init() - t.process( _event( "values", @@ -915,21 +748,34 @@ class TestMessagesTransformer: t = MessagesTransformer() t.init() t._log._bind(is_async=False) + t._bind_pump(lambda: False) it = iter(t._log) - - t.process(_event("messages", ("chunk", {"meta": True}))) + meta = {"langgraph_node": "llm", "run_id": "run-1"} + for evt in ( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + {"event": "message-finish", "reason": "stop"}, + ): + t.process(_event("messages", (evt, meta))) t._log.close() items = list(it) assert len(items) == 1 - assert items[0] == ("chunk", {"meta": True}) + assert hasattr(items[0], "dispatch") + assert items[0].message_id == "run-1" def test_ignores_non_root_namespace(self) -> None: t = MessagesTransformer() t.init() t._log._bind(is_async=False) + t._bind_pump(lambda: False) it = iter(t._log) - - t.process(_event("messages", ("chunk", {}), namespace=["sub"])) + meta = {"langgraph_node": "llm", "run_id": "run-1"} + t.process( + _event( + "messages", + ({"event": "message-start", "message_id": "run-1"}, meta), + namespace=["sub"], + ) + ) t._log.close() assert list(it) == [] @@ -938,9 +784,7 @@ class TestMessagesTransformer: t.init() t._log._bind(is_async=False) it = iter(t._log) - - result = t.process(_event("values", {"v": 1})) - assert result is True + assert t.process(_event("values", {"v": 1})) is True t._log.close() assert list(it) == [] @@ -955,17 +799,12 @@ class TestMessagesTransformer: # --------------------------------------------------------------------------- -# StreamMux resilience tests +# StreamMux resilience: close/fail continue cleanup on transformer errors # --------------------------------------------------------------------------- class TestStreamMuxResilience: - """StreamMux.close() and fail() must complete cleanup even if a transformer raises.""" - def test_close_continues_after_finalize_error(self) -> None: - """If a transformer's finalize() raises, the main event log and - remaining transformers should still be closed/finalized.""" - class BrokenFinalizer(StreamTransformer): def init(self) -> dict[str, Any]: return {} @@ -991,19 +830,13 @@ class TestStreamMuxResilience: good = GoodTransformer() mux = StreamMux([BrokenFinalizer(), good]) - mux.push(_event("values")) - with pytest.raises(RuntimeError, match="finalize broke"): mux.close() - assert good.finalized assert mux._events._closed def test_fail_continues_after_transformer_error(self) -> None: - """If a transformer's fail() raises, the main event log and - remaining transformers should still be failed.""" - class BrokenFailer(StreamTransformer): def init(self) -> dict[str, Any]: return {} @@ -1029,16 +862,12 @@ class TestStreamMuxResilience: good = GoodTransformer() mux = StreamMux([BrokenFailer(), good]) - original_error = ValueError("original") mux.fail(original_error) - assert good.failed_with is original_error assert mux._events._error is original_error - def test_close_still_closes_channels_after_finalize_error(self) -> None: - """Channels should be closed even if a transformer's finalize raises.""" - + def test_channels_closed_after_finalize_error(self) -> None: class BrokenWithChannel(StreamTransformer): def __init__(self) -> None: self._channel: StreamChannel[str] = StreamChannel("ch") @@ -1054,17 +883,18 @@ class TestStreamMuxResilience: t = BrokenWithChannel() mux = StreamMux([t]) - with pytest.raises(RuntimeError, match="finalize broke"): mux.close() - assert t._channel._log._closed +# --------------------------------------------------------------------------- +# Custom transformer tests +# --------------------------------------------------------------------------- + + class TestCustomTransformer: def test_extension_transformer_with_stream_channel(self) -> None: - """User transformer with StreamChannel appears in extensions.""" - class CounterTransformer(StreamTransformer): def __init__(self) -> None: super().__init__() @@ -1080,22 +910,18 @@ class TestCustomTransformer: self._channel.push(self._count) return True - graph = _build_simple_graph() - handler = graph counter_t = CounterTransformer() - run = handler.stream_v2({"value": "x", "items": []}, transformers=[counter_t]) + run = _build_simple_graph().stream_v2( + {"value": "x", "items": []}, transformers=[counter_t] + ) assert "counter" in run.extensions - # Subscribe before driving the run so channel pushes are retained. counter_iter = iter(run.extensions["counter"]) _ = run.output counts = list(counter_iter) assert len(counts) > 0 - # Non-native transformer should not set direct attributes. - assert not hasattr(run, "counter") + assert not hasattr(run, "counter") # non-native: no direct attribute def test_native_transformer_gets_direct_attr(self) -> None: - """A transformer with _native=True gets its keys as run attributes.""" - class FooTransformer(StreamTransformer): _native = True @@ -1111,22 +937,17 @@ class TestCustomTransformer: self._log.push("saw_values") return True - graph = _build_simple_graph() - handler = graph foo_t = FooTransformer() - run = handler.stream_v2({"value": "x", "items": []}, transformers=[foo_t]) - # Subscribe before driving the run. + run = _build_simple_graph().stream_v2( + {"value": "x", "items": []}, transformers=[foo_t] + ) foo_iter = iter(run.foo) _ = run.output - # foo should be both in extensions and as a direct attribute. - assert "foo" in run.extensions - assert hasattr(run, "foo") - assert run.foo is run.extensions["foo"] - items = list(foo_iter) - assert "saw_values" in items + assert "foo" in run.extensions and run.foo is run.extensions["foo"] + assert "saw_values" in list(foo_iter) def test_stream_channel_auto_forward(self) -> None: - """StreamChannel pushes inject ProtocolEvents into main log.""" + """StreamChannel pushes inject ProtocolEvents into the main log.""" class EmitterTransformer(StreamTransformer): def __init__(self) -> None: @@ -1141,23 +962,16 @@ class TestCustomTransformer: self._channel.push("emitted") return True - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2( + run = _build_simple_graph().stream_v2( {"value": "x", "items": []}, transformers=[EmitterTransformer()] ) - events = list(run) - custom_events = [e for e in events if e["method"] == "custom:emitter"] + custom_events = [e for e in run if e["method"] == "custom:emitter"] assert len(custom_events) > 0 assert custom_events[0]["params"]["data"] == "emitted" def test_stream_channel_seq_ordering(self) -> None: - """Seq numbers in the main event log must be monotonically increasing. - - When a transformer pushes to a StreamChannel during process(), the - auto-forwarded event enters the main log before the original event. - The seq numbers must still be in order. - """ + """Seq numbers must be monotonically increasing even when a channel push + auto-forwards an event mid-pipeline.""" class ChannelPusher(StreamTransformer): def __init__(self) -> None: @@ -1168,27 +982,19 @@ class TestCustomTransformer: return {"ch": self._channel} def process(self, event: ProtocolEvent) -> bool: - # Push to channel during process — this triggers auto-forward - # which injects an event into the main log mid-pipeline. self._channel.push(f"saw:{event['method']}") return True mux = StreamMux([ChannelPusher()]) it = iter(mux._events) - mux.push(_event("values")) mux.push(_event("updates")) mux.close() - - events = list(it) - seqs = [e["seq"] for e in events] - # Seq numbers must be strictly increasing. + seqs = [e["seq"] for e in it] for i in range(1, len(seqs)): assert seqs[i] > seqs[i - 1], f"Seq out of order at index {i}: {seqs}" def test_projection_key_conflict_raises(self) -> None: - """User transformer that collides with a built-in key should raise.""" - class ConflictTransformer(StreamTransformer): def __init__(self) -> None: super().__init__() @@ -1200,22 +1006,19 @@ class TestCustomTransformer: def process(self, event: ProtocolEvent) -> bool: return True - graph = _build_simple_graph() - handler = graph - with pytest.raises( - ValueError, - match=r"conflict.*'values'.*ValuesTransformer", - ): - handler.stream_v2( - {"value": "x", "items": []}, - transformers=[ConflictTransformer()], + with pytest.raises(ValueError, match=r"conflict.*'values'.*ValuesTransformer"): + _build_simple_graph().stream_v2( + {"value": "x", "items": []}, transformers=[ConflictTransformer()] ) +# --------------------------------------------------------------------------- +# EventLog auto-lifecycle via StreamMux +# --------------------------------------------------------------------------- + + class TestEventLogAutoLifecycle: def test_mux_auto_closes_event_logs(self) -> None: - """EventLogs in projections should be auto-closed by mux.close().""" - class SimpleTransformer(StreamTransformer): def __init__(self) -> None: super().__init__() @@ -1230,17 +1033,11 @@ class TestEventLogAutoLifecycle: mux = StreamMux([SimpleTransformer()]) it = iter(mux._events) - mux.push(_event("values")) mux.close() - - # The log should have been auto-closed — iteration should work. - items = list(it) - assert len(items) == 1 + assert len(list(it)) == 1 def test_mux_auto_fails_event_logs(self) -> None: - """EventLogs in projections should be auto-failed by mux.fail().""" - class SimpleTransformer(StreamTransformer): def __init__(self) -> None: super().__init__() @@ -1256,17 +1053,12 @@ class TestEventLogAutoLifecycle: t = SimpleTransformer() mux = StreamMux([t]) it = iter(t._log) - mux.push(_event("values")) mux.fail(ValueError("boom")) - - # The transformer's log should have been auto-failed. with pytest.raises(ValueError, match="boom"): list(it) def test_no_double_close_if_transformer_closes_own_log(self) -> None: - """If a transformer closes its log in finalize(), mux should not error.""" - class ManualCloseTransformer(StreamTransformer): def __init__(self) -> None: super().__init__() @@ -1282,13 +1074,9 @@ class TestEventLogAutoLifecycle: self._log.close() mux = StreamMux([ManualCloseTransformer()]) - # Should not raise even though the log is closed by both - # the transformer and the mux. - mux.close() + mux.close() # should not raise even with double-close def test_transformer_without_finalize_works(self) -> None: - """Transformer with only init+process should work end-to-end.""" - class MinimalTransformer(StreamTransformer): def __init__(self) -> None: super().__init__() @@ -1302,30 +1090,49 @@ class TestEventLogAutoLifecycle: self._log.push("got_it") return True - graph = _build_simple_graph() - handler = graph t = MinimalTransformer() - run = handler.stream_v2({"value": "x", "items": []}, transformers=[t]) + run = _build_simple_graph().stream_v2( + {"value": "x", "items": []}, transformers=[t] + ) minimal_iter = iter(run.extensions["minimal"]) _ = run.output - items = list(minimal_iter) - assert len(items) > 0 + assert len(list(minimal_iter)) > 0 + + +class TestStreamTransformerSchedule: + def test_schedule_without_running_loop_raises(self) -> None: + class Sched(StreamTransformer): + requires_async = True + + def init(self) -> dict[str, Any]: + return {} + + def process(self, event: ProtocolEvent) -> bool: + return True + + t = Sched() + + async def noop() -> None: + pass + + coro = noop() + try: + with pytest.raises(RuntimeError, match="requires a running event loop"): + t.schedule(coro) + finally: + coro.close() # --------------------------------------------------------------------------- -# Async transformer lane — aprocess / afinalize / afail / schedule() +# Async transformer lane # --------------------------------------------------------------------------- +@pytest.mark.anyio class TestAsyncTransformerLane: - @pytest.mark.anyio async def test_aprocess_is_awaited_before_next_transformer(self) -> None: - """aprocess must complete before the next transformer sees the event. - - This is the load-bearing guarantee for mutating transformers - like PII redaction: the downstream transformer reads the mutated - event synchronously. - """ + """aprocess must complete before the next transformer sees the event — + load-bearing guarantee for mutating transformers like PII redaction.""" order: list[str] = [] class RedactTransformer(StreamTransformer): @@ -1350,13 +1157,10 @@ class TestAsyncTransformerLane: return True mux = StreamMux([RedactTransformer(), ObserverTransformer()], is_async=True) - await mux.apush(_event("values", {"secret": "x"})) await mux.aclose() - assert order == ["redact", "observe:True"] - @pytest.mark.anyio async def test_schedule_joins_tasks_before_afinalize(self) -> None: """Every scheduled task must complete before afinalize runs.""" phase: list[str] = [] @@ -1387,19 +1191,13 @@ class TestAsyncTransformerLane: t = SchedTransformer() mux = StreamMux([t], is_async=True) - await mux.apush(_event("values", {})) await mux.apush(_event("values", {})) await mux.aclose() - - # Both tasks ran before afinalize; the log holds both pushes. assert phase.count("task") == 2 assert phase[-1] == "afinalize" - @pytest.mark.anyio async def test_sync_stream_rejects_async_transformer(self) -> None: - """Registering a requires_async transformer on a sync mux raises.""" - class NeedsAsync(StreamTransformer): requires_async = True @@ -1412,10 +1210,7 @@ class TestAsyncTransformerLane: with pytest.raises(RuntimeError, match="requires an async run"): StreamMux([NeedsAsync()], is_async=False) - @pytest.mark.anyio async def test_sync_stream_rejects_aprocess_override(self) -> None: - """Overriding aprocess also marks the transformer as async-required.""" - class HasAprocess(StreamTransformer): def init(self) -> dict[str, Any]: return {} @@ -1426,34 +1221,7 @@ class TestAsyncTransformerLane: with pytest.raises(RuntimeError, match="requires an async run"): StreamMux([HasAprocess()], is_async=False) - def test_schedule_without_running_loop_raises(self) -> None: - """schedule() called outside an event loop fails with a clear message.""" - - class Sched(StreamTransformer): - requires_async = True - - def init(self) -> dict[str, Any]: - return {} - - def process(self, event: ProtocolEvent) -> bool: - return True - - t = Sched() - - async def noop() -> None: - pass - - coro = noop() - try: - with pytest.raises(RuntimeError, match="requires a running event loop"): - t.schedule(coro) - finally: - coro.close() - - @pytest.mark.anyio async def test_schedule_on_error_log_swallows_exceptions(self) -> None: - """on_error="log" (default) keeps the run alive when a task fails.""" - class Bad(StreamTransformer): requires_async = True @@ -1474,15 +1242,10 @@ class TestAsyncTransformerLane: self._log.close() mux = StreamMux([Bad()], is_async=True) - await mux.apush(_event("values", {})) - # Should not raise; the scheduled task's exception is logged. - await mux.aclose() + await mux.aclose() # should not raise; exception is logged - @pytest.mark.anyio async def test_schedule_on_error_raise_fails_the_run(self) -> None: - """on_error="raise" propagates the exception through aclose.""" - class Strict(StreamTransformer): requires_async = True @@ -1497,14 +1260,11 @@ class TestAsyncTransformerLane: return True mux = StreamMux([Strict()], is_async=True) - await mux.apush(_event("values", {})) with pytest.raises(ValueError, match="strict boom"): await mux.aclose() - @pytest.mark.anyio async def test_afail_cancels_pending_scheduled_tasks(self) -> None: - """When the run fails, outstanding scheduled tasks are cancelled.""" cancelled = asyncio.Event() class Sched(StreamTransformer): @@ -1525,19 +1285,13 @@ class TestAsyncTransformerLane: return True mux = StreamMux([Sched()], is_async=True) - await mux.apush(_event("values", {})) - # Yield so the scheduled task actually starts before we cancel it; - # otherwise it's cancelled before its first step and the `except` - # inside work() never runs. + # Yield so the task actually starts before we cancel it. await asyncio.sleep(0) await mux.afail(RuntimeError("run died")) - assert cancelled.is_set() - @pytest.mark.anyio async def test_mixed_sync_and_async_transformers(self) -> None: - """Sync and async transformers coexist under astream.""" seen_sync: list[str] = [] class SyncOne(StreamTransformer): @@ -1568,19 +1322,13 @@ class TestAsyncTransformerLane: async_t = AsyncOne() mux = StreamMux([SyncOne(), async_t], is_async=True) seen_cursor = aiter(async_t._log) - await mux.apush(_event("values", {})) await mux.apush(_event("updates", {})) await mux.aclose() - assert seen_sync == ["values", "updates"] - items = [x async for x in seen_cursor] - assert items == ["values", "updates"] + assert [x async for x in seen_cursor] == ["values", "updates"] - @pytest.mark.anyio async def test_handler_astream_with_scheduled_work(self) -> None: - """End-to-end: transformer schedules work during an astream run.""" - class Scorer(StreamTransformer): requires_async = True @@ -1603,13 +1351,9 @@ class TestAsyncTransformerLane: async def afinalize(self) -> None: self._log.close() - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2( - {"value": "x", "items": []}, - transformers=[Scorer()], + run = await _build_simple_graph().astream_v2( + {"value": "x", "items": []}, transformers=[Scorer()] ) - # Subscribe before driving the run so scheduled pushes are retained. scores_cursor = aiter(run.extensions["scores"]) _ = await run.output() scores = [x async for x in scores_cursor] @@ -1617,21 +1361,14 @@ class TestAsyncTransformerLane: # --------------------------------------------------------------------------- -# Drain-on-consume semantics — bounded backpressure, single subscriber +# Memory bounds: drain-on-consume semantics # --------------------------------------------------------------------------- +@NEEDS_CONTEXTVARS class TestMemoryBounds: - """Drain-on-consume guarantees memory stays bounded for the common - access patterns. These tests lock in the property — if a change - re-introduces retention, they should fail.""" - def test_sync_subscribed_buffer_stays_at_most_one_between_yields(self) -> None: - """With a single sync consumer, the pump produces exactly one event - per cursor advance, so the buffer never holds more than one.""" - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) events_iter = iter(run) max_buffered = 0 count = 0 @@ -1640,53 +1377,31 @@ class TestMemoryBounds: count += 1 assert count > 0 assert max_buffered == 0, ( - f"Subscribed buffer should hold 0 items after each yield " - f"(drain-on-consume), observed max {max_buffered}" + f"drain-on-consume violated, observed max {max_buffered}" ) def test_unsubscribed_projections_never_accumulate(self) -> None: - """Projections without a subscriber drop pushes silently — - their buffers stay empty regardless of run length.""" - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) - # Subscribe to main events only; leave values and messages unsubscribed. + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) list(run) values_log = run.extensions["values"] messages_log = run.extensions["messages"] - assert len(values_log._items) == 0 - assert len(messages_log._items) == 0 - assert values_log._subscribed is False - assert messages_log._subscribed is False + assert len(values_log._items) == 0 and not values_log._subscribed + assert len(messages_log._items) == 0 and not messages_log._subscribed def test_output_path_does_not_retain_values(self) -> None: - """`run.output` is a scalar accessor — it updates `_latest` from - process() without populating the log, so the values log buffer - stays empty even across a full run.""" - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) _ = run.output values_log = run.extensions["values"] - assert len(values_log._items) == 0 - assert values_log._subscribed is False + assert len(values_log._items) == 0 and not values_log._subscribed def test_drained_subscriber_buffer_returns_to_empty(self) -> None: - """After fully draining a subscribed log, the internal deque is empty.""" - graph = _build_simple_graph() - handler = graph - run = handler.stream_v2({"value": "x", "items": []}) - values_log = run.extensions["values"] + run = _build_simple_graph().stream_v2({"value": "x", "items": []}) list(run.values) - assert len(values_log._items) == 0 + assert len(run.extensions["values"]._items) == 0 @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_async_single_consumer_buffer_stays_at_most_one(self) -> None: - """Same drain-on-consume guarantee for the async lane.""" - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) max_buffered = 0 count = 0 async for _ in run: @@ -1696,19 +1411,18 @@ class TestMemoryBounds: assert max_buffered == 0 @pytest.mark.anyio - @NEEDS_CONTEXTVARS async def test_async_unsubscribed_projections_never_accumulate(self) -> None: - """Projections with no async subscriber stay empty under astream.""" - graph = _build_simple_graph() - handler = graph - run = await handler.astream_v2({"value": "x", "items": []}) + run = await _build_simple_graph().astream_v2({"value": "x", "items": []}) _ = await run.output() values_log = run.extensions["values"] messages_log = run.extensions["messages"] - assert len(values_log._items) == 0 - assert len(messages_log._items) == 0 - assert values_log._subscribed is False - assert messages_log._subscribed is False + assert len(values_log._items) == 0 and not values_log._subscribed + assert len(messages_log._items) == 0 and not messages_log._subscribed + + +# --------------------------------------------------------------------------- +# DrainOnConsume: EventLog capacity semantics +# --------------------------------------------------------------------------- class TestDrainOnConsume: @@ -1719,8 +1433,7 @@ class TestDrainOnConsume: EventLog(maxlen=-3) def test_push_unbounded_by_design(self) -> None: - """Push is non-blocking and doesn't enforce capacity — the - caller-driven pump bounds memory via iteration pace.""" + """Push is non-blocking; the caller-driven pump bounds memory via iteration pace.""" log: EventLog[int] = EventLog() log._bind(is_async=False) it = iter(log) @@ -1729,20 +1442,6 @@ class TestDrainOnConsume: log.close() assert list(it) == list(range(100)) - @pytest.mark.anyio - async def test_atee_fans_out(self) -> None: - """atee provides the documented fan-out for concurrent consumers.""" - log: EventLog[int] = EventLog() - log._bind(is_async=True) - a, b = log.atee(2) - for i in range(3): - log.push(i) - log.close() - items_a = [x async for x in a] - items_b = [x async for x in b] - assert items_a == [0, 1, 2] - assert items_b == [0, 1, 2] - def test_tee_fans_out_sync(self) -> None: log: EventLog[int] = EventLog() log._bind(is_async=False) @@ -1752,3 +1451,14 @@ class TestDrainOnConsume: log.close() assert list(a) == [0, 1, 2] assert list(b) == [0, 1, 2] + + @pytest.mark.anyio + async def test_atee_fans_out(self) -> None: + log: EventLog[int] = EventLog() + log._bind(is_async=True) + a, b = log.atee(2) + for i in range(3): + log.push(i) + log.close() + assert [x async for x in a] == [0, 1, 2] + assert [x async for x in b] == [0, 1, 2] diff --git a/libs/langgraph/tests/test_stream_messages_transformer.py b/libs/langgraph/tests/test_stream_messages_transformer.py new file mode 100644 index 000000000..b2f4f9c65 --- /dev/null +++ b/libs/langgraph/tests/test_stream_messages_transformer.py @@ -0,0 +1,872 @@ +"""Tests for MessagesTransformer: protocol event routing, whole-message fallback, +legacy v1 chunk filtering, and end-to-end via stream_v2 / astream_v2.""" + +from __future__ import annotations + +import time +from typing import Any + +import pytest +from langchain_core.language_models import GenericFakeChatModel +from langchain_core.language_models.chat_model_stream import ( + AsyncChatModelStream, + ChatModelStream, +) +from langchain_core.messages import AIMessage, AIMessageChunk +from langchain_core.runnables import RunnableConfig +from typing_extensions import TypedDict + +from langgraph.constants import END, START +from langgraph.graph import MessagesState, StateGraph +from langgraph.stream._event_log import EventLog +from langgraph.stream._mux import StreamMux +from langgraph.stream.run_stream import GraphRunStream +from langgraph.stream.transformers import MessagesTransformer, ValuesTransformer + +TS = int(time.time() * 1000) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _proto_event( + event: dict[str, Any], + *, + run_id: str = "run-1", + node: str = "llm", +) -> dict[str, Any]: + """Build a messages ProtocolEvent carrying a protocol event dict (v2 path).""" + return { + "type": "event", + "method": "messages", + "params": { + "namespace": [], + "timestamp": TS, + "data": (event, {"langgraph_node": node, "run_id": run_id}), + }, + } + + +def _v1_chunk( + text: str, + msg_id: str = "msg-1", + *, + finish: bool = False, + node: str = "llm", +) -> dict[str, Any]: + """Build a messages ProtocolEvent carrying a v1 AIMessageChunk tuple.""" + rm: dict[str, Any] = {"finish_reason": "stop"} if finish else {} + return { + "type": "event", + "method": "messages", + "params": { + "namespace": [], + "timestamp": TS, + "data": ( + AIMessageChunk(content=text, id=msg_id, response_metadata=rm), + {"langgraph_node": node}, + ), + }, + } + + +def _whole_msg( + text: str, + msg_id: str = "msg-10", + *, + node: str = "node", +) -> dict[str, Any]: + """Build a messages ProtocolEvent carrying a completed AIMessage.""" + return { + "type": "event", + "method": "messages", + "params": { + "namespace": [], + "timestamp": TS, + "data": (AIMessage(content=text, id=msg_id), {"langgraph_node": node}), + }, + } + + +def _make_sync_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelStream]]: + t = MessagesTransformer() + log: EventLog[ChatModelStream] = t.init()["messages"] + log._bind(is_async=False) + # Subscribe up front so pushes during process() are retained. + log._subscribed = True + t._bind_pump(lambda: False) + return t, log + + +def _make_async_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelStream]]: + t = MessagesTransformer() + log: EventLog[ChatModelStream] = t.init()["messages"] + log._bind(is_async=True) + log._subscribed = True + return t, log + + +def _lifecycle( + *, text: str = "hello world", message_id: str = "run-1" +) -> list[dict[str, Any]]: + """Produce a valid protocol event lifecycle: start, delta, finish.""" + half = len(text) // 2 + first, second = text[:half], text[half:] + return [ + {"event": "message-start", "role": "ai", "message_id": message_id}, + { + "event": "content-block-start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + { + "event": "content-block-delta", + "index": 0, + "content_block": {"type": "text", "text": first}, + }, + { + "event": "content-block-delta", + "index": 0, + "content_block": {"type": "text", "text": second}, + }, + { + "event": "content-block-finish", + "index": 0, + "content_block": {"type": "text", "text": text}, + }, + {"event": "message-finish", "reason": "stop"}, + ] + + +def _simple_graph(): + def call_model(state: MessagesState) -> dict[str, Any]: + model = GenericFakeChatModel(messages=iter(["hello world"])) + stream = model.stream_v2(state["messages"]) + return {"messages": stream.output} + + return ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + +# --------------------------------------------------------------------------- +# Protocol event routing +# --------------------------------------------------------------------------- + + +class TestProtocolEventRouting: + def test_message_start_creates_stream(self) -> None: + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + run_id="run-1", + ) + ) + log.close() + (stream,) = list(log._items) + assert isinstance(stream, ChatModelStream) + assert stream.message_id == "run-1" + + def test_full_lifecycle_yields_done_stream(self) -> None: + t, log = _make_sync_transformer() + for evt in _lifecycle(text="hello world"): + t.process(_proto_event(evt, run_id="run-1")) + log.close() + (stream,) = list(log._items) + assert stream.done + assert stream.output.text == "hello world" + + def test_message_finish_cleans_up_routing(self) -> None: + t, log = _make_sync_transformer() + for evt in _lifecycle(): + t.process(_proto_event(evt, run_id="run-1")) + assert t._by_run == {} + + def test_events_without_prior_start_are_ignored(self) -> None: + t, log = _make_sync_transformer() + t.process( + _proto_event( + { + "event": "content-block-delta", + "index": 0, + "content_block": {"type": "text", "text": "orphan"}, + }, + run_id="unknown", + ) + ) + log.close() + assert list(log._items) == [] + + def test_concurrent_streams_routed_by_run_id(self) -> None: + t, log = _make_sync_transformer() + life_a = _lifecycle(text="aaaa", message_id="run-a") + life_b = _lifecycle(text="bbbb", message_id="run-b") + for a, b in zip(life_a, life_b): + t.process(_proto_event(a, run_id="run-a")) + t.process(_proto_event(b, run_id="run-b")) + log.close() + streams = list(log._items) + assert len(streams) == 2 + by_id = {s.message_id: s for s in streams} + assert by_id["run-a"].output.text == "aaaa" + assert by_id["run-b"].output.text == "bbbb" + + def test_text_deltas_accumulated_on_stream(self) -> None: + t, log = _make_sync_transformer() + for evt in _lifecycle(text="abcdef"): + t.process(_proto_event(evt)) + log.close() + (stream,) = list(log._items) + assert "".join(stream._text_proj._deltas) == "abcdef" + + def test_stream_pushed_on_message_start_not_finish(self) -> None: + # Consumer can see the stream before message-finish arrives. + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + run_id="run-1", + ) + ) + assert len(log._items) == 1 + + def test_node_metadata_set_on_stream(self) -> None: + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + run_id="run-1", + node="my_llm", + ) + ) + (stream,) = list(log._items) + assert stream.node == "my_llm" + + +# --------------------------------------------------------------------------- +# Whole-message fallback +# --------------------------------------------------------------------------- + + +class TestWholeMessageFallback: + def test_whole_ai_message_produces_complete_stream(self) -> None: + t, log = _make_sync_transformer() + t.process(_whole_msg("the full answer")) + log.close() + (stream,) = list(log._items) + assert stream.done + assert stream.output.text == "the full answer" + + def test_whole_message_has_full_lifecycle(self) -> None: + t, log = _make_sync_transformer() + t.process(_whole_msg("full")) + log.close() + (stream,) = list(log._items) + assert [e["event"] for e in stream._events] == [ + "message-start", + "content-block-start", + "content-block-delta", + "content-block-finish", + "message-finish", + ] + + +# --------------------------------------------------------------------------- +# Filtering +# --------------------------------------------------------------------------- + + +class TestFiltering: + def test_non_messages_events_pass_through(self) -> None: + t, _ = _make_sync_transformer() + assert ( + t.process( + { + "type": "event", + "method": "values", + "params": {"namespace": [], "timestamp": TS, "data": {"x": 1}}, + } + ) + is True + ) + + def test_subgraph_namespace_dropped(self) -> None: + t, log = _make_sync_transformer() + t.process( + { + "type": "event", + "method": "messages", + "params": { + "namespace": ["subgraph"], + "timestamp": TS, + "data": ( + {"event": "message-start", "message_id": "run-x"}, + {"run_id": "run-x"}, + ), + }, + } + ) + log.close() + assert list(log._items) == [] + + def test_legacy_v1_chunks_ignored(self) -> None: + # v1 AIMessageChunk tuples (from on_llm_new_token) are not streamed + # into this projection; callers must migrate to stream_v2. + t, log = _make_sync_transformer() + t.process(_v1_chunk("hello")) + t.process(_v1_chunk(" world", finish=True)) + log.close() + assert list(log._items) == [] + + +# --------------------------------------------------------------------------- +# Lifecycle: fail / finalize +# --------------------------------------------------------------------------- + + +class TestLifecycle: + def test_fail_propagates_to_open_streams(self) -> None: + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "message_id": "run-1"}, run_id="run-1" + ) + ) + streams = list(log._items) + err = RuntimeError("graph died") + t.fail(err) + assert t._by_run == {} + assert streams[0]._error is err + + def test_finalize_clears_routing_state(self) -> None: + t, _ = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "message_id": "run-1"}, run_id="run-1" + ) + ) + assert "run-1" in t._by_run + t.finalize() + assert t._by_run == {} + + +# --------------------------------------------------------------------------- +# Async mode +# --------------------------------------------------------------------------- + + +class TestAsyncMode: + def test_async_mode_creates_async_stream(self) -> None: + t, log = _make_async_transformer() + for evt in _lifecycle(text="async stream"): + t.process(_proto_event(evt)) + assert isinstance(list(log._items)[0], AsyncChatModelStream) + + @pytest.mark.anyio + async def test_text_projection_yields_deltas(self) -> None: + t, log = _make_async_transformer() + for evt in _lifecycle(text="hello world"): + t.process(_proto_event(evt)) + (stream,) = list(log._items) + assert isinstance(stream, AsyncChatModelStream) + assert "".join([d async for d in stream.text]) == "hello world" + + @pytest.mark.anyio + async def test_output_awaitable(self) -> None: + t, log = _make_async_transformer() + for evt in _lifecycle(text="async"): + t.process(_proto_event(evt)) + (stream,) = list(log._items) + assert (await stream.output).text == "async" + + +# --------------------------------------------------------------------------- +# GraphRunStream integration +# --------------------------------------------------------------------------- + + +class TestWireRequestMore: + def test_bind_pump_called_on_wire(self) -> None: + values_t = ValuesTransformer() + messages_t = MessagesTransformer() + mux = StreamMux([values_t, messages_t], is_async=False) + + assert messages_t._pump_fn is None + run = GraphRunStream(iter([]), mux, values_t) + assert messages_t._pump_fn is not None + assert messages_t._pump_fn() is False + assert run._exhausted + + def test_created_streams_have_request_more(self) -> None: + values_t = ValuesTransformer() + messages_t = MessagesTransformer() + mux = StreamMux([values_t, messages_t], is_async=False) + GraphRunStream(iter([]), mux, values_t) + + log: EventLog[ChatModelStream] = mux.extensions["messages"] + log._subscribed = True + for evt in _lifecycle(): + messages_t.process(_proto_event(evt)) + + (stream,) = list(log._items) + assert stream._request_more is messages_t._pump_fn + + +# --------------------------------------------------------------------------- +# End-to-end via StreamMux +# --------------------------------------------------------------------------- + + +class TestViaMux: + def _make_mux( + self, + ) -> tuple[MessagesTransformer, StreamMux, EventLog[ChatModelStream]]: + t = MessagesTransformer() + v = ValuesTransformer() + mux = StreamMux([v, t], is_async=False) + t._bind_pump(lambda: False) + log: EventLog[ChatModelStream] = mux.extensions["messages"] + log._subscribed = True + return t, mux, log + + def test_streaming_via_mux(self) -> None: + t, mux, log = self._make_mux() + for evt in _lifecycle(text="mux stream"): + mux.push(_proto_event(evt)) + mux.close() + (stream,) = list(log._items) + assert stream.output.text == "mux stream" + + def test_whole_message_via_mux(self) -> None: + t, mux, log = self._make_mux() + mux.push(_whole_msg("result")) + mux.close() + (stream,) = list(log._items) + assert stream.output.text == "result" + + @pytest.mark.anyio + async def test_async_streaming_via_mux(self) -> None: + t = MessagesTransformer() + v = ValuesTransformer() + mux = StreamMux([v, t], is_async=True) + log: EventLog[ChatModelStream] = mux.extensions["messages"] + log._subscribed = True + + for evt in _lifecycle(text="async mux"): + await mux.apush(_proto_event(evt)) + + (stream,) = list(log._items) + assert (await stream.output).text == "async mux" + await mux.aclose() + + +# --------------------------------------------------------------------------- +# End-to-end: graph → stream_v2 → run.messages (node calls stream_v2) +# --------------------------------------------------------------------------- + + +class TestEndToEnd: + """stream_v2 path: node calls model.stream_v2() explicitly.""" + + def test_node_calling_stream_v2_populates_messages(self) -> None: + model = GenericFakeChatModel(messages=iter(["hello world"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + stream = model.stream_v2(state["messages"]) + return {"messages": stream.output} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + run = graph.stream_v2({"messages": "hi"}) + (stream,) = list(run.messages) + assert isinstance(stream, ChatModelStream) + assert stream.output.text == "hello world" + + def test_node_stream_v2_text_deltas_iterate(self) -> None: + """Consumer can iterate `.text` on the streamed message in real time.""" + model = GenericFakeChatModel(messages=iter(["streamed answer"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + stream = model.stream_v2(state["messages"]) + return {"messages": stream.output} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + run = graph.stream_v2({"messages": "go"}) + (stream,) = list(run.messages) + assert "".join(stream.text) == "streamed answer" + + def test_non_llm_message_returned_from_node(self) -> None: + """Whole-message fallback: node returns a finalized AIMessage directly.""" + + def return_message(state: MessagesState) -> dict[str, Any]: + return {"messages": AIMessage(content="hardcoded", id="msg-abc")} + + graph = ( + StateGraph(MessagesState) + .add_node("return_message", return_message) + .add_edge(START, "return_message") + .add_edge("return_message", END) + .compile() + ) + + run = graph.stream_v2({"messages": "hi"}) + (stream,) = list(run.messages) + assert stream.output.text == "hardcoded" + + @pytest.mark.anyio + async def test_async_node_calling_astream_v2(self) -> None: + model = GenericFakeChatModel(messages=iter(["async answer"])) + + async def call_model(state: MessagesState) -> dict[str, Any]: + stream = await model.astream_v2(state["messages"]) + return {"messages": await stream} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + run = await graph.astream_v2({"messages": "hi"}) + streams = [s async for s in run.messages] + assert len(streams) == 1 + assert isinstance(streams[0], AsyncChatModelStream) + assert (await streams[0].output).text == "async answer" + + @pytest.mark.anyio + async def test_nested_async_iteration_yields_text_deltas(self) -> None: + """Inner stream.text drives the shared graph pump via the async pump binding.""" + import asyncio + + model = GenericFakeChatModel(messages=iter(["hello world"])) + + async def call_model(state: MessagesState) -> dict[str, Any]: + stream = await model.astream_v2(state["messages"]) + return {"messages": await stream} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + run = await graph.astream_v2({"messages": "hi"}) + + async def consume() -> list[str]: + collected: list[str] = [] + async for stream in run.messages: + async for delta in stream.text: + collected.append(delta) + return collected + + assert "".join(await asyncio.wait_for(consume(), timeout=2.0)) == "hello world" + + +# --------------------------------------------------------------------------- +# End-to-end: graph → stream_v2 → run.messages (node calls invoke) +# --------------------------------------------------------------------------- + + +class TestEndToEndV2Invoke: + """Auto-routing path: stream_v2 injects CONFIG_KEY_STREAM_MESSAGES_V2, + causing BaseChatModel to drive the v2 protocol event generator even for + model.invoke().""" + + def _graph(self, model): + def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + return ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + def test_invoke_populates_messages(self) -> None: + run = self._graph( + GenericFakeChatModel(messages=iter(["hello world"])) + ).stream_v2({"messages": "hi"}) + (stream,) = list(run.messages) + assert isinstance(stream, ChatModelStream) + assert stream.output.text == "hello world" + + def test_invoke_emits_protocol_events(self) -> None: + """Iterating the stream yields the full v2 lifecycle, not v1 chunks.""" + run = self._graph( + GenericFakeChatModel(messages=iter(["streamed answer"])) + ).stream_v2({"messages": "go"}) + (stream,) = list(run.messages) + + events = list(stream) + event_types = [e.get("event") for e in events] + assert "message-start" in event_types + assert "content-block-start" in event_types + assert "content-block-delta" in event_types + assert "content-block-finish" in event_types + assert "message-finish" in event_types + # Sanity: every event is a dict carrying an "event" key — not an + # AIMessageChunk tuple from the v1 path. + for event in events: + assert isinstance(event, dict) + assert "event" in event + # Typed projection still assembles the final text. + assert stream.output.text == "streamed answer" + + def test_invoke_text_deltas_iterate(self) -> None: + run = self._graph( + GenericFakeChatModel(messages=iter(["delta streaming works"])) + ).stream_v2({"messages": "hi"}) + (stream,) = list(run.messages) + assert "".join(stream.text) == "delta streaming works" + + def test_invoke_two_nodes_two_streams(self) -> None: + model_a = GenericFakeChatModel(messages=iter(["alpha"])) + model_b = GenericFakeChatModel(messages=iter(["beta"])) + + def node_a(state: MessagesState) -> dict[str, Any]: + return {"messages": model_a.invoke(state["messages"])} + + def node_b(state: MessagesState) -> dict[str, Any]: + return {"messages": model_b.invoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("node_a", node_a) + .add_node("node_b", node_b) + .add_edge(START, "node_a") + .add_edge("node_a", "node_b") + .add_edge("node_b", END) + .compile() + ) + + streams = list(graph.stream_v2({"messages": "hi"}).messages) + assert len(streams) == 2 + assert {s.output.text for s in streams} == {"alpha", "beta"} + + def test_invoke_plus_constructed_message_two_streams(self) -> None: + """Live-streamed node + constructed-message node → two ChatModelStreams.""" + model = GenericFakeChatModel(messages=iter(["live stream"])) + + def streaming_node(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + def constructed_node(state: MessagesState) -> dict[str, Any]: + return {"messages": [AIMessage(content="hardcoded", id="constructed-1")]} + + graph = ( + StateGraph(MessagesState) + .add_node("streaming_node", streaming_node) + .add_node("constructed_node", constructed_node) + .add_edge(START, "streaming_node") + .add_edge("streaming_node", "constructed_node") + .add_edge("constructed_node", END) + .compile() + ) + + run = graph.stream_v2({"messages": "hi"}) + streams = list(run.messages) + assert len(streams) == 2 + assert streams[0].node == "streaming_node" + assert streams[0].output.text == "live stream" + assert streams[1].node == "constructed_node" + assert streams[1].output.text == "hardcoded" + assert streams[1].message_id == "constructed-1" + + @pytest.mark.anyio + async def test_ainvoke_populates_messages(self) -> None: + model = GenericFakeChatModel(messages=iter(["async invoke"])) + + async def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": await model.ainvoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + run = await graph.astream_v2({"messages": "hi"}) + streams = [s async for s in run.messages] + assert len(streams) == 1 + assert isinstance(streams[0], AsyncChatModelStream) + assert (await streams[0].output).text == "async invoke" + + +# --------------------------------------------------------------------------- +# Regression: direct stream_mode="messages" must stay v1 +# --------------------------------------------------------------------------- + + +class TestDirectMessagesModeStaysV1: + def test_direct_graph_stream_messages_yields_ai_message_chunks(self) -> None: + """graph.stream(stream_mode="messages") must not leak v2 event dicts — + the v2 flag is only injected by stream_v2 / astream_v2.""" + model = GenericFakeChatModel(messages=iter(["legacy path"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + parts = list(graph.stream({"messages": "hi"}, stream_mode="messages")) + assert parts, "expected stream_mode='messages' to emit tuples" + for payload, _metadata in parts: + assert isinstance(payload, AIMessageChunk) + assert ( + "".join(p[0].content for p in parts if isinstance(p[0].content, str)) + == "legacy path" + ) + + def test_nested_graph_stream_messages_stays_v1_under_outer_stream_v2(self) -> None: + """An outer `stream_v2()` run must not flip an inner direct + `stream_mode="messages"` call onto the v2 event protocol.""" + model = GenericFakeChatModel(messages=iter(["nested legacy path"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + inner = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + class OuterState(TypedDict, total=False): + saw_only_chunks: bool + first_payload_type: str + text: str + + def call_subgraph(state: OuterState, config: RunnableConfig) -> dict[str, Any]: + parts = list( + inner.stream( + {"messages": "hi"}, + config, + stream_mode="messages", + ) + ) + assert parts + payloads = [payload for payload, _metadata in parts] + return { + "saw_only_chunks": all( + isinstance(payload, AIMessageChunk) for payload in payloads + ), + "first_payload_type": type(payloads[0]).__name__, + "text": "".join( + payload.content + for payload in payloads + if isinstance(payload, AIMessageChunk) + and isinstance(payload.content, str) + ), + } + + outer = ( + StateGraph(OuterState) + .add_node("call_subgraph", call_subgraph) + .add_edge(START, "call_subgraph") + .add_edge("call_subgraph", END) + .compile() + ) + + result = outer.stream_v2({}).output + + assert result is not None + assert result["saw_only_chunks"] is True + assert result["first_payload_type"] == "AIMessageChunk" + assert result["text"] == "nested legacy path" + + +# --------------------------------------------------------------------------- +# StreamMessagesHandlerV2 unit +# --------------------------------------------------------------------------- + + +class TestStreamMessagesHandlerV2Unit: + def test_on_llm_new_token_is_noop(self) -> None: + """v2 handler must not emit v1 chunks even when on_llm_new_token fires.""" + from uuid import uuid4 + + from langchain_core.outputs import ChatGenerationChunk + + from langgraph.pregel._messages import StreamMessagesHandlerV2 + + emitted: list[Any] = [] + handler = StreamMessagesHandlerV2(emitted.append, subgraphs=False) + run_id = uuid4() + handler.metadata[run_id] = ((), {"langgraph_node": "x"}) + + handler.on_llm_new_token( + "hello", + chunk=ChatGenerationChunk(message=AIMessageChunk(content="hello")), + run_id=run_id, + ) + + assert emitted == [] + + def test_on_llm_end_dedupes_when_final_message_id_differs(self) -> None: + """A streamed v2 message should not be emitted again from the final + AIMessage fallback when its final id does not match `message-start`.""" + from uuid import uuid4 + + from langchain_core.outputs import ChatGeneration, LLMResult + + from langgraph.pregel._messages import StreamMessagesHandlerV2 + + emitted: list[Any] = [] + handler = StreamMessagesHandlerV2(emitted.append, subgraphs=False) + run_id = uuid4() + handler.metadata[run_id] = ((), {"langgraph_node": "x"}) + + handler.on_stream_event( + {"event": "message-start", "message_id": "stream-msg-1"}, + run_id=run_id, + ) + handler.on_llm_end( + LLMResult( + generations=[ + [ + ChatGeneration( + message=AIMessage(content="hello", id="final-msg-1") + ) + ] + ] + ), + run_id=run_id, + ) + + assert len(emitted) == 1 diff --git a/libs/langgraph/uv.lock b/libs/langgraph/uv.lock index a18fc116b..2a63650d7 100644 --- a/libs/langgraph/uv.lock +++ b/libs/langgraph/uv.lock @@ -1348,10 +1348,11 @@ wheels = [ [[package]] name = "langchain-core" -version = "1.3.0a2" +version = "1.3.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jsonpatch" }, + { name = "langchain-protocol" }, { name = "langsmith" }, { name = "packaging" }, { name = "pydantic" }, @@ -1360,9 +1361,21 @@ dependencies = [ { name = "typing-extensions" }, { name = "uuid-utils" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a8/03/7219502e8ca728d65eb44d7a3eb60239230742a70dbfc9241b9bfd61c4ab/langchain_core-1.3.2.tar.gz", hash = "sha256:fd7a50b2f28ba561fd9d7f5d2760bc9e06cf00cdf820a3ccafe88a94ffa8d5b7", size = 911813, upload-time = "2026-04-24T15:49:23.699Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" }, + { url = "https://files.pythonhosted.org/packages/7d/d5/8fa4431007cbb7cfed7590f4d6a5dea3ad724f4174d248f6642ef5ce7d05/langchain_core-1.3.2-py3-none-any.whl", hash = "sha256:d44a66127f9f8db735bdfd0ab9661bccb47a97113cfd3f2d89c74864422b7274", size = 542390, upload-time = "2026-04-24T15:49:21.991Z" }, +] + +[[package]] +name = "langchain-protocol" +version = "0.0.11" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d0/bb/38b5eaefa41c67735eedd9f9a2568b11c9eb376fa129a5edd7cc3dcde071/langchain_protocol-0.0.11.tar.gz", hash = "sha256:c276e2373b5ac691fc7ac9a72019d55182444ce8e89385c3f7e9f0185d0aace7", size = 6622, upload-time = "2026-04-23T22:13:16.771Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0b/fa/6a8ecad8472b182f2caf9d83fd89f40fc1590cb96546d90089b7869b7f5e/langchain_protocol-0.0.11-py3-none-any.whl", hash = "sha256:364da1faf6f5d3001413bede792c1a822c0f23ae55d1ce1266ca7d8e80e79011", size = 6778, upload-time = "2026-04-23T22:13:15.712Z" }, ] [[package]] @@ -1439,7 +1452,7 @@ test = [ [package.metadata] requires-dist = [ - { name = "langchain-core", specifier = "==1.3.0a2" }, + { name = "langchain-core", specifier = "==1.3.2" }, { name = "langgraph-checkpoint", editable = "../checkpoint" }, { name = "langgraph-prebuilt", editable = "../prebuilt" }, { name = "langgraph-sdk", editable = "../sdk-py" }, diff --git a/libs/prebuilt/uv.lock b/libs/prebuilt/uv.lock index f49a86e4a..9b828181d 100644 --- a/libs/prebuilt/uv.lock +++ b/libs/prebuilt/uv.lock @@ -249,10 +249,11 @@ wheels = [ [[package]] name = "langchain-core" -version = "1.3.0a2" +version = "1.3.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jsonpatch" }, + { name = "langchain-protocol" }, { name = "langsmith" }, { name = "packaging" }, { name = "pydantic" }, @@ -261,9 +262,21 @@ dependencies = [ { name = "typing-extensions" }, { name = "uuid-utils" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a8/03/7219502e8ca728d65eb44d7a3eb60239230742a70dbfc9241b9bfd61c4ab/langchain_core-1.3.2.tar.gz", hash = "sha256:fd7a50b2f28ba561fd9d7f5d2760bc9e06cf00cdf820a3ccafe88a94ffa8d5b7", size = 911813, upload-time = "2026-04-24T15:49:23.699Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" }, + { url = "https://files.pythonhosted.org/packages/7d/d5/8fa4431007cbb7cfed7590f4d6a5dea3ad724f4174d248f6642ef5ce7d05/langchain_core-1.3.2-py3-none-any.whl", hash = "sha256:d44a66127f9f8db735bdfd0ab9661bccb47a97113cfd3f2d89c74864422b7274", size = 542390, upload-time = "2026-04-24T15:49:21.991Z" }, +] + +[[package]] +name = "langchain-protocol" +version = "0.0.11" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d0/bb/38b5eaefa41c67735eedd9f9a2568b11c9eb376fa129a5edd7cc3dcde071/langchain_protocol-0.0.11.tar.gz", hash = "sha256:c276e2373b5ac691fc7ac9a72019d55182444ce8e89385c3f7e9f0185d0aace7", size = 6622, upload-time = "2026-04-23T22:13:16.771Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0b/fa/6a8ecad8472b182f2caf9d83fd89f40fc1590cb96546d90089b7869b7f5e/langchain_protocol-0.0.11-py3-none-any.whl", hash = "sha256:364da1faf6f5d3001413bede792c1a822c0f23ae55d1ce1266ca7d8e80e79011", size = 6778, upload-time = "2026-04-23T22:13:15.712Z" }, ] [[package]] @@ -281,7 +294,7 @@ dependencies = [ [package.metadata] requires-dist = [ - { name = "langchain-core", specifier = "==1.3.0a2" }, + { name = "langchain-core", specifier = "==1.3.2" }, { name = "langgraph-checkpoint", editable = "../checkpoint" }, { name = "langgraph-prebuilt", editable = "." }, { name = "langgraph-sdk", editable = "../sdk-py" }, diff --git a/libs/sdk-py/uv.lock b/libs/sdk-py/uv.lock index 767980369..7046da656 100644 --- a/libs/sdk-py/uv.lock +++ b/libs/sdk-py/uv.lock @@ -262,10 +262,11 @@ wheels = [ [[package]] name = "langchain-core" -version = "1.3.0a2" +version = "1.3.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jsonpatch" }, + { name = "langchain-protocol" }, { name = "langsmith" }, { name = "packaging" }, { name = "pydantic" }, @@ -274,9 +275,21 @@ dependencies = [ { name = "typing-extensions" }, { name = "uuid-utils" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a8/03/7219502e8ca728d65eb44d7a3eb60239230742a70dbfc9241b9bfd61c4ab/langchain_core-1.3.2.tar.gz", hash = "sha256:fd7a50b2f28ba561fd9d7f5d2760bc9e06cf00cdf820a3ccafe88a94ffa8d5b7", size = 911813, upload-time = "2026-04-24T15:49:23.699Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" }, + { url = "https://files.pythonhosted.org/packages/7d/d5/8fa4431007cbb7cfed7590f4d6a5dea3ad724f4174d248f6642ef5ce7d05/langchain_core-1.3.2-py3-none-any.whl", hash = "sha256:d44a66127f9f8db735bdfd0ab9661bccb47a97113cfd3f2d89c74864422b7274", size = 542390, upload-time = "2026-04-24T15:49:21.991Z" }, +] + +[[package]] +name = "langchain-protocol" +version = "0.0.11" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d0/bb/38b5eaefa41c67735eedd9f9a2568b11c9eb376fa129a5edd7cc3dcde071/langchain_protocol-0.0.11.tar.gz", hash = "sha256:c276e2373b5ac691fc7ac9a72019d55182444ce8e89385c3f7e9f0185d0aace7", size = 6622, upload-time = "2026-04-23T22:13:16.771Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0b/fa/6a8ecad8472b182f2caf9d83fd89f40fc1590cb96546d90089b7869b7f5e/langchain_protocol-0.0.11-py3-none-any.whl", hash = "sha256:364da1faf6f5d3001413bede792c1a822c0f23ae55d1ce1266ca7d8e80e79011", size = 6778, upload-time = "2026-04-23T22:13:15.712Z" }, ] [[package]] @@ -294,7 +307,7 @@ dependencies = [ [package.metadata] requires-dist = [ - { name = "langchain-core", specifier = "==1.3.0a2" }, + { name = "langchain-core", specifier = "==1.3.2" }, { name = "langgraph-checkpoint", editable = "../checkpoint" }, { name = "langgraph-prebuilt", editable = "../prebuilt" }, { name = "langgraph-sdk", editable = "." },