fix: sequential interrupt handling w/ functional API (#6863)

In the original proposed fix
https://github.com/langchain-ai/langgraph/pull/6863, we were guarding
this specific case by only checking for null resume in the scratchpad.

But I noticed the tests passed on main for the sync case, indicating an
inconsistency between the two paths.

Seems like we were making a redundant put_writes call when replaying the
async tasks

---------

Co-authored-by: William FH <13333726+hinthornw@users.noreply.github.com>
This commit is contained in:
Sydney Runkle
2026-02-19 07:24:18 -08:00
committed by GitHub
co-authored by William FH
parent 5ac837d7cd
commit ca26805b5f
3 changed files with 568 additions and 1 deletions
@@ -728,7 +728,6 @@ async def _acall_impl(
)
else:
fut.set_result(None)
futures()[fut] = next_task # type: ignore[index]
else:
# schedule the next task
fut = cast(
+280
View File
@@ -5800,6 +5800,286 @@ def test_multiple_interrupts_functional_cache(
assert counter == 6
def test_task_before_interrupt_resume(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly when a @task runs
before interrupt-producing tasks in an @entrypoint.
The @task wrapper on both setup and ask is essential to reproduce the bug:
- @task on setup triggers a mid-step put_writes (creating a new pending_writes list)
- @task on ask means interrupt() runs in a child scratchpad that must
delegate to the parent for null resume consumption tracking
"""
@entrypoint(checkpointer=sync_checkpointer)
def workflow(number_of_topics: int) -> dict:
@task
def setup() -> int:
return number_of_topics
@task
def ask(question: str) -> str:
return interrupt(question)
n = setup().result()
answers = []
for i in range(n):
q = f"Whats the answer for topic {i + 1}?"
answers.append(ask(q).result())
return {"answers": answers}
config = {"configurable": {"thread_id": "1"}}
# First invocation - should get first interrupt
result = workflow.invoke(2, config=config)
assert "__interrupt__" in result
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 1?"
# Resume with answer for topic 1 - should get second interrupt
result = workflow.invoke(Command(resume="answer1"), config=config)
assert "__interrupt__" in result, f"Expected interrupt for topic 2, got: {result}"
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 2?"
# Resume with answer for topic 2 - should get final result
result = workflow.invoke(Command(resume="answer2"), config=config)
assert result == {"answers": ["answer1", "answer2"]}
def test_multiple_tasks_before_interrupt_resume(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly when multiple @tasks
run before an interrupt-producing task in an @entrypoint."""
@entrypoint(checkpointer=sync_checkpointer)
def workflow(inputs: dict) -> dict:
@task
def step_a(x: int) -> int:
return x + 1
@task
def step_b(x: int) -> int:
return x * 2
@task
def ask(question: str) -> str:
return interrupt(question)
a = step_a(inputs["x"]).result()
b = step_b(a).result()
answer = ask(f"Result so far is {b}. What next?").result()
return {"computed": b, "answer": answer}
config = {"configurable": {"thread_id": "1"}}
# First invocation - should get interrupt
result = workflow.invoke({"x": 5}, config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "Result so far is 12. What next?"
# Resume
result = workflow.invoke(Command(resume="continue"), config=config)
assert result == {"computed": 12, "answer": "continue"}
def test_no_redundant_put_writes_for_cached_task(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Cached @tasks on resume must not trigger redundant put_writes."""
from unittest.mock import patch
from langgraph.pregel._loop import PregelLoop
@task
def setup(x: int) -> int:
return x
@task
def ask(question: str) -> str:
return interrupt(question)
@entrypoint(checkpointer=sync_checkpointer)
def workflow(x: int) -> dict:
n = setup(x).result()
answer = ask(f"q{n}").result()
return {"answer": answer}
config = {"configurable": {"thread_id": "1"}}
result = workflow.invoke(1, config=config)
assert "__interrupt__" in result
put_writes_task_ids: list[str] = []
orig = PregelLoop.put_writes
def spy(self, task_id, writes):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes)
with patch.object(PregelLoop, "put_writes", spy):
result = workflow.invoke(Command(resume="ans"), config=config)
assert result == {"answer": "ans"}
# Count unique non-null task IDs that got put_writes.
# Should be exactly 2: the ask task and the entrypoint task.
# If 3, the cached setup task is being redundantly re-committed.
non_null = set(
tid for tid in put_writes_task_ids if not tid.startswith("00000000")
)
assert len(non_null) == 2, (
f"Expected 2 task IDs in put_writes (ask + entrypoint), got {len(non_null)}"
)
def test_node_before_interrupt_resume_graph_api(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly in a StateGraph when a
node runs before a node that calls interrupt(). This is the graph-API
analog of test_task_before_interrupt_resume (entrypoint API)."""
class State(TypedDict):
topics: list[str]
answers: Annotated[list[str], operator.add]
def setup(state: State) -> dict:
return {"topics": [f"topic {i + 1}" for i in range(len(state["topics"]))]}
def ask(state: State) -> dict:
answers = []
for topic in state["topics"]:
answer = interrupt(f"Whats the answer for {topic}?")
answers.append(answer)
return {"answers": answers}
graph = (
StateGraph(State)
.add_node("setup", setup)
.add_node("ask", ask)
.add_edge(START, "setup")
.add_edge("setup", "ask")
.add_edge("ask", END)
.compile(checkpointer=sync_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First invocation - setup runs, then ask interrupts on the first topic
result = graph.invoke({"topics": ["a", "b"], "answers": []}, config=config)
assert "__interrupt__" in result
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 1?"
# Resume with answer for topic 1 - should get second interrupt
result = graph.invoke(Command(resume="answer1"), config=config)
assert "__interrupt__" in result, f"Expected interrupt for topic 2, got: {result}"
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 2?"
# Resume with answer for topic 2 - should complete
result = graph.invoke(Command(resume="answer2"), config=config)
assert result == {
"topics": ["topic 1", "topic 2"],
"answers": ["answer1", "answer2"],
}
def test_multiple_nodes_before_interrupt_resume_graph_api(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly in a StateGraph when
multiple nodes run before a node that calls interrupt(). This is the
graph-API analog of test_multiple_tasks_before_interrupt_resume."""
class State(TypedDict):
value: int
answer: str
def step_a(state: State) -> dict:
return {"value": state["value"] + 1}
def step_b(state: State) -> dict:
return {"value": state["value"] * 2}
def ask(state: State) -> dict:
answer = interrupt(f"Result so far is {state['value']}. What next?")
return {"answer": answer}
graph = (
StateGraph(State)
.add_node("step_a", step_a)
.add_node("step_b", step_b)
.add_node("ask", ask)
.add_edge(START, "step_a")
.add_edge("step_a", "step_b")
.add_edge("step_b", "ask")
.add_edge("ask", END)
.compile(checkpointer=sync_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First invocation - step_a and step_b run, then ask interrupts
result = graph.invoke({"value": 5, "answer": ""}, config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "Result so far is 12. What next?"
# Resume - should complete
result = graph.invoke(Command(resume="continue"), config=config)
assert result == {"value": 12, "answer": "continue"}
def test_node_before_multiple_interrupt_cycles_graph_api(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that a node running before an interrupt node does not interfere
with multiple interrupt/resume cycles in a StateGraph."""
class State(TypedDict):
count: int
data: str
def prepare(state: State) -> dict:
return {"count": state["count"] + 10}
def multi_interrupt(state: State) -> dict:
first = interrupt("First question?")
second = interrupt("Second question?")
return {"data": f"{first},{second}"}
graph = (
StateGraph(State)
.add_node("prepare", prepare)
.add_node("multi_interrupt", multi_interrupt)
.add_edge(START, "prepare")
.add_edge("prepare", "multi_interrupt")
.add_edge("multi_interrupt", END)
.compile(checkpointer=sync_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First invocation - prepare runs, multi_interrupt hits first interrupt
result = graph.invoke({"count": 0, "data": ""}, config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "First question?"
# Resume first interrupt - hits second interrupt
result = graph.invoke(Command(resume="first_answer"), config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "Second question?"
# Resume second interrupt - completes
result = graph.invoke(Command(resume="second_answer"), config=config)
assert result == {"count": 10, "data": "first_answer,second_answer"}
def test_double_interrupt_subgraph(sync_checkpointer: BaseCheckpointSaver) -> None:
class AgentState(TypedDict):
input: str
+288
View File
@@ -7920,6 +7920,294 @@ async def test_interrupts_in_tasks_surfaced_once(
assert result[1] == "Added Will!"
@NEEDS_CONTEXTVARS
async def test_task_before_interrupt_resume(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly when a @task runs
before interrupt-producing tasks in an @entrypoint.
The @task wrapper on both setup and ask is essential to reproduce the bug:
- @task on setup triggers a mid-step put_writes (creating a new pending_writes list)
- @task on ask means interrupt() runs in a child scratchpad that must
delegate to the parent for null resume consumption tracking
"""
@entrypoint(checkpointer=async_checkpointer)
async def workflow(number_of_topics: int) -> dict:
@task
async def setup() -> int:
return number_of_topics
@task
async def ask(question: str) -> str:
return interrupt(question)
n = await setup()
answers = []
for i in range(n):
q = f"Whats the answer for topic {i + 1}?"
answers.append(await ask(q))
return {"answers": answers}
config = {"configurable": {"thread_id": "1"}}
# First invocation - should get first interrupt
result = await workflow.ainvoke(2, config=config)
assert "__interrupt__" in result
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 1?"
# Resume with answer for topic 1 - should get second interrupt
result = await workflow.ainvoke(Command(resume="answer1"), config=config)
assert "__interrupt__" in result, f"Expected interrupt for topic 2, got: {result}"
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 2?"
# Resume with answer for topic 2 - should get final result
result = await workflow.ainvoke(Command(resume="answer2"), config=config)
assert result == {"answers": ["answer1", "answer2"]}
@NEEDS_CONTEXTVARS
async def test_multiple_tasks_before_interrupt_resume(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly when multiple @tasks
run before an interrupt-producing task in an @entrypoint."""
@entrypoint(checkpointer=async_checkpointer)
async def workflow(inputs: dict) -> dict:
@task
async def step_a(x: int) -> int:
return x + 1
@task
async def step_b(x: int) -> int:
return x * 2
@task
async def ask(question: str) -> str:
return interrupt(question)
a = await step_a(inputs["x"])
b = await step_b(a)
answer = await ask(f"Result so far is {b}. What next?")
return {"computed": b, "answer": answer}
config = {"configurable": {"thread_id": "1"}}
# First invocation - should get interrupt
result = await workflow.ainvoke({"x": 5}, config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "Result so far is 12. What next?"
# Resume
result = await workflow.ainvoke(Command(resume="continue"), config=config)
assert result == {"computed": 12, "answer": "continue"}
@NEEDS_CONTEXTVARS
async def test_no_redundant_put_writes_for_cached_task(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Cached @tasks on resume must not trigger redundant put_writes."""
from unittest.mock import patch
from langgraph.pregel._loop import PregelLoop
@task
async def setup(x: int) -> int:
return x
@task
async def ask(question: str) -> str:
return interrupt(question)
@entrypoint(checkpointer=async_checkpointer)
async def workflow(x: int) -> dict:
n = await setup(x)
answer = await ask(f"q{n}")
return {"answer": answer}
config = {"configurable": {"thread_id": "1"}}
result = await workflow.ainvoke(1, config=config)
assert "__interrupt__" in result
put_writes_task_ids: list[str] = []
orig = PregelLoop.put_writes
def spy(self, task_id, writes):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes)
with patch.object(PregelLoop, "put_writes", spy):
result = await workflow.ainvoke(Command(resume="ans"), config=config)
assert result == {"answer": "ans"}
# Count unique non-null task IDs that got put_writes.
# Should be exactly 2: the ask task and the entrypoint task.
# If 3, the cached setup task is being redundantly re-committed.
non_null = set(
tid for tid in put_writes_task_ids if not tid.startswith("00000000")
)
assert len(non_null) == 2, (
f"Expected 2 task IDs in put_writes (ask + entrypoint), got {len(non_null)}"
)
@NEEDS_CONTEXTVARS
async def test_node_before_interrupt_resume_graph_api(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly in a StateGraph when a
node runs before a node that calls interrupt(). This is the graph-API
analog of test_task_before_interrupt_resume (entrypoint API)."""
class State(TypedDict):
topics: list[str]
answers: Annotated[list[str], operator.add]
def setup(state: State) -> dict:
return {"topics": [f"topic {i + 1}" for i in range(len(state["topics"]))]}
def ask(state: State) -> dict:
answers = []
for topic in state["topics"]:
answer = interrupt(f"Whats the answer for {topic}?")
answers.append(answer)
return {"answers": answers}
graph = (
StateGraph(State)
.add_node("setup", setup)
.add_node("ask", ask)
.add_edge(START, "setup")
.add_edge("setup", "ask")
.add_edge("ask", END)
.compile(checkpointer=async_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First invocation - setup runs, then ask interrupts on the first topic
result = await graph.ainvoke(
{"topics": ["a", "b"], "answers": []}, config=config
)
assert "__interrupt__" in result
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 1?"
# Resume with answer for topic 1 - should get second interrupt
result = await graph.ainvoke(Command(resume="answer1"), config=config)
assert "__interrupt__" in result, f"Expected interrupt for topic 2, got: {result}"
assert len(result["__interrupt__"]) == 1
assert result["__interrupt__"][0].value == "Whats the answer for topic 2?"
# Resume with answer for topic 2 - should complete
result = await graph.ainvoke(Command(resume="answer2"), config=config)
assert result == {
"topics": ["topic 1", "topic 2"],
"answers": ["answer1", "answer2"],
}
@NEEDS_CONTEXTVARS
async def test_multiple_nodes_before_interrupt_resume_graph_api(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that Command(resume=value) works correctly in a StateGraph when
multiple nodes run before a node that calls interrupt(). This is the
graph-API analog of test_multiple_tasks_before_interrupt_resume."""
class State(TypedDict):
value: int
answer: str
def step_a(state: State) -> dict:
return {"value": state["value"] + 1}
def step_b(state: State) -> dict:
return {"value": state["value"] * 2}
def ask(state: State) -> dict:
answer = interrupt(f"Result so far is {state['value']}. What next?")
return {"answer": answer}
graph = (
StateGraph(State)
.add_node("step_a", step_a)
.add_node("step_b", step_b)
.add_node("ask", ask)
.add_edge(START, "step_a")
.add_edge("step_a", "step_b")
.add_edge("step_b", "ask")
.add_edge("ask", END)
.compile(checkpointer=async_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First invocation - step_a and step_b run, then ask interrupts
result = await graph.ainvoke({"value": 5, "answer": ""}, config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "Result so far is 12. What next?"
# Resume - should complete
result = await graph.ainvoke(Command(resume="continue"), config=config)
assert result == {"value": 12, "answer": "continue"}
@NEEDS_CONTEXTVARS
async def test_node_before_multiple_interrupt_cycles_graph_api(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that a node running before an interrupt node does not interfere
with multiple interrupt/resume cycles in a StateGraph."""
class State(TypedDict):
count: int
data: str
def prepare(state: State) -> dict:
return {"count": state["count"] + 10}
def multi_interrupt(state: State) -> dict:
first = interrupt("First question?")
second = interrupt("Second question?")
return {"data": f"{first},{second}"}
graph = (
StateGraph(State)
.add_node("prepare", prepare)
.add_node("multi_interrupt", multi_interrupt)
.add_edge(START, "prepare")
.add_edge("prepare", "multi_interrupt")
.add_edge("multi_interrupt", END)
.compile(checkpointer=async_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First invocation - prepare runs, multi_interrupt hits first interrupt
result = await graph.ainvoke({"count": 0, "data": ""}, config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "First question?"
# Resume first interrupt - hits second interrupt
result = await graph.ainvoke(Command(resume="first_answer"), config=config)
assert "__interrupt__" in result
assert result["__interrupt__"][0].value == "Second question?"
# Resume second interrupt - completes
result = await graph.ainvoke(Command(resume="second_answer"), config=config)
assert result == {"count": 10, "data": "first_answer,second_answer"}
async def test_pregel_loop_refcount():
gc.collect()
try: