From 2eb3b316c968e861d3a6aadaf80295cff818aaa6 Mon Sep 17 00:00:00 2001 From: Eugene Yurtsev Date: Tue, 5 Aug 2025 14:06:48 -0400 Subject: [PATCH] x --- libs/langgraph/tests/test_interruption.py | 87 +++++++++++++++++++++++ 1 file changed, 87 insertions(+) diff --git a/libs/langgraph/tests/test_interruption.py b/libs/langgraph/tests/test_interruption.py index 2abd0749a..293d50f28 100644 --- a/libs/langgraph/tests/test_interruption.py +++ b/libs/langgraph/tests/test_interruption.py @@ -165,3 +165,90 @@ def test_interrupt_with_send_payloads(sync_checkpointer: BaseCheckpointSaver) -> # Map node runs 3 times initially (item1 completes, 2 dangerous_items interrupt), # then 2 times on resume assert node_counter["map_node"] == 5 + + +def test_interrupt_with_send_payloads_sequential_resume( + sync_checkpointer: BaseCheckpointSaver, +) -> None: + """Test interruption in map node with Send payloads and sequential resume.""" + + # Global counter to track node executions + node_counter = {"entry": 0, "map_node": 0} + + class State(TypedDict): + items: list[str] + processed: Annotated[list[str], operator.add] + + def entry_node(state: State): + node_counter["entry"] += 1 + return {} # No state updates in entry node + + def send_to_map(state: State): + return [Send("map_node", {"item": item}) for item in state["items"]] + + def map_node(state: State): + node_counter["map_node"] += 1 + if "dangerous" in state["item"]: + value = interrupt({"processing": state["item"]}) + return {"processed": [f"processed_{value}"]} + else: + return {"processed": [f"processed_{state['item']}_auto"]} + + builder = StateGraph(State) + builder.add_node("entry", entry_node) + builder.add_node("map_node", map_node) + builder.add_edge(START, "entry") + builder.add_conditional_edges("entry", send_to_map, ["map_node"]) + builder.add_edge("map_node", END) + + graph = builder.compile(checkpointer=sync_checkpointer) + + config = {"configurable": {"thread_id": "test_interrupt_send_sequential"}} + + # Run until interrupts + result = graph.invoke( + {"items": ["item1", "dangerous_item1", "dangerous_item2"]}, config=config + ) + + # Verify we have interrupts + interrupts = result.get("__interrupt__", []) + assert len(interrupts) == 2 + assert "dangerous_item" in interrupts[0].value["processing"] + + # Resume first interrupt only + first_interrupt = interrupts[0] + first_resume_map = { + first_interrupt.interrupt_id: f"human_input_{first_interrupt.value['processing']}" + } + + partial_result = graph.invoke(Command(resume=first_resume_map), config=config) + + # Verify we still have one pending interrupt + remaining_interrupts = partial_result.get("__interrupt__", []) + assert len(remaining_interrupts) == 1 + + # Resume second interrupt + second_interrupt = remaining_interrupts[0] + second_resume_map = { + second_interrupt.interrupt_id: f"human_input_{second_interrupt.value['processing']}" + } + + final_result = graph.invoke(Command(resume=second_resume_map), config=config) + + # Verify final result contains processed items + assert "processed" in final_result + processed_items = final_result["processed"] + assert len(processed_items) == 3 + assert "processed_item1_auto" in processed_items # item1 processed automatically + assert any( + "processed_human_input_dangerous_item1" in item for item in processed_items + ) # dangerous_item1 processed after interrupt + assert any( + "processed_human_input_dangerous_item2" in item for item in processed_items + ) # dangerous_item2 processed after interrupt + + # Verify node execution counts + assert node_counter["entry"] == 1 # Entry node runs once + # Map node runs 3 times initially (item1 completes, 2 dangerous_items interrupt), + # then 1 time on first resume, then 1 time on second resume + assert node_counter["map_node"] == 5