This commit is contained in:
Nuno Campos
2025-01-22 16:02:09 -08:00
parent 020d10138d
commit 3a997be088
9 changed files with 87 additions and 132 deletions
+13 -9
View File
@@ -1,3 +1,4 @@
import functools
import sys
from collections import defaultdict, deque
from functools import partial
@@ -46,7 +47,6 @@ from langgraph.constants import (
EMPTY_SEQ,
ERROR,
INTERRUPT,
MISSING,
NO_WRITES,
NS_END,
NS_SEP,
@@ -341,7 +341,7 @@ def prepare_next_tasks(
@overload
def prepare_next_tasks(
checkpoint: Checkpoint,
pending_writes: Sequence[PendingWrite],
pending_writes: list[PendingWrite],
processes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
managed: ManagedValueMapping,
@@ -357,7 +357,7 @@ def prepare_next_tasks(
def prepare_next_tasks(
checkpoint: Checkpoint,
pending_writes: Sequence[PendingWrite],
pending_writes: list[PendingWrite],
processes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
managed: ManagedValueMapping,
@@ -418,7 +418,7 @@ def prepare_single_task(
task_id_checksum: Optional[str],
*,
checkpoint: Checkpoint,
pending_writes: Sequence[PendingWrite],
pending_writes: list[PendingWrite],
processes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
managed: ManagedValueMapping,
@@ -761,9 +761,13 @@ def prepare_single_task(
def _scratchpad(
pending_writes: Sequence[PendingWrite],
pending_writes: list[PendingWrite],
task_id: str,
) -> PregelScratchpad:
null_resume_write = next(
(w for w in pending_writes if w[0] == NULL_TASK_ID and w[1] == RESUME), None
)
return PregelScratchpad(
# call
call_counter=0,
@@ -772,10 +776,10 @@ def _scratchpad(
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,
),
null_resume=null_resume_write[2] if null_resume_write is not None else None,
_consume_null_resume=functools.partial(pending_writes.remove, null_resume_write)
if null_resume_write is not None
else lambda: None,
# subgraph
subgraph_counter=0,
)
+1 -1
View File
@@ -89,7 +89,7 @@ def map_command(
raise TypeError(
f"In Command.goto, expected Send/str, got {type(send).__name__}"
)
if cmd.resume:
if cmd.resume is not None:
if isinstance(cmd.resume, dict) and all(is_task_id(k) for k in cmd.resume):
for tid, resume in cmd.resume.items():
existing: list[Any] = next(
+8 -7
View File
@@ -54,7 +54,6 @@ from langgraph.constants import (
ERROR,
INPUT,
INTERRUPT,
MISSING,
NS_SEP,
NULL_TASK_ID,
PUSH,
@@ -230,19 +229,19 @@ class PregelLoop(LoopProtocol):
self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM])
scratchpad: Optional[PregelScratchpad] = config[CONF].get(CONFIG_KEY_SCRATCHPAD)
if not self.config[CONF].get(CONFIG_KEY_DELEGATE) and scratchpad is not None:
if scratchpad["subgraph_counter"]:
if scratchpad.subgraph_counter:
self.config = patch_configurable(
self.config,
{
CONFIG_KEY_CHECKPOINT_NS: NS_SEP.join(
(
config[CONF][CONFIG_KEY_CHECKPOINT_NS],
str(scratchpad["subgraph_counter"]),
str(scratchpad.subgraph_counter),
)
)
},
)
scratchpad["subgraph_counter"] += 1
scratchpad.subgraph_counter += 1
if not self.is_nested and config[CONF].get(CONFIG_KEY_CHECKPOINT_NS):
self.config = patch_configurable(
self.config,
@@ -563,9 +562,11 @@ 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"])])
if scratchpad := cast(
Optional[PregelScratchpad], configurable.get(CONFIG_KEY_SCRATCHPAD)
):
if scratchpad.null_resume is not None:
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:
+7 -9
View File
@@ -39,7 +39,7 @@ from langgraph.errors import GraphBubbleUp, GraphInterrupt
from langgraph.pregel.algo import Call
from langgraph.pregel.executor import Submit
from langgraph.pregel.retry import arun_with_retry, run_with_retry
from langgraph.types import PregelExecutableTask, RetryPolicy
from langgraph.types import PregelExecutableTask, PregelScratchpad, RetryPolicy
from langgraph.utils.future import chain_future
F = TypeVar("F", concurrent.futures.Future, asyncio.Future)
@@ -135,8 +135,7 @@ class PregelRunner:
return task.config[CONF][CONFIG_KEY_SEND](writes)
# schedule PUSH tasks, collect futures
scratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
scratchpad.setdefault("call_counter", 0)
scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
rtn: dict[int, Optional[concurrent.futures.Future]] = {}
for idx, w in enumerate(writes):
# bail if not a PUSH write
@@ -144,8 +143,8 @@ class PregelRunner:
continue
# schedule the next task, if the callback returns one
wcall = calls[idx] if calls else None
cnt = scratchpad["call_counter"]
scratchpad["call_counter"] += 1
cnt = scratchpad.call_counter
scratchpad.call_counter += 1
if next_task := self.schedule_task(task, cnt, wcall):
if fut := next(
(
@@ -324,8 +323,7 @@ class PregelRunner:
return task.config[CONF][CONFIG_KEY_SEND](writes)
# schedule PUSH tasks, collect futures
scratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
scratchpad.setdefault("call_counter", 0)
scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
rtn: dict[int, Optional[asyncio.Future]] = {}
for idx, w in enumerate(writes):
# bail if not a PUSH write
@@ -333,8 +331,8 @@ class PregelRunner:
continue
# schedule the next task, if the callback returns one
wcall = calls[idx] if calls is not None else None
cnt = scratchpad["call_counter"]
scratchpad["call_counter"] += 1
cnt = scratchpad.call_counter
scratchpad.call_counter += 1
if next_task := self.schedule_task(task, cnt, wcall):
# if the parent task was retried,
# the next task might already be running
+24 -15
View File
@@ -19,7 +19,7 @@ from typing import (
)
from langchain_core.runnables import Runnable, RunnableConfig
from typing_extensions import Self, TypedDict
from typing_extensions import Self
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
@@ -339,16 +339,26 @@ class LoopProtocol:
self.stop = stop
class PregelScratchpad(TypedDict):
@dataclasses.dataclass(**{**_DC_KWARGS, "frozen": False})
class PregelScratchpad:
# call
call_counter: int
# interrupt
interrupt_counter: int
resume: list[Any]
null_resume: Any
null_resume: Optional[Any]
_consume_null_resume: Callable[[], None]
# subgraph
subgraph_counter: int
def consume_null_resume(self) -> Any:
if self.null_resume is not None:
value = self.null_resume
self._consume_null_resume()
self.null_resume = None
return value
raise ValueError("No null resume to consume")
def interrupt(value: Any) -> Any:
"""Interrupt the graph with a resumable exception from within a node.
@@ -449,7 +459,6 @@ def interrupt(value: Any) -> Any:
CONFIG_KEY_CHECKPOINT_NS,
CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_SEND,
MISSING,
NS_SEP,
RESUME,
)
@@ -459,19 +468,19 @@ def interrupt(value: Any) -> Any:
conf = get_config()["configurable"]
# track interrupt index
scratchpad: PregelScratchpad = conf[CONFIG_KEY_SCRATCHPAD]
scratchpad["interrupt_counter"] += 1
idx = scratchpad["interrupt_counter"]
scratchpad.interrupt_counter += 1
idx = scratchpad.interrupt_counter
# find previous resume values
if scratchpad["resume"]:
if idx < len(scratchpad["resume"]):
return scratchpad["resume"][idx]
if scratchpad.resume:
if idx < len(scratchpad.resume):
return scratchpad.resume[idx]
# find current resume value
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"])])
if scratchpad.null_resume is not None:
assert len(scratchpad.resume) == idx, (scratchpad.resume, idx)
print("consume null resume", scratchpad.null_resume)
v = scratchpad.consume_null_resume()
scratchpad.resume.append(v)
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
return v
# no resume value found
raise GraphInterrupt(
+4 -5
View File
@@ -5082,7 +5082,7 @@ def test_interrupt_task_functional(
config = {"configurable": {"thread_id": "1"}}
# First run, interrupted at bar
graph.invoke({"a": ""}, config)
assert not graph.invoke({"a": ""}, config)
# Resume with an answer
res = graph.invoke(Command(resume="bar"), config)
assert res == {"a": "foobar"}
@@ -5098,11 +5098,10 @@ def test_interrupt_task_functional(
return baz_result
# First run, interrupted at bar
graph.invoke({"a": ""}, config)
assert not graph.invoke({"a": ""}, config)
# Provide resumes
graph.invoke(Command(resume="bar"), config)
res_2 = graph.invoke(Command(resume="baz"), config)
assert res_2 == {"a": "foobarbaz"} # Errors, produces "foobarbar" instead
assert not graph.invoke(Command(resume="bar"), config)
assert graph.invoke(Command(resume="baz"), config) == {"a": "foobarbaz"}
def test_root_mixed_return() -> None:
+16
View File
@@ -6408,6 +6408,22 @@ async def test_interrupt_task_functional(checkpointer_name: str) -> None:
res = await graph.ainvoke(Command(resume="bar"), config)
assert res == {"a": "foobar"}
# Test that we can interrupt the same task multiple times
config = {"configurable": {"thread_id": "2"}}
@entrypoint(checkpointer=checkpointer)
async def graph(inputs: dict) -> dict:
foo_result = await foo(inputs)
bar_result = await bar(foo_result)
baz_result = await bar(bar_result)
return baz_result
# First run, interrupted at bar
assert not await graph.ainvoke({"a": ""}, config)
# Provide resumes
assert not await graph.ainvoke(Command(resume="bar"), config)
assert await graph.ainvoke(Command(resume="baz"), config) == {"a": "foobarbaz"}
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_command_with_static_breakpoints(checkpointer_name: str) -> None:
+7 -43
View File
@@ -15,7 +15,7 @@ from langgraph.graph.state import StateGraph
from langgraph.pregel import Pregel
from langgraph.scheduler.kafka import serde
from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics
from tests.any import AnyDict, AnyInt
from tests.any import AnyDict
from tests.drain import drain_topics_async
from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage
@@ -198,13 +198,7 @@ async def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": None,
"checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"]
@@ -271,13 +265,7 @@ async def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"]
@@ -374,13 +362,7 @@ async def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"]
@@ -487,13 +469,7 @@ async def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": None,
"checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"]
@@ -555,13 +531,7 @@ async def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"]
@@ -679,13 +649,7 @@ async def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"]
@@ -15,7 +15,7 @@ from langgraph.pregel import Pregel
from langgraph.scheduler.kafka import serde
from langgraph.scheduler.kafka.default_sync import DefaultProducer
from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics
from tests.any import AnyDict, AnyInt
from tests.any import AnyDict
from tests.drain import drain_topics
from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage
@@ -197,13 +197,7 @@ def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": None,
"checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"]
@@ -270,13 +264,7 @@ def test_subgraph_w_interrupt(
"__pregel_resuming": False,
"__pregel_previous": None,
"__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"]
@@ -373,13 +361,7 @@ def test_subgraph_w_interrupt(
"__pregel_resuming": False,
"__pregel_task_id": history[0].tasks[0].id,
"__pregel_previous": None,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"]
@@ -485,13 +467,7 @@ def test_subgraph_w_interrupt(
"__pregel_resuming": True,
"__pregel_previous": None,
"__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": None,
"checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"]
@@ -553,13 +529,7 @@ def test_subgraph_w_interrupt(
"__pregel_resuming": True,
"__pregel_previous": None,
"__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"]
@@ -677,13 +647,7 @@ def test_subgraph_w_interrupt(
"__pregel_previous": None,
"__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {
"subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"__pregel_scratchpad": None,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"]