mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-30 19:59:40 +02:00
Support stream_mode=updates in invoke()
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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")]})
|
||||
] == [
|
||||
|
||||
Reference in New Issue
Block a user