This commit is contained in:
Tat Dat Duong
2025-04-30 20:04:48 +02:00
parent b3371a1d63
commit be1af772f5
2 changed files with 17 additions and 24 deletions
+2 -2
View File
@@ -1,4 +1,4 @@
from typing import Any, Literal, Optional, Union
from typing import Any, Literal, Optional, Union, cast
from uuid import uuid4
from langchain_core.messages import AnyMessage
@@ -194,7 +194,7 @@ def ui_message_reducer(
else:
ids_to_remove.discard(msg_id)
if msg.get("metadata", {}).get("merge", False):
if cast(UIMessage, msg).get("metadata", {}).get("merge", False):
prev_msg = merged[existing_idx]
msg = msg.copy()
msg["props"] = {**prev_msg["props"], **msg["props"]}
+15 -22
View File
@@ -6351,9 +6351,7 @@ def test_double_interrupt_subgraph(
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_multi_resume(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
def test_multi_resume(request: pytest.FixtureRequest, checkpointer_name: str) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class ChildState(TypedDict):
@@ -6362,11 +6360,11 @@ def test_multi_resume(
human_inputs: list[str]
def get_human_input(state: ChildState):
human_input = interrupt(state['prompt'])
human_input = interrupt(state["prompt"])
return {
'human_input': human_input,
'human_inputs': [human_input],
"human_input": human_input,
"human_inputs": [human_input],
}
child_graph = (
@@ -6385,13 +6383,13 @@ def test_multi_resume(
return [
Send(
"child_graph",
{'prompt': prompt},
{"prompt": prompt},
)
for prompt in state['prompts']
for prompt in state["prompts"]
]
def cleanup(state: ParentState):
assert len(state['human_inputs']) == len(state["prompts"])
assert len(state["human_inputs"]) == len(state["prompts"])
parent_graph = (
StateGraph(ParentState)
@@ -6404,21 +6402,19 @@ def test_multi_resume(
)
thread_config: RunnableConfig = {
'configurable': {
'thread_id': uuid.uuid4(),
"configurable": {
"thread_id": uuid.uuid4(),
},
}
prompts = ['a', 'b', 'c', 'd', 'e']
prompts = ["a", "b", "c", "d", "e"]
events = parent_graph.invoke(
{'prompts': prompts},
thread_config,
stream_mode='values'
{"prompts": prompts}, thread_config, stream_mode="values"
)
assert len(events['__interrupt__']) == len(prompts)
interrupt_values = {i.value for i in events['__interrupt__']}
assert len(events["__interrupt__"]) == len(prompts)
interrupt_values = {i.value for i in events["__interrupt__"]}
assert interrupt_values == set(prompts)
resume_map: dict[str, str] = {
@@ -6428,11 +6424,8 @@ def test_multi_resume(
result = parent_graph.invoke(Command(resume=resume_map), thread_config)
assert result == {
'prompts': prompts,
'human_inputs': [
f"human input for prompt {prompt}"
for prompt in prompts
],
"prompts": prompts,
"human_inputs": [f"human input for prompt {prompt}" for prompt in prompts],
}