Support stream_mode=updates in invoke()

This commit is contained in:
Nuno Campos
2024-04-10 17:12:09 -07:00
parent 8725839492
commit 8788a6adfb
2 changed files with 115 additions and 8 deletions
+28 -8
View File
@@ -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(
+87
View File
@@ -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")]})
] == [