fix(langgraph): do NOT re-execute nodes on invoke(None, ...). fix tests

This commit is contained in:
Caspar Broekhuizen
2025-10-16 11:33:43 -07:00
parent 8b42793d30
commit dd02a773ab
4 changed files with 16 additions and 16 deletions
+1 -1
View File
@@ -500,7 +500,7 @@ class PregelLoop:
)
resume_map = self.config.get(CONF, {}).get(CONFIG_KEY_RESUME_MAP, {})
if resume_map:
if resume_map or self.input is None:
skipped_interrupt_ids = self._pending_interrupts() - set(resume_map)
self.skipped_task_ids = {
task_id
+13 -13
View File
@@ -479,12 +479,12 @@ def test_node_with_multiple_interrupts_requires_full_resume(
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with None resume. this should execute the node
# invoke with None resume. this should NOT execute the node
partial = graph.invoke(None, config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 3
assert node_counter == 2
# invoke with nonspecific resume. this should execute the node
partial = graph.invoke(Command(resume="human_second"), config=config)
@@ -492,14 +492,14 @@ def test_node_with_multiple_interrupts_requires_full_resume(
assert len(remaining_interrupts) == 1
print("REMAINING INTERRUPTS: ", remaining_interrupts)
assert remaining_interrupts[0].value == "third"
assert node_counter == 4
assert node_counter == 3
# finally, invoke with an interrupt map that matches double_interrupt_node.
# this should execute the node and all interrupts should be resolved
final_result = graph.invoke(Command(resume="human_third"), config=config)
assert "input" in final_result
assert final_result["input"] == "human_first-human_second-human_third"
assert node_counter == 5
assert node_counter == 4
@NEEDS_CONTEXTVARS
@@ -561,12 +561,12 @@ async def test_node_with_multiple_interrupts_requires_full_resume_async(
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with None resume. this should execute the node
# invoke with None resume. this should NOT execute the node
partial = await graph.ainvoke(None, config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 3
assert node_counter == 2
# invoke with nonspecific resume. this should execute the node
partial = await graph.ainvoke(Command(resume="human_second"), config=config)
@@ -574,14 +574,14 @@ async def test_node_with_multiple_interrupts_requires_full_resume_async(
assert len(remaining_interrupts) == 1
print("REMAINING INTERRUPTS: ", remaining_interrupts)
assert remaining_interrupts[0].value == "third"
assert node_counter == 4
assert node_counter == 3
# finally, invoke with an interrupt map that matches double_interrupt_node.
# this should execute the node and all interrupts should be resolved
final_result = await graph.ainvoke(Command(resume="human_third"), config=config)
assert "input" in final_result
assert final_result["input"] == "human_first-human_second-human_third"
assert node_counter == 5
assert node_counter == 4
def test_invoke_interrupted_graph_with_none(
@@ -616,14 +616,14 @@ def test_invoke_interrupted_graph_with_none(
assert len(interrupts) == 1
assert node_counter == 1
# invoke with None. this should execute the node and the history should
# invoke with None. this should NOT execute the node and the history should
# look the same as the first run
partial = graph.invoke(None, config=config)
second_history = list(graph.get_state_history(config))
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "first"
assert node_counter == 2
assert node_counter == 1
# history should look the same for tasks and interrupts
print("first_history[0].interrupts: ", first_history[0].interrupts)
@@ -640,7 +640,7 @@ def test_invoke_interrupted_graph_with_none(
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 3
assert node_counter == 2
# invoke with None again. the history should look the same as
# the third run
@@ -649,7 +649,7 @@ def test_invoke_interrupted_graph_with_none(
fourth_history = list(graph.get_state_history(config))
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert node_counter == 4
assert node_counter == 2
print("\nthird_history[0].interrupts: ", third_history[0].interrupts)
print("fourth_history[0].interrupts: ", fourth_history[0].interrupts)
@@ -662,4 +662,4 @@ def test_invoke_interrupted_graph_with_none(
partial = graph.invoke(Command(resume="bix"), config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 0
assert node_counter == 5
assert node_counter == 3
+1 -1
View File
@@ -4726,7 +4726,7 @@ def test_send_dedupe_on_resume(
assert len(history) == (4 if durability != "exit" else 1)
# resume execution
assert graph.invoke(None, thread1, durability=durability) == [
assert graph.invoke(Command(resume=""), thread1, durability=durability) == [
"0",
"1",
"3.1",
+1 -1
View File
@@ -2545,7 +2545,7 @@ async def test_send_dedupe_on_resume(
assert builder.nodes["2"].runnable.func.ticks == 3
assert builder.nodes["flaky"].runnable.func.ticks == 1
# resume execution
assert await graph.ainvoke(None, thread1, durability=durability) == [
assert await graph.ainvoke(Command(resume=""), thread1, durability=durability) == [
"0",
"1",
"3.1",