diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 29a2d3ccb..2ba8af3a6 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -116,7 +116,7 @@ def DuplexStream(*streams: StreamProtocol) -> StreamProtocol: def __call__(value: StreamChunk) -> None: for stream in streams: if value[1] in stream.modes: - stream(value) # type: ignore + stream(value) return StreamProtocol(__call__, {mode for s in streams for mode in s.modes}) @@ -587,7 +587,7 @@ class PregelLoop(LoopProtocol): if mode not in self.stream.modes: return for v in values(*args, **kwargs): - self.stream((self.checkpoint_ns, mode, v)) # type: ignore + self.stream((self.checkpoint_ns, mode, v)) def _output_writes( self, task_id: str, writes: Sequence[tuple[str, Any]], *, cached: bool = False diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index e47f5f50b..ccf4306f2 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -31,11 +31,16 @@ from langgraph_sdk.schema import StreamMode as StreamModeSDK from typing_extensions import Self from langgraph.checkpoint.base import CheckpointMetadata -from langgraph.constants import INTERRUPT +from langgraph.constants import ( + CONF, + CONFIG_KEY_CHECKPOINT_NS, + CONFIG_KEY_STREAM, + INTERRUPT, +) from langgraph.errors import GraphInterrupt from langgraph.pregel.protocol import PregelProtocol from langgraph.pregel.types import All, PregelTask, StateSnapshot, StreamMode -from langgraph.types import Interrupt +from langgraph.types import Interrupt, StreamProtocol from langgraph.utils.config import merge_configs @@ -507,7 +512,7 @@ class RemoteGraph(PregelProtocol): 'updates' mode is added to the list of stream modes so that interrupts can be detected in the remote graph. """ - updated_stream_modes: list[StreamMode] = [] + updated_stream_modes: list[StreamModeSDK] = [] req_updates = False req_single = True # coerce to list, or add default stream mode @@ -519,6 +524,10 @@ class RemoteGraph(PregelProtocol): updated_stream_modes.extend(stream_mode) else: updated_stream_modes.append(default) + # map "messages" to "messages-tuple" + if "messages" in updated_stream_modes: + updated_stream_modes.remove("messages") + updated_stream_modes.append("messages-tuple") # add 'updates' mode if not present if "updates" in updated_stream_modes: req_updates = True @@ -557,18 +566,33 @@ class RemoteGraph(PregelProtocol): merged_config = merge_configs(self.config, config) sanitized_config = self._sanitize_config(merged_config) stream_modes, req_updates, req_single = self._get_stream_modes(stream_mode) + stream: Optional[StreamProtocol] = ( + (config or {}).get(CONF, {}).get(CONFIG_KEY_STREAM) + ) + stream_modes_ext: list[StreamModeSDK] = ( + [*stream_modes, *stream.modes] if stream else stream_modes + ) for chunk in sync_client.runs.stream( thread_id=sanitized_config["configurable"].get("thread_id"), assistant_id=self.name, input=input, config=sanitized_config, - stream_mode=stream_modes, + stream_mode=stream_modes_ext, interrupt_before=interrupt_before, interrupt_after=interrupt_after, - stream_subgraphs=subgraphs, + stream_subgraphs=subgraphs or stream is not None, if_not_exists="create", ): + if "|" in chunk.event: + mode, ns_ = chunk.event.split("|", 1) + ns = tuple(ns_.split("|")) + else: + mode, ns = chunk.event, () + if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS): + ns = caller_ns + ns + if stream is not None and chunk.event in stream.modes: + stream((ns, mode, chunk.data)) if chunk.event.startswith("updates"): if isinstance(chunk.data, dict) and INTERRUPT in chunk.data: raise GraphInterrupt(chunk.data[INTERRUPT]) @@ -576,6 +600,8 @@ class RemoteGraph(PregelProtocol): continue elif chunk.event.startswith("error"): raise RemoteException(chunk.data) + if chunk.event.split("|", 1)[0] not in stream_modes: + continue if subgraphs: if "|" in chunk.event: mode, ns_ = chunk.event.split("|", 1) @@ -622,18 +648,33 @@ class RemoteGraph(PregelProtocol): merged_config = merge_configs(self.config, config) sanitized_config = self._sanitize_config(merged_config) stream_modes, req_updates, req_single = self._get_stream_modes(stream_mode) + stream: Optional[StreamProtocol] = ( + (config or {}).get(CONF, {}).get(CONFIG_KEY_STREAM) + ) + stream_modes_ext: list[StreamModeSDK] = ( + [*stream_modes, *stream.modes] if stream else stream_modes + ) async for chunk in client.runs.stream( thread_id=sanitized_config["configurable"].get("thread_id"), assistant_id=self.name, input=input, config=sanitized_config, - stream_mode=stream_modes, + stream_mode=stream_modes_ext, interrupt_before=interrupt_before, interrupt_after=interrupt_after, stream_subgraphs=subgraphs, if_not_exists="create", ): + if "|" in chunk.event: + mode, ns_ = chunk.event.split("|", 1) + ns = tuple(ns_.split("|")) + else: + mode, ns = chunk.event, () + if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS): + ns = caller_ns + ns + if stream is not None and chunk.event in stream.modes: + stream((ns, mode, chunk.data)) if chunk.event.startswith("updates"): if isinstance(chunk.data, dict) and INTERRUPT in chunk.data: raise GraphInterrupt(chunk.data[INTERRUPT]) @@ -641,12 +682,9 @@ class RemoteGraph(PregelProtocol): continue elif chunk.event.startswith("error"): raise RemoteException(chunk.data) + if chunk.event.split("|", 1)[0] not in stream_modes: + continue if subgraphs: - if "|" in chunk.event: - mode, ns_ = chunk.event.split("|", 1) - ns = tuple(ns_.split("|")) - else: - mode, ns = chunk.event, () if req_single: yield ns, chunk.data else: diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index c29cc1480..279594dec 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -10,9 +10,11 @@ from typing import ( Sequence, Type, Union, + cast, ) from langchain_core.runnables import Runnable, RunnableConfig +from typing_extensions import Self from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata @@ -227,14 +229,14 @@ class StreamProtocol: modes: set[StreamMode] - __call__: Callable[[StreamChunk], None] + __call__: Callable[[Self, StreamChunk], None] def __init__( self, __call__: Callable[[StreamChunk], None], modes: set[StreamMode], ) -> None: - self.__call__ = __call__ + self.__call__ = cast(Callable[[Self, StreamChunk], None], __call__) self.modes = modes diff --git a/libs/sdk-py/langgraph_sdk/schema.py b/libs/sdk-py/langgraph_sdk/schema.py index 6b40dd9e9..43e218669 100644 --- a/libs/sdk-py/langgraph_sdk/schema.py +++ b/libs/sdk-py/langgraph_sdk/schema.py @@ -25,7 +25,9 @@ Represents the status of a thread: - "error": An exception occurred during task processing. """ -StreamMode = Literal["values", "messages", "updates", "events", "debug", "custom"] +StreamMode = Literal[ + "values", "messages", "updates", "events", "debug", "custom", "messages-tuple" +] """ Defines the mode of streaming: - "values": Stream only the values.