From 8788a6adfbf6bb3245dd9b04f3f34cd5b6fe0210 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 10 Apr 2024 17:12:09 -0700 Subject: [PATCH] Support stream_mode=updates in invoke() --- langgraph/pregel/__init__.py | 36 +++++++++++---- tests/test_pregel.py | 87 ++++++++++++++++++++++++++++++++++++ 2 files changed, 115 insertions(+), 8 deletions(-) diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index 999017ed3..c766b6873 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -964,6 +964,7 @@ class Pregel( input: Union[dict[str, Any], Any], config: Optional[RunnableConfig] = None, *, + stream_mode: StreamMode = "values", output_keys: Optional[Union[str, Sequence[str]]] = None, input_keys: Optional[Union[str, Sequence[str]]] = None, interrupt_before_nodes: Optional[Sequence[str]] = None, @@ -973,11 +974,14 @@ class Pregel( ) -> Union[dict[str, Any], Any]: output_keys = output_keys if output_keys is not None else self.output_channels output_is_dict = not isinstance(output_keys, str) - latest: Union[dict[str, Any], Any] = {} if output_is_dict else None + if stream_mode == "values": + latest: Union[dict[str, Any], Any] = {} if output_is_dict else None + else: + chunks = [] for chunk in self.stream( input, config, - stream_mode="values", + stream_mode=stream_mode, output_keys=output_keys, input_keys=input_keys, interrupt_before_nodes=interrupt_before_nodes, @@ -985,14 +989,21 @@ class Pregel( debug=debug, **kwargs, ): - latest = {**latest, **chunk} if output_is_dict else chunk - return latest + if stream_mode == "values": + latest = {**latest, **chunk} if output_is_dict else chunk + else: + chunks.append(chunk) + if stream_mode == "values": + return latest + else: + return chunks async def ainvoke( self, input: Union[dict[str, Any], Any], config: Optional[RunnableConfig] = None, *, + stream_mode: StreamMode = "values", output_keys: Optional[Union[str, Sequence[str]]] = None, input_keys: Optional[Union[str, Sequence[str]]] = None, interrupt_before_nodes: Optional[Sequence[str]] = None, @@ -1002,11 +1013,14 @@ class Pregel( ) -> Union[dict[str, Any], Any]: output_keys = output_keys if output_keys is not None else self.output_channels output_is_dict = not isinstance(output_keys, str) - latest: Union[dict[str, Any], Any] = {} if output_is_dict else None + if stream_mode == "values": + latest: Union[dict[str, Any], Any] = {} if output_is_dict else None + else: + chunks = [] async for chunk in self.astream( input, config, - stream_mode="values", + stream_mode=stream_mode, output_keys=output_keys, input_keys=input_keys, interrupt_before_nodes=interrupt_before_nodes, @@ -1014,8 +1028,14 @@ class Pregel( debug=debug, **kwargs, ): - latest = {**latest, **chunk} if output_is_dict else chunk - return latest + if stream_mode == "values": + latest = {**latest, **chunk} if output_is_dict else chunk + else: + chunks.append(chunk) + if stream_mode == "values": + return latest + else: + return chunks def _panic_or_proceed( diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 8e091aa09..6be844479 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -1939,6 +1939,93 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None: ] } + assert app.invoke( + {"messages": [HumanMessage(content="what is weather in sf")]}, + stream_mode="updates", + ) == [ + { + "agent": { + "messages": [ + AIMessage( + content="", + additional_kwargs={ + "tool_calls": [ + { + "id": "tool_call123", + "type": "function", + "function": { + "name": "search_api", + "arguments": '"query"', + }, + } + ] + }, + id=AnyStr(), + ) + ] + } + }, + { + "action": { + "messages": [ + ToolMessage(content="result for query", tool_call_id="tool_call123") + ] + } + }, + { + "agent": { + "messages": [ + AIMessage( + content="", + additional_kwargs={ + "tool_calls": [ + { + "id": "tool_call234", + "type": "function", + "function": { + "name": "search_api", + "arguments": '"another"', + }, + }, + { + "id": "tool_call567", + "type": "function", + "function": { + "name": "search_api", + "arguments": '"a third one"', + }, + }, + ] + }, + id=AnyStr(), + ) + ] + } + }, + { + "action": { + "messages": [ + ToolMessage( + content="result for another", tool_call_id="tool_call234" + ), + ToolMessage( + content="result for a third one", tool_call_id="tool_call567" + ), + ] + } + }, + { + "agent": { + "messages": [ + AIMessage( + content="answer", + id=AnyStr(), + ) + ] + } + }, + ] + assert [ *app.stream({"messages": [HumanMessage(content="what is weather in sf")]}) ] == [