mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 05:25:05 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a9a10dedaf | ||
|
|
89ff2d33de |
@@ -47,6 +47,16 @@ def exit_delta_task_id(step: int, task_id: str) -> str:
|
||||
return f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
|
||||
|
||||
|
||||
def exit_delta_late_task_id(step: int, task_id: str) -> str:
|
||||
"""Synthetic task id for exit-mode writes of a superstep after the anchor's own.
|
||||
|
||||
Sorts after every real task id, in step order, so replay keeps them after
|
||||
the anchor's own superstep whether a saver orders by task path or task id.
|
||||
"""
|
||||
parts = str(uuid.UUID(task_id)).split("-")
|
||||
return f"ffffffff-{step >> 16:04x}-{step & 0xFFFF:04x}-{parts[3]}-{parts[4]}"
|
||||
|
||||
|
||||
def delta_channels_to_snapshot(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
counters_since_delta_snapshot: Mapping[str, tuple[int, int]],
|
||||
|
||||
@@ -103,6 +103,7 @@ from langgraph.pregel._checkpoint import (
|
||||
create_checkpoint,
|
||||
delta_channels_to_snapshot,
|
||||
empty_checkpoint,
|
||||
exit_delta_late_task_id,
|
||||
exit_delta_task_id,
|
||||
)
|
||||
from langgraph.pregel._executor import (
|
||||
@@ -217,10 +218,15 @@ class PregelLoop:
|
||||
# `after_tick`). At exit, `_put_exit_delta_writes` filters out channels
|
||||
# that will snapshot, then persists the rest under an anchor parent.
|
||||
# `None` when not in exit mode (so the capture sites are no-ops).
|
||||
# Each tuple is `(step, task_id, channel, value)` — `step` drives the
|
||||
# synthetic step-prefixed task_id used to preserve chronological order
|
||||
# under the saver's `ORDER BY task_id, idx` sorting.
|
||||
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
|
||||
# Each tuple is `(step, task_id, task_path, channel, value)`; see
|
||||
# `_put_exit_delta_writes` for how they are ordered.
|
||||
_exit_delta_writes: list[tuple[int, str, str, str, Any]] | None = None
|
||||
|
||||
# The pending writes loaded with the checkpoint, already stored on it, kept
|
||||
# alive so their ids stay unique; and the checkpoint's own superstep, the
|
||||
# first one this run ticks.
|
||||
_loaded_write_ids: dict[int, tuple[str, str, Any]]
|
||||
_exit_first_step: int | None = None
|
||||
|
||||
# Delta channels that saw an Overwrite since the last checkpoint. These
|
||||
# channels must snapshot after live update applies overwrite semantics so
|
||||
@@ -707,9 +713,18 @@ class PregelLoop:
|
||||
)
|
||||
# capture delta-channel writes for exit-mode accumulator before clearing
|
||||
if self._exit_delta_writes is not None:
|
||||
for tid, ch, v in self.checkpoint_pending_writes:
|
||||
if isinstance(self.specs.get(ch), DeltaChannel):
|
||||
self._exit_delta_writes.append((self.step, tid, ch, v))
|
||||
if self._exit_first_step is None:
|
||||
self._exit_first_step = self.step
|
||||
for w in self.checkpoint_pending_writes:
|
||||
tid, ch, v = w
|
||||
if not isinstance(self.specs.get(ch), DeltaChannel):
|
||||
continue
|
||||
if id(w) in self._loaded_write_ids:
|
||||
continue
|
||||
task = self.tasks.get(tid)
|
||||
path = task_path_str(task.path) if task else ""
|
||||
self._exit_delta_writes.append((self.step, tid, path, ch, v))
|
||||
self._loaded_write_ids = {}
|
||||
# clear pending writes
|
||||
self.checkpoint_pending_writes.clear()
|
||||
# only replay (re-execute) done tasks on the first tick
|
||||
@@ -848,6 +863,7 @@ class PregelLoop:
|
||||
def _first(
|
||||
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
|
||||
) -> set[str] | None:
|
||||
self._loaded_write_ids = {id(w): w for w in self.checkpoint_pending_writes}
|
||||
# Resuming from a previous checkpoint requires two things:
|
||||
# 1. A prior checkpoint exists (channel_versions is non-empty)
|
||||
# 2. The input signals continuation (not a fresh run with new input)
|
||||
@@ -1017,7 +1033,9 @@ class PregelLoop:
|
||||
if self._exit_delta_writes is not None:
|
||||
for c, v in input_writes:
|
||||
if isinstance(self.specs.get(c), DeltaChannel):
|
||||
self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v))
|
||||
self._exit_delta_writes.append(
|
||||
(self.step, NULL_TASK_ID, "", c, v)
|
||||
)
|
||||
# Persist delta-channel input writes so sub-freq inputs are
|
||||
# recoverable via ancestor walk (mirrors the Command input path).
|
||||
if self.durability != "exit":
|
||||
@@ -1243,9 +1261,7 @@ class PregelLoop:
|
||||
)
|
||||
|
||||
pending = [
|
||||
(step, tid, ch, v)
|
||||
for (step, tid, ch, v) in self._exit_delta_writes
|
||||
if ch not in channels_to_snapshot
|
||||
w for w in self._exit_delta_writes if w[3] not in channels_to_snapshot
|
||||
]
|
||||
if not pending:
|
||||
return
|
||||
@@ -1280,11 +1296,19 @@ class PregelLoop:
|
||||
# sees the stub as its parent.
|
||||
self.checkpoint_config = anchor_config
|
||||
|
||||
# Step-prefixed synthetic task_id preserves chronological superstep
|
||||
# order under the saver's ORDER BY task_id, idx sorting.
|
||||
grouped: dict[tuple[int, str], list[tuple[str, Any]]] = {}
|
||||
for step, tid, ch, v in pending:
|
||||
grouped.setdefault((step, tid), []).append((ch, v))
|
||||
# The checkpoint's own superstep is stored as sync durability stores
|
||||
# it, so it interleaves with the writes a resume loaded from it. Later
|
||||
# supersteps sort after every real task path and task id, in step
|
||||
# order, so this holds whether a saver orders by path or by id.
|
||||
grouped: dict[tuple[str, str], list[tuple[str, Any]]] = {}
|
||||
for step, tid, path, ch, v in pending:
|
||||
if tid == NULL_TASK_ID:
|
||||
key = (exit_delta_task_id(step, tid), "")
|
||||
elif step == self._exit_first_step:
|
||||
key = (tid, path)
|
||||
else:
|
||||
key = (exit_delta_late_task_id(step, tid), f"~~{step:010d}{path}")
|
||||
grouped.setdefault(key, []).append((ch, v))
|
||||
anchor_write_config = patch_configurable(
|
||||
anchor_config,
|
||||
{
|
||||
@@ -1294,22 +1318,21 @@ class PregelLoop:
|
||||
CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID],
|
||||
},
|
||||
)
|
||||
for (step, tid), entries in grouped.items():
|
||||
synth_tid = exit_delta_task_id(step, tid)
|
||||
for (tid, path), entries in grouped.items():
|
||||
if self.checkpointer_put_writes_accepts_task_path:
|
||||
fut = self.submit(
|
||||
self.checkpointer_put_writes,
|
||||
anchor_write_config,
|
||||
entries,
|
||||
synth_tid,
|
||||
"",
|
||||
tid,
|
||||
path,
|
||||
)
|
||||
else:
|
||||
fut = self.submit(
|
||||
self.checkpointer_put_writes,
|
||||
anchor_write_config,
|
||||
entries,
|
||||
synth_tid,
|
||||
tid,
|
||||
)
|
||||
if self._delta_write_futs is not None:
|
||||
self._delta_write_futs.append(fut)
|
||||
|
||||
@@ -6,11 +6,13 @@ channel), lazy stub creation when no parent exists, and proper read-path
|
||||
reconstruction via ancestor walks.
|
||||
"""
|
||||
|
||||
import operator
|
||||
import uuid
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import TypedDict
|
||||
@@ -19,6 +21,7 @@ from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
from langgraph.pregel._checkpoint import exit_delta_task_id
|
||||
from langgraph.types import Command, Durability, interrupt
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
@@ -389,3 +392,112 @@ async def test_exit_snapshot_then_tail_deltas() -> None:
|
||||
assert "seed-msg" in contents
|
||||
assert "tail-msg" in contents
|
||||
assert contents.index("seed-msg") < contents.index("tail-msg")
|
||||
|
||||
|
||||
def _append(current: list, writes: list) -> list:
|
||||
out = list(current)
|
||||
for write in writes:
|
||||
out.extend(write)
|
||||
return out
|
||||
|
||||
|
||||
class _ResumeState(TypedDict):
|
||||
log: Annotated[list, DeltaChannel(_append)]
|
||||
plain: Annotated[list, operator.add]
|
||||
|
||||
|
||||
def _both(marker: str) -> dict:
|
||||
return {"log": [marker], "plain": [marker]}
|
||||
|
||||
|
||||
def _ask(marker: str) -> Any:
|
||||
def ask(state: _ResumeState) -> dict:
|
||||
interrupt("approve?")
|
||||
return _both(marker)
|
||||
|
||||
return ask
|
||||
|
||||
|
||||
@pytest.mark.parametrize("addressed", [False, True])
|
||||
def test_resume_after_a_parallel_interrupt_replays_in_live_order(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("done", lambda state: _both("done"))
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_node("after", lambda state: _both("after"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
builder.add_edge("ask", "after")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
head = graph.get_state(config).config
|
||||
|
||||
graph.invoke(
|
||||
Command(resume="yes"), head if addressed else config, durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"]
|
||||
assert sorted(state.values["log"]) == ["after", "ask", "done", "in"]
|
||||
|
||||
|
||||
def test_resume_interleaves_the_resumed_superstep_by_task_path(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("z_done", lambda state: _both("z"))
|
||||
builder.add_node("a_asks", _ask("a"))
|
||||
builder.add_edge(START, "z_done")
|
||||
builder.add_edge(START, "a_asks")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
|
||||
graph.invoke(Command(resume="yes"), config, durability=durability)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "a", "z"]
|
||||
|
||||
|
||||
class _TaskIdOrderSaver(InMemorySaver):
|
||||
"""Replays each checkpoint's writes by task id, as savers without task path
|
||||
ordering do."""
|
||||
|
||||
def get_tuple(self, config: Any) -> Any:
|
||||
tup = super().get_tuple(config)
|
||||
if tup and tup.pending_writes:
|
||||
tup = tup._replace(pending_writes=sorted(tup.pending_writes))
|
||||
return tup
|
||||
|
||||
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
|
||||
|
||||
|
||||
def test_exit_run_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("a", lambda state: _both("a"))
|
||||
builder.add_node("b", lambda state: _both("b"))
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
graph = builder.compile(checkpointer=_TaskIdOrderSaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
graph.invoke(_both("in"), config, durability="exit")
|
||||
|
||||
assert graph.get_state(config).values["log"] == ["in", "a", "b"]
|
||||
|
||||
|
||||
def test_exit_resume_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_node("after", lambda state: _both("after"))
|
||||
builder.add_edge(START, "ask")
|
||||
builder.add_edge("ask", "after")
|
||||
graph = builder.compile(checkpointer=_TaskIdOrderSaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability="exit")
|
||||
|
||||
graph.invoke(Command(resume="yes"), config, durability="exit")
|
||||
|
||||
assert graph.get_state(config).values["log"] == ["in", "ask", "after"]
|
||||
|
||||
Reference in New Issue
Block a user