Fix unexpected re-use of null resume value by subgraphs

- Also stop exposing writes in config, in favor of scratchpad
This commit is contained in:
Nuno Campos
2025-01-15 13:31:49 -08:00
parent aab6fdf3f3
commit 3626478029
7 changed files with 321 additions and 85 deletions
+1 -3
View File
@@ -75,9 +75,7 @@ CONFIG_KEY_CHECKPOINT_ID = sys.intern("checkpoint_id")
CONFIG_KEY_CHECKPOINT_NS = sys.intern("checkpoint_ns")
# holds the current checkpoint_ns, "" for root graph
CONFIG_KEY_NODE_FINISHED = sys.intern("__pregel_node_finished")
# holds the value that "answers" an interrupt() call
CONFIG_KEY_WRITES = sys.intern("__pregel_writes")
# read-only list of existing task writes
# holds a callback to be called when a node is finished
CONFIG_KEY_SCRATCHPAD = sys.intern("__pregel_scratchpad")
# holds a mutable dict for temporary storage scoped to the current task
+33 -22
View File
@@ -42,10 +42,10 @@ from langgraph.constants import (
CONFIG_KEY_SEND,
CONFIG_KEY_STORE,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_WRITES,
EMPTY_SEQ,
ERROR,
INTERRUPT,
MISSING,
NO_WRITES,
NS_END,
NS_SEP,
@@ -71,6 +71,7 @@ from langgraph.types import (
All,
LoopProtocol,
PregelExecutableTask,
PregelScratchpad,
PregelTask,
RetryPolicy,
)
@@ -502,13 +503,10 @@ def prepare_single_task(
},
CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_WRITES: [
w
for w in pending_writes
+ configurable.get(CONFIG_KEY_WRITES, [])
if w[0] in (NULL_TASK_ID, task_id)
],
CONFIG_KEY_SCRATCHPAD: {},
CONFIG_KEY_SCRATCHPAD: _scratchpad(
pending_writes,
task_id,
),
},
),
triggers,
@@ -614,13 +612,10 @@ def prepare_single_task(
},
CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_WRITES: [
w
for w in pending_writes
+ configurable.get(CONFIG_KEY_WRITES, [])
if w[0] in (NULL_TASK_ID, task_id)
],
CONFIG_KEY_SCRATCHPAD: {},
CONFIG_KEY_SCRATCHPAD: _scratchpad(
pending_writes,
task_id,
),
},
),
triggers,
@@ -738,13 +733,10 @@ def prepare_single_task(
},
CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_WRITES: [
w
for w in pending_writes
+ configurable.get(CONFIG_KEY_WRITES, [])
if w[0] in (NULL_TASK_ID, task_id)
],
CONFIG_KEY_SCRATCHPAD: {},
CONFIG_KEY_SCRATCHPAD: _scratchpad(
pending_writes,
task_id,
),
},
),
triggers,
@@ -758,6 +750,25 @@ def prepare_single_task(
return PregelTask(task_id, name, task_path[:3])
def _scratchpad(
pending_writes: Sequence[PendingWrite],
task_id: str,
) -> PregelScratchpad:
return PregelScratchpad(
# call
call_counter=0,
# interrupt
interrupt_counter=-1,
resume=next(
(w[2] for w in pending_writes if w[0] == task_id and w[1] == RESUME), []
),
null_resume=next(
(w[2] for w in pending_writes if w[0] == NULL_TASK_ID and w[1] == RESUME),
MISSING,
),
)
def _proc_input(
proc: PregelNode,
managed: ManagedValueMapping,
+10
View File
@@ -47,12 +47,14 @@ from langgraph.constants import (
CONFIG_KEY_DELEGATE,
CONFIG_KEY_ENSURE_LATEST,
CONFIG_KEY_RESUMING,
CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_STREAM,
CONFIG_KEY_TASK_ID,
EMPTY_SEQ,
ERROR,
INPUT,
INTERRUPT,
MISSING,
NS_SEP,
NULL_TASK_ID,
PUSH,
@@ -556,8 +558,16 @@ class PregelLoop(LoopProtocol):
)
)
# take resume value from parent
if scratchpad := configurable.get(CONFIG_KEY_SCRATCHPAD):
if scratchpad["null_resume"] is not MISSING:
self.put_writes(NULL_TASK_ID, [(RESUME, scratchpad["null_resume"])])
# map command to writes
if isinstance(self.input, Command):
if self.input.resume is not None and not self.checkpointer:
raise RuntimeError(
"Cannot use Command(resume=...) without checkpointer"
)
writes: defaultdict[str, list[tuple[str, Any]]] = defaultdict(list)
# group writes by task ID
for tid, c, v in map_command(self.input, self.checkpoint_pending_writes):
+10 -12
View File
@@ -107,12 +107,11 @@ class PregelRunner:
elif next_task.writes:
# if it already ran, return the result
fut = concurrent.futures.Future()
if (
val := next(
(v for c, v in next_task.writes if c == RETURN), MISSING
)
) and val is not MISSING:
fut.set_result(val)
ret = next(
(v for c, v in next_task.writes if c == RETURN), MISSING
)
if ret is not MISSING:
fut.set_result(ret)
elif exc := next(
(v for c, v in next_task.writes if c == ERROR), None
):
@@ -295,12 +294,11 @@ class PregelRunner:
elif next_task.writes:
# if it already ran, return the result
fut = asyncio.Future()
if (
val := next(
(v for c, v in next_task.writes if c == RETURN), MISSING
)
) and val is not MISSING:
fut.set_result(val)
ret = next(
(v for c, v in next_task.writes if c == RETURN), MISSING
)
if ret is not MISSING:
fut.set_result(ret)
elif exc := next(
(v for c, v in next_task.writes if c == ERROR), None
):
+16 -30
View File
@@ -21,11 +21,7 @@ from typing import (
from langchain_core.runnables import Runnable, RunnableConfig
from typing_extensions import Self, TypedDict
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
CheckpointMetadata,
PendingWrite,
)
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
if TYPE_CHECKING:
from langgraph.store.base import BaseStore
@@ -341,13 +337,13 @@ class LoopProtocol:
self.stop = stop
class PregelScratchpad(TypedDict, total=False):
# interrupt
interrupt_counter: int
used_null_resume: bool
resume: list[Any]
class PregelScratchpad(TypedDict):
# call
call_counter: int
# interrupt
interrupt_counter: int
resume: list[Any]
null_resume: Any
def interrupt(value: Any) -> Any:
@@ -449,10 +445,8 @@ def interrupt(value: Any) -> Any:
CONFIG_KEY_CHECKPOINT_NS,
CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_SEND,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_WRITES,
MISSING,
NS_SEP,
NULL_TASK_ID,
RESUME,
)
from langgraph.errors import GraphInterrupt
@@ -461,29 +455,21 @@ def interrupt(value: Any) -> Any:
conf = get_config()["configurable"]
# track interrupt index
scratchpad: PregelScratchpad = conf[CONFIG_KEY_SCRATCHPAD]
if "interrupt_counter" not in scratchpad:
scratchpad["interrupt_counter"] = 0
else:
scratchpad["interrupt_counter"] += 1
print("interrupt", scratchpad)
scratchpad["interrupt_counter"] += 1
idx = scratchpad["interrupt_counter"]
# find previous resume values
task_id = conf[CONFIG_KEY_TASK_ID]
writes: list[PendingWrite] = conf[CONFIG_KEY_WRITES]
scratchpad.setdefault(
"resume", next((w[2] for w in writes if w[0] == task_id and w[1] == RESUME), [])
)
if scratchpad["resume"]:
if idx < len(scratchpad["resume"]):
return scratchpad["resume"][idx]
# find current resume value
if not scratchpad.get("used_null_resume"):
scratchpad["used_null_resume"] = True
for tid, c, v in sorted(writes, key=lambda x: x[0], reverse=True):
if tid == NULL_TASK_ID and c == RESUME:
assert len(scratchpad["resume"]) == idx, (scratchpad["resume"], idx)
scratchpad["resume"].append(v)
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad["resume"])])
return v
if scratchpad["null_resume"] is not MISSING:
assert len(scratchpad["resume"]) == idx, (scratchpad["resume"], idx)
v = scratchpad["null_resume"]
scratchpad["null_resume"] = MISSING
scratchpad["resume"].append(v)
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad["resume"])])
return v
# no resume value found
raise GraphInterrupt(
(
+122 -8
View File
@@ -5262,9 +5262,10 @@ def test_multiple_updates() -> None:
]
def test_falsy_return_from_task() -> None:
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_name: str):
"""Test with a falsy return from a task."""
checkpointer = MemorySaver()
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
@task
def falsy_task() -> bool:
@@ -5276,17 +5277,18 @@ def test_falsy_return_from_task() -> None:
falsy_task().result()
interrupt("test")
configurable = {"configurable": {"thread_id": uuid.uuid4()}}
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
graph.invoke({"a": 5}, configurable)
graph.invoke(Command(resume="123"), configurable)
def test_multiple_interrupts_imperative() -> None:
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_multiple_interrupts_imperative(
request: pytest.FixtureRequest, checkpointer_name: str
):
"""Test multiple interrupts with an imperative API."""
from langgraph.checkpoint.memory import MemorySaver
from langgraph.func import entrypoint, task
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
checkpointer = MemorySaver()
counter = 0
@task
@@ -5307,7 +5309,7 @@ def test_multiple_interrupts_imperative() -> None:
return {"values": values}
configurable = {"configurable": {"thread_id": uuid.uuid4()}}
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
graph.invoke({}, configurable)
graph.invoke(Command(resume="a"), configurable)
graph.invoke(Command(resume="b"), configurable)
@@ -5317,3 +5319,115 @@ def test_multiple_interrupts_imperative() -> None:
"values": [2, "a", 4, "b", 6, "c"],
}
assert counter == 3
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_double_interrupt_subgraph(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class AgentState(TypedDict):
input: str
def node_1(state: AgentState):
result = interrupt("interrupt node 1")
return {"input": result}
def node_2(state: AgentState):
result = interrupt("interrupt node 2")
return {"input": result}
subgraph_builder = (
StateGraph(AgentState)
.add_node("node_1", node_1)
.add_node("node_2", node_2)
.add_edge(START, "node_1")
.add_edge("node_1", "node_2")
.add_edge("node_2", END)
)
# invoke the sub graph
subgraph = subgraph_builder.compile(checkpointer=checkpointer)
thread = {"configurable": {"thread_id": str(uuid.uuid4())}}
assert [c for c in subgraph.stream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [c for c in subgraph.stream(Command(resume="123"), thread)] == [
{
"node_1": {"input": "123"},
},
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("node_2:")],
when="during",
),
)
},
]
# resume from the second interrupt
assert [c for c in subgraph.stream(Command(resume="123"), thread)] == [
{
"node_2": {"input": "123"},
},
]
subgraph = subgraph_builder.compile()
def invoke_sub_agent(state: AgentState):
return subgraph.invoke(state)
parent_agent = (
StateGraph(AgentState)
.add_node("invoke_sub_agent", invoke_sub_agent)
.add_edge(START, "invoke_sub_agent")
.add_edge("invoke_sub_agent", END)
.compile(checkpointer=checkpointer)
)
assert [c for c in parent_agent.stream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [c for c in parent_agent.stream(Command(resume=True), thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_2:")],
when="during",
),
)
}
]
# resume from 2nd interrupt
assert [c for c in parent_agent.stream(Command(resume=True), thread)] == [
{
"invoke_sub_agent": {"input": True},
},
]
+129 -10
View File
@@ -6696,23 +6696,25 @@ async def test_multiple_updates() -> None:
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
async def test_falsy_return_from_task() -> None:
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_falsy_return_from_task(checkpointer_name: str) -> None:
"""Test with a falsy return from a task."""
checkpointer = MemorySaver()
@task
async def falsy_task() -> bool:
return False
@entrypoint(checkpointer=checkpointer)
async def graph(state: dict) -> dict:
"""React tool."""
await falsy_task()
interrupt("test")
async with awith_checkpointer(checkpointer_name) as checkpointer:
configurable = {"configurable": {"thread_id": uuid.uuid4()}}
await graph.ainvoke({"a": 5}, configurable)
await graph.ainvoke(Command(resume="123"), configurable)
@entrypoint(checkpointer=checkpointer)
async def graph(state: dict) -> dict:
"""React tool."""
await falsy_task()
interrupt("test")
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
await graph.ainvoke({"a": 5}, configurable)
await graph.ainvoke(Command(resume="123"), configurable)
@pytest.mark.skipif(
@@ -6756,3 +6758,120 @@ async def test_multiple_interrupts_imperative(checkpointer_name: str) -> None:
"values": [2, "a", 4, "b", 6, "c"],
}
assert counter == 3
@pytest.mark.skipif(
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_double_interrupt_subgraph(checkpointer_name: str) -> None:
class AgentState(TypedDict):
input: str
def node_1(state: AgentState):
result = interrupt("interrupt node 1")
return {"input": result}
def node_2(state: AgentState):
result = interrupt("interrupt node 2")
return {"input": result}
subgraph_builder = (
StateGraph(AgentState)
.add_node("node_1", node_1)
.add_node("node_2", node_2)
.add_edge(START, "node_1")
.add_edge("node_1", "node_2")
.add_edge("node_2", END)
)
async with awith_checkpointer(checkpointer_name) as checkpointer:
# invoke the sub graph
subgraph = subgraph_builder.compile(checkpointer=checkpointer)
thread = {"configurable": {"thread_id": str(uuid.uuid4())}}
assert [c async for c in subgraph.astream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [c async for c in subgraph.astream(Command(resume="123"), thread)] == [
{
"node_1": {"input": "123"},
},
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("node_2:")],
when="during",
),
)
},
]
# resume from the second interrupt
assert [c async for c in subgraph.astream(Command(resume="123"), thread)] == [
{
"node_2": {"input": "123"},
},
]
subgraph = subgraph_builder.compile()
def invoke_sub_agent(state: AgentState):
return subgraph.invoke(state)
parent_agent = (
StateGraph(AgentState)
.add_node("invoke_sub_agent", invoke_sub_agent)
.add_edge(START, "invoke_sub_agent")
.add_edge("invoke_sub_agent", END)
.compile(checkpointer=checkpointer)
)
assert [c async for c in parent_agent.astream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [
c async for c in parent_agent.astream(Command(resume=True), thread)
] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_2:")],
when="during",
),
)
}
]
# resume from 2nd interrupt
assert [
c async for c in parent_agent.astream(Command(resume=True), thread)
] == [
{
"invoke_sub_agent": {"input": True},
},
]