From 5071a6cd975325c84e8e4084395a3120b0b10c43 Mon Sep 17 00:00:00 2001 From: vbarda Date: Fri, 11 Apr 2025 13:55:49 -0400 Subject: [PATCH 1/4] langgraph: support streaming messages from Command.update --- libs/langgraph/langgraph/pregel/messages.py | 57 +++++++++++---------- libs/langgraph/tests/test_pregel.py | 46 +++++++++++++++++ libs/langgraph/tests/test_pregel_async.py | 47 +++++++++++++++++ 3 files changed, 123 insertions(+), 27 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/messages.py b/libs/langgraph/langgraph/pregel/messages.py index 5766c2b63..6ac01fb93 100644 --- a/libs/langgraph/langgraph/pregel/messages.py +++ b/libs/langgraph/langgraph/pregel/messages.py @@ -18,7 +18,7 @@ from langchain_core.messages import BaseMessage from langchain_core.outputs import ChatGenerationChunk, LLMResult from langgraph.constants import NS_SEP, TAG_HIDDEN, TAG_NOSTREAM -from langgraph.types import StreamChunk +from langgraph.types import Command, StreamChunk try: from langchain_core.tracers._streaming import _StreamingCallbackHandler @@ -153,32 +153,35 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): **kwargs: Any, ) -> Any: if meta := self.metadata.pop(run_id, None): - if isinstance(response, BaseMessage): - self._emit(meta, response, dedupe=True) - elif isinstance(response, Sequence): - for value in response: - if isinstance(value, BaseMessage): - self._emit(meta, value, dedupe=True) - elif isinstance(response, dict): - for value in response.values(): - if isinstance(value, BaseMessage): - self._emit(meta, value, dedupe=True) - elif isinstance(value, Sequence): - for item in value: - if isinstance(item, BaseMessage): - self._emit(meta, item, dedupe=True) - elif hasattr(response, "__dir__") and callable(response.__dir__): - for key in dir(response): - try: - value = getattr(response, key) - if isinstance(value, BaseMessage): - self._emit(meta, value, dedupe=True) - elif isinstance(value, Sequence): - for item in value: - if isinstance(item, BaseMessage): - self._emit(meta, item, dedupe=True) - except AttributeError: - pass + self._process_response(response, meta, set()) + + def _process_response(self, response: Any, meta: Meta, visited: set[int]) -> None: + """Recursively process a response to find and emit BaseMessage instances.""" + if response is None or id(response) in visited: + return + + visited.add(id(response)) + + if isinstance(response, BaseMessage): + self._emit(meta, response, dedupe=True) + elif isinstance(response, dict): + for value in response.values(): + self._process_response(value, meta, visited) + elif isinstance(response, Sequence) and not isinstance(response, str): + for item in response: + self._process_response(item, meta, visited) + elif hasattr(response, "__dir__") and callable(response.__dir__): + for key in dir(response): + # Skip magic methods and properties to reduce recursion depth + if key.startswith("__") and key.endswith("__"): + continue + try: + value = getattr(response, key) + self._process_response(value, meta, visited) + except AttributeError: + pass + elif isinstance(response, Command): + self._process_response(response.update, meta, visited) def on_chain_error( self, diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index be3f94ec7..83c8b4a7e 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7168,6 +7168,52 @@ def test_tags_stream_mode_messages() -> None: ] +def test_stream_mode_messages_command() -> None: + from langchain_core.messages import HumanMessage + + def my_node(state): + return {"messages": HumanMessage(content="foo")} + + def my_other_node(state): + return Command(update={"messages": HumanMessage(content="bar")}) + + graph = ( + StateGraph(MessagesState) + .add_sequence([my_node, my_other_node]) + .add_edge(START, "my_node") + .compile() + ) + assert list( + graph.stream( + { + "messages": [], + }, + stream_mode="messages", + ) + ) == [ + ( + _AnyIdHumanMessage(content="foo"), + { + "langgraph_step": 1, + "langgraph_node": "my_node", + "langgraph_triggers": ("branch:to:my_node",), + "langgraph_path": ("__pregel_pull", "my_node"), + "langgraph_checkpoint_ns": AnyStr("my_node:"), + }, + ), + ( + _AnyIdHumanMessage(content="bar"), + { + "langgraph_step": 2, + "langgraph_node": "my_other_node", + "langgraph_triggers": ("branch:to:my_other_node",), + "langgraph_path": ("__pregel_pull", "my_other_node"), + "langgraph_checkpoint_ns": AnyStr("my_other_node:"), + }, + ), + ] + + def test_node_destinations() -> None: class State(TypedDict): foo: Annotated[str, operator.add] diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 0e74f0b74..987732b90 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -7894,6 +7894,53 @@ async def test_tags_stream_mode_messages() -> None: ] +async def test_stream_mode_messages_command() -> None: + from langchain_core.messages import HumanMessage + + async def my_node(state): + return {"messages": HumanMessage(content="foo")} + + async def my_other_node(state): + return Command(update={"messages": HumanMessage(content="bar")}) + + graph = ( + StateGraph(MessagesState) + .add_sequence([my_node, my_other_node]) + .add_edge(START, "my_node") + .compile() + ) + assert [ + c + async for c in graph.astream( + { + "messages": [], + }, + stream_mode="messages", + ) + ] == [ + ( + _AnyIdHumanMessage(content="foo"), + { + "langgraph_step": 1, + "langgraph_node": "my_node", + "langgraph_triggers": ("branch:to:my_node",), + "langgraph_path": ("__pregel_pull", "my_node"), + "langgraph_checkpoint_ns": AnyStr("my_node:"), + }, + ), + ( + _AnyIdHumanMessage(content="bar"), + { + "langgraph_step": 2, + "langgraph_node": "my_other_node", + "langgraph_triggers": ("branch:to:my_other_node",), + "langgraph_path": ("__pregel_pull", "my_other_node"), + "langgraph_checkpoint_ns": AnyStr("my_other_node:"), + }, + ), + ] + + async def test_stream_messages_dedupe_inputs() -> None: from langchain_core.messages import AIMessage From b526fe0a4b34618520980d7553dae3f1be9e8980 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 14 Apr 2025 13:21:00 -0400 Subject: [PATCH 2/4] set max recursion depth --- libs/langgraph/langgraph/pregel/messages.py | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/messages.py b/libs/langgraph/langgraph/pregel/messages.py index 6ac01fb93..300a882a1 100644 --- a/libs/langgraph/langgraph/pregel/messages.py +++ b/libs/langgraph/langgraph/pregel/messages.py @@ -153,23 +153,25 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): **kwargs: Any, ) -> Any: if meta := self.metadata.pop(run_id, None): - self._process_response(response, meta, set()) + self._process_response(response, meta, 0) - def _process_response(self, response: Any, meta: Meta, visited: set[int]) -> None: + def _process_response(self, response: Any, meta: Meta, depth: int = 0) -> None: """Recursively process a response to find and emit BaseMessage instances.""" - if response is None or id(response) in visited: + if response is None: return - visited.add(id(response)) + # Cap recursion depth at 5 + if depth >= 5: + return if isinstance(response, BaseMessage): self._emit(meta, response, dedupe=True) elif isinstance(response, dict): for value in response.values(): - self._process_response(value, meta, visited) + self._process_response(value, meta, depth + 1) elif isinstance(response, Sequence) and not isinstance(response, str): for item in response: - self._process_response(item, meta, visited) + self._process_response(item, meta, depth + 1) elif hasattr(response, "__dir__") and callable(response.__dir__): for key in dir(response): # Skip magic methods and properties to reduce recursion depth @@ -177,11 +179,11 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): continue try: value = getattr(response, key) - self._process_response(value, meta, visited) + self._process_response(value, meta, depth + 1) except AttributeError: pass elif isinstance(response, Command): - self._process_response(response.update, meta, visited) + self._process_response(response.update, meta, depth + 1) def on_chain_error( self, From 07ca03ff152b40b606e6c0605604b658896b97c0 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 14 Apr 2025 14:12:40 -0400 Subject: [PATCH 3/4] lower depth --- libs/langgraph/langgraph/pregel/messages.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/messages.py b/libs/langgraph/langgraph/pregel/messages.py index 300a882a1..dfb710d2a 100644 --- a/libs/langgraph/langgraph/pregel/messages.py +++ b/libs/langgraph/langgraph/pregel/messages.py @@ -160,8 +160,8 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): if response is None: return - # Cap recursion depth at 5 - if depth >= 5: + # Cap recursion depth at 3 + if depth >= 3: return if isinstance(response, BaseMessage): From 04dd69b1cda74a49032e472e2e5f23f963727223 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 14 Apr 2025 14:53:08 -0400 Subject: [PATCH 4/4] simplify --- libs/langgraph/langgraph/pregel/messages.py | 64 +++++++++++---------- 1 file changed, 35 insertions(+), 29 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/messages.py b/libs/langgraph/langgraph/pregel/messages.py index dfb710d2a..16d0904db 100644 --- a/libs/langgraph/langgraph/pregel/messages.py +++ b/libs/langgraph/langgraph/pregel/messages.py @@ -153,37 +153,43 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): **kwargs: Any, ) -> Any: if meta := self.metadata.pop(run_id, None): - self._process_response(response, meta, 0) + if isinstance(response, Command): + response = response.update - def _process_response(self, response: Any, meta: Meta, depth: int = 0) -> None: - """Recursively process a response to find and emit BaseMessage instances.""" - if response is None: - return + if isinstance(response, Sequence) and any( + isinstance(value, Command) for value in response + ): + response = [ + value.update if isinstance(value, Command) else value + for value in response + ] - # Cap recursion depth at 3 - if depth >= 3: - return - - if isinstance(response, BaseMessage): - self._emit(meta, response, dedupe=True) - elif isinstance(response, dict): - for value in response.values(): - self._process_response(value, meta, depth + 1) - elif isinstance(response, Sequence) and not isinstance(response, str): - for item in response: - self._process_response(item, meta, depth + 1) - elif hasattr(response, "__dir__") and callable(response.__dir__): - for key in dir(response): - # Skip magic methods and properties to reduce recursion depth - if key.startswith("__") and key.endswith("__"): - continue - try: - value = getattr(response, key) - self._process_response(value, meta, depth + 1) - except AttributeError: - pass - elif isinstance(response, Command): - self._process_response(response.update, meta, depth + 1) + if isinstance(response, BaseMessage): + self._emit(meta, response, dedupe=True) + elif isinstance(response, Sequence): + for value in response: + if isinstance(value, BaseMessage): + self._emit(meta, value, dedupe=True) + elif isinstance(response, dict): + for value in response.values(): + if isinstance(value, BaseMessage): + self._emit(meta, value, dedupe=True) + elif isinstance(value, Sequence): + for item in value: + if isinstance(item, BaseMessage): + self._emit(meta, item, dedupe=True) + elif hasattr(response, "__dir__") and callable(response.__dir__): + for key in dir(response): + try: + value = getattr(response, key) + if isinstance(value, BaseMessage): + self._emit(meta, value, dedupe=True) + elif isinstance(value, Sequence): + for item in value: + if isinstance(item, BaseMessage): + self._emit(meta, item, dedupe=True) + except AttributeError: + pass def on_chain_error( self,