Compare commits

...
Author SHA1 Message Date
Elior Nataf LackritzandGitHub 1431900f7e Merge branch 'main' into fix/delta-replay-task-path-order 2026-08-07 09:50:07 -04:00
e095d9c41a fix: order delta channel replay by task path
DeltaChannel reconstructs its value by replaying ancestor writes through
the reducer. Every saver ordered a checkpoint's writes by (task_id, idx),
but live execution applies them in task-path order: apply_writes sorts a
super-step's tasks by task_path_str(task.path[:3]) before calling
channel.update. task_id is a hash of the path, so the two orders are
unrelated, and two or more tasks writing one DeltaChannel in a single
super-step replayed in an arbitrary permutation.

Reducers are only required to be batching-invariant, not order-invariant,
so the permutation changes the value: get_state disagreed with what invoke
returned, and continuing the thread persisted the reordered replay as the
base for later writes.

Replay now orders by (task_path, task_id, idx), following the precedent
already set for the Send channel by SELECT_PENDING_SENDS_SQL. InMemorySaver
and the postgres savers already persisted task_path and only needed the
sort key; sqlite accepted task_path on put_writes and dropped it, so the
writes table gains the column plus a probe-guarded ALTER for databases
created by earlier versions.

Writes stored without a task_path sort first within their checkpoint, which
is where live execution applies the task-less input writes that carry "".

The graph-level regression tests run against the full async_checkpointer
matrix (memory, sqlite, postgres in three pool modes) because each saver
reconstructs delta channels through its own override; 20 of the 25 cases
fail on main, and the 5 that pass are the sequential control.

Co-authored-by: ErenAta16 <149434812+ErenAta16@users.noreply.github.com>
Co-authored-by: ragnarok268 <58264829+ragnarok268@users.noreply.github.com>
2026-08-05 20:44:07 -04:00
10 changed files with 651 additions and 36 deletions
@@ -208,6 +208,93 @@ async def test_history_migration_plain_value_as_seed(
assert values == [2], f"Expected [2], got {values}"
# Task ids used by the ordering tests below. `build_delta_chain` tags its own
# writes with a `uuid4`, whose hex digits are all <= "f", so "aaaa..." sorts
# before every fixture task id and "zzzz..." sorts after every one of them.
# That makes the expected order fully determined rather than dependent on which
# uuid4 the fixture happened to draw.
TASK_ID_SORTS_FIRST = "aaaaaaaa-0000-0000-0000-000000000000"
TASK_ID_SORTS_LAST = "zzzzzzzz-0000-0000-0000-000000000000"
async def test_history_orders_parallel_writes_by_task_path(
saver: BaseCheckpointSaver,
) -> None:
"""Writes from several tasks in one super-step replay in task_path order.
Live execution sorts a super-step's tasks by `task_path_str(path[:3])`
before applying their values, so replay has to recover that order rather
than `task_id` order — `task_id` is a hash of the path, so the two
disagree, and reducers are only required to be batching-invariant, not
order-invariant.
The two task_ids are assigned so they sort in the *opposite* order from
their task_paths. A saver ordering by `(task_id, idx)` therefore returns
these writes reversed, rather than passing by happening to agree.
"""
configs = await build_delta_chain(
saver,
thread_id=str(uuid4()),
channel="ch",
snapshots_at_steps=[0],
total_steps=3,
)
# The chain is: step 0 snapshot (seed), step 1 write, step 2 write.
# `aget_delta_channel_history` walks from the head's parent back to the
# seed, so it collects step 1's writes only — step 0 terminates the walk
# and step 2 is the head, whose own writes are pending for the next
# super-step and excluded. So step 1 is where these writes have to go.
step_1, head = configs[1], configs[2]
await saver.aput_writes(
step_1, [("ch", "second")], TASK_ID_SORTS_FIRST, "~pull, 02"
)
await saver.aput_writes(step_1, [("ch", "first")], TASK_ID_SORTS_LAST, "~pull, 01")
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
values = [w[2] for w in result["ch"]["writes"]]
# 1 is the fixture's own write at step 1. It carries no task_path, so it
# sorts ahead of both writes added above.
assert values == [1, "first", "second"], (
f"Expected task_path order [1, 'first', 'second'], got {values}. "
"Ordering by (task_id, idx) alone yields [1, 'second', 'first']."
)
async def test_history_orders_pathless_writes_first(
saver: BaseCheckpointSaver,
) -> None:
"""Writes stored without a task_path sort ahead of path-carrying ones.
A task-less write (graph input) persists `task_path=""`, as does any row
written before a saver recorded the column. `""` precedes every
`task_path_str` output because that function prefixes tuples with `~`, so
those writes replay first — where live execution applies graph input.
"""
configs = await build_delta_chain(
saver,
thread_id=str(uuid4()),
channel="ch",
snapshots_at_steps=[0],
total_steps=3,
)
# Same chain shape as above: step 1 is the only step the walk collects.
step_1, head = configs[1], configs[2]
# Committed in the opposite order to the one they must replay in, so the
# assertion cannot pass on insertion order alone.
await saver.aput_writes(
step_1, [("ch", "from_node")], TASK_ID_SORTS_FIRST, "~pull, a"
)
await saver.aput_writes(step_1, [("ch", "from_input")], TASK_ID_SORTS_LAST)
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
values = [w[2] for w in result["ch"]["writes"]]
# Both 1 (the fixture's write) and "from_input" are pathless, so they sort
# by task_id among themselves and both precede the path-carrying write.
assert values == [1, "from_input", "from_node"], (
f"Expected pathless writes first, got {values}"
)
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
test_history_returns_writes_oldest_first,
test_history_seed_is_nearest_snapshot,
@@ -216,6 +303,8 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
test_history_empty_channels_returns_empty,
test_history_walk_to_root_no_seed,
test_history_migration_plain_value_as_seed,
test_history_orders_parallel_writes_by_task_path,
test_history_orders_pathless_writes_first,
]
@@ -168,6 +168,7 @@ class _DeltaStage2Row(TypedDict, total=False):
type: str | None
blob: bytes | None
task_id: str | None # "w" rows only
task_path: str | None # "w" rows only
idx: int | None # "w" rows only
version: str | None # "b" rows only
@@ -319,7 +320,7 @@ def _build_delta_stage2_sql(
branches.append(
"SELECT 'w'::text AS _kind, "
"checkpoint_id, channel, "
"type, blob, task_id, idx, NULL::text AS version "
"type, blob, task_id, task_path, idx, NULL::text AS version "
"FROM checkpoint_writes "
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
"AND checkpoint_id = ANY(%s)"
@@ -327,7 +328,8 @@ def _build_delta_stage2_sql(
for _ in channels_with_seed:
branches.append(
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
"type, blob, NULL::text AS task_id, NULL::int AS idx, version "
"type, blob, NULL::text AS task_id, NULL::text AS task_path, "
"NULL::int AS idx, version "
"FROM checkpoint_blobs "
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
"AND version = %s"
@@ -492,10 +494,11 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
stored value, or when the seed blob is sentinel "empty" — in both cases
the consumer treats absence as "start empty".
"""
# writes_by_ch_by_cid[channel][cid] = list of (type, blob, task_id, idx)
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
ch: {} for ch in channels
}
# writes_by_ch_by_cid[channel][cid] = list of
# (type, blob, task_id, idx, task_path)
writes_by_ch_by_cid: dict[
str, dict[str, list[tuple[str, bytes, str, int, str]]]
] = {ch: {} for ch in channels}
# seed_blob_by_ver[(channel, version)] = (type, blob)
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
@@ -506,8 +509,17 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
cid = cast(str, r["checkpoint_id"])
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
cast(
"tuple[str, bytes, str, int]",
(r["type"], r["blob"], r["task_id"], r["idx"]),
"tuple[str, bytes, str, int, str]",
(
r["type"],
r["blob"],
r["task_id"],
r["idx"],
# `task_path` is NOT NULL DEFAULT '' on "w" rows;
# it is nullable on `_DeltaStage2Row` only because
# the seed branch selects NULL for it.
r["task_path"],
),
)
)
else: # kind == "b"
@@ -516,10 +528,12 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
"tuple[str, bytes]", (r["type"], r["blob"])
)
# Sort writes per (channel, cid) newest-first by (task_id, idx)
# Sort writes per (channel, cid) newest-first by
# (task_path, task_id, idx) — the order `apply_writes` applied them
# in live, and the order documented on `DeltaChannelHistory`.
for cid_map in writes_by_ch_by_cid.values():
for ws in cid_map.values():
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
ws.sort(key=lambda w: (w[4], w[2], w[3]), reverse=True)
result: dict[str, DeltaChannelHistory] = {}
for ch in channels:
@@ -529,7 +543,9 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
collected: list[PendingWrite] = []
cid_writes = writes_by_ch_by_cid.get(ch, {})
for cid in chain_cids:
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, []):
for type_tag, write_blob, task_id, _idx, _path in cid_writes.get(
cid, []
):
val = self.serde.loads_typed((type_tag, write_blob))
collected.append((task_id, ch, val))
collected.reverse()
@@ -29,6 +29,11 @@ from langgraph.checkpoint.sqlite._delta import (
build_delta_stage2_sql,
step_walk_with_row,
)
from langgraph.checkpoint.sqlite._schema import (
ADD_WRITES_TASK_PATH_SQL,
DUPLICATE_COLUMN_ERROR,
HAS_WRITES_TASK_PATH_SQL,
)
from langgraph.checkpoint.sqlite.utils import search_where
_AIO_ERROR_MSG = (
@@ -154,6 +159,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
checkpoint_ns TEXT NOT NULL DEFAULT '',
checkpoint_id TEXT NOT NULL,
task_id TEXT NOT NULL,
task_path TEXT NOT NULL DEFAULT '',
idx INTEGER NOT NULL,
channel TEXT NOT NULL,
type TEXT,
@@ -162,6 +168,12 @@ class SqliteSaver(BaseCheckpointSaver[str]):
);
"""
)
if not self.conn.execute(HAS_WRITES_TASK_PATH_SQL).fetchone():
try:
self.conn.execute(ADD_WRITES_TASK_PATH_SQL)
except sqlite3.OperationalError as exc:
if DUPLICATE_COLUMN_ERROR not in str(exc):
raise
self.is_setup = True
@@ -460,9 +472,9 @@ class SqliteSaver(BaseCheckpointSaver[str]):
task_path: Path of the task creating the writes.
"""
query = (
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
if all(w[0] in WRITES_IDX_MAP for w in writes)
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
)
with self.cursor() as cur:
cur.executemany(
@@ -473,6 +485,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
str(config["configurable"]["checkpoint_ns"]),
str(config["configurable"]["checkpoint_id"]),
task_id,
task_path,
WRITES_IDX_MAP.get(channel, idx),
channel,
*self.serde.dumps_typed(value),
@@ -568,7 +581,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
)
cur.execute(stage2_sql, stage2_params)
stage2_rows = cast(
"list[tuple[str, str, str, int, str, bytes]]", cur.fetchall()
"list[tuple[str, str, str, int, str, bytes, str]]", cur.fetchall()
)
else:
stage2_rows = []
@@ -57,7 +57,7 @@ def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
for n in chain_lens:
cid_placeholders = ",".join("?" * n)
branches.append(
"SELECT checkpoint_id, channel, task_id, idx, type, value "
"SELECT checkpoint_id, channel, task_id, idx, type, value, task_path "
"FROM writes "
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
f"AND checkpoint_id IN ({cid_placeholders})"
@@ -130,29 +130,33 @@ def build_delta_channels_writes_history(
chain_by_ch: Mapping[str, list[str]],
seed_val_by_ch: Mapping[str, Any],
seeded: set[str],
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes]],
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes, str]],
serde: Any,
) -> dict[str, DeltaChannelHistory]:
"""Demux stage-2 rows per channel; produce per-channel histories.
Stage-2 rows are `(checkpoint_id, channel, task_id, idx, type, value)`.
Final write order is oldest→newest globally and `(task_id, idx)` within
a checkpoint, matching the contract on `DeltaChannelHistory.writes`.
Stage-2 rows are
`(checkpoint_id, channel, task_id, idx, type, value, task_path)`.
Final write order is oldest→newest globally and
`(task_path, task_id, idx)` within a checkpoint, matching the contract
on `DeltaChannelHistory.writes` — that is the order `apply_writes`
applied them in live, which `(task_id, idx)` alone does not recover
for parallel tasks writing one channel in a single super-step.
`seed` is omitted when the walk reached a true root with no snapshot
found (channel never entered `seeded`); consumers treat absence as
"start empty".
"""
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
ch: {} for ch in channels
}
for cid, ch, task_id, idx, type_tag, value_blob in stage2_rows:
writes_by_ch_by_cid: dict[
str, dict[str, list[tuple[str, bytes, str, int, str]]]
] = {ch: {} for ch in channels}
for cid, ch, task_id, idx, type_tag, value_blob, task_path in stage2_rows:
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
(type_tag, value_blob, task_id, idx)
(type_tag, value_blob, task_id, idx, task_path)
)
for cid_map in writes_by_ch_by_cid.values():
for ws in cid_map.values():
ws.sort(key=lambda w: (w[2], w[3]))
ws.sort(key=lambda w: (w[4], w[2], w[3]))
result: dict[str, DeltaChannelHistory] = {}
for ch in channels:
@@ -161,7 +165,7 @@ def build_delta_channels_writes_history(
collected: list[PendingWrite] = []
# Chain is newest-first; iterate oldest-first for the public order.
for cid in reversed(chain_cids):
for type_tag, value_blob, task_id, _idx in cid_writes.get(cid, []):
for type_tag, value_blob, task_id, _idx, _path in cid_writes.get(cid, []):
collected.append(
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
)
@@ -0,0 +1,37 @@
"""Additive schema migrations shared by the sqlite savers.
`SqliteSaver.setup` and `AsyncSqliteSaver.setup` create their tables with
`CREATE TABLE IF NOT EXISTS`, which leaves a database created by an earlier
version on the earlier schema. Sqlite has no `ADD COLUMN IF NOT EXISTS`
(the postgres savers rely on that form), and re-running a plain
`ALTER TABLE ... ADD COLUMN` raises `OperationalError: duplicate column
name`. So each migration pairs an `ALTER` with a probe against
`pragma_table_info` that tells us whether this database still needs it.
Databases created fresh already carry every column from the `CREATE TABLE`
statements, so the probe finds the column and the `ALTER` never runs.
"""
from __future__ import annotations
# `writes.task_path` records the path of the task that produced a write.
# Delta channel replay orders a checkpoint's writes by
# (task_path, task_id, idx) to reproduce the order `apply_writes` applied
# them in live; without the column, replay can only order by
# (task_id, idx), which permutes writes made by parallel tasks in the same
# super-step. Rows written before this migration keep the `''` default and
# so sort ahead of path-carrying rows within their checkpoint.
HAS_WRITES_TASK_PATH_SQL = (
"SELECT 1 FROM pragma_table_info('writes') WHERE name = 'task_path'"
)
ADD_WRITES_TASK_PATH_SQL = (
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
)
# Substring of the `OperationalError` sqlite raises when the column is already
# there. The probe above is not enough on its own: two connections opening the
# same file can both pass it and both issue the `ALTER`, and unlike
# `CREATE TABLE IF NOT EXISTS` the loser of that race raises. Callers treat it
# as success — whoever won did the same migration.
DUPLICATE_COLUMN_ERROR = "duplicate column name"
@@ -30,6 +30,11 @@ from langgraph.checkpoint.sqlite._delta import (
build_delta_stage2_sql,
step_walk_with_row,
)
from langgraph.checkpoint.sqlite._schema import (
ADD_WRITES_TASK_PATH_SQL,
DUPLICATE_COLUMN_ERROR,
HAS_WRITES_TASK_PATH_SQL,
)
from langgraph.checkpoint.sqlite.utils import search_where
T = TypeVar("T", bound=Callable)
@@ -331,6 +336,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
checkpoint_ns TEXT NOT NULL DEFAULT '',
checkpoint_id TEXT NOT NULL,
task_id TEXT NOT NULL,
task_path TEXT NOT NULL DEFAULT '',
idx INTEGER NOT NULL,
channel TEXT NOT NULL,
type TEXT,
@@ -341,6 +347,16 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
):
await self.conn.commit()
async with self.conn.execute(HAS_WRITES_TASK_PATH_SQL) as cur:
has_task_path = await cur.fetchone() is not None
if not has_task_path:
try:
await self.conn.execute(ADD_WRITES_TASK_PATH_SQL)
except aiosqlite.OperationalError as exc:
if DUPLICATE_COLUMN_ERROR not in str(exc):
raise
await self.conn.commit()
self.is_setup = True
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
@@ -576,9 +592,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
task_path: Path of the task creating the writes.
"""
query = (
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
if all(w[0] in WRITES_IDX_MAP for w in writes)
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
)
await self.setup()
async with self.lock, self.conn.cursor() as cur:
@@ -590,6 +606,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
str(config["configurable"]["checkpoint_ns"]),
str(config["configurable"]["checkpoint_id"]),
task_id,
task_path,
WRITES_IDX_MAP.get(channel, idx),
channel,
*self.serde.dumps_typed(value),
@@ -681,7 +698,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
)
await cur.execute(stage2_sql, stage2_params)
stage2_rows = cast(
"list[tuple[str, str, str, int, str, bytes]]",
"list[tuple[str, str, str, int, str, bytes, str]]",
await cur.fetchall(),
)
else:
@@ -0,0 +1,214 @@
"""Tests for the additive `writes.task_path` migration (#8382).
`task_path` records the path of the task that produced a write, so delta
channel replay can restore the order `apply_writes` applied a super-step's
writes in. The sqlite savers previously accepted `task_path` on `put_writes`
and dropped it, so the column has to be added to databases created by earlier
versions as well as to fresh ones.
Sqlite has no `ADD COLUMN IF NOT EXISTS`, so `setup()` probes
`pragma_table_info` before issuing the `ALTER` — these tests pin that the
probe makes the migration both effective and repeatable.
"""
from __future__ import annotations
import sqlite3
from pathlib import Path
import aiosqlite
import pytest
from langgraph.checkpoint.base import empty_checkpoint
from langgraph.checkpoint.sqlite import SqliteSaver
from langgraph.checkpoint.sqlite._schema import ADD_WRITES_TASK_PATH_SQL
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
# The `writes` table as created before `task_path` existed.
LEGACY_SCHEMA = """
CREATE TABLE checkpoints (
thread_id TEXT NOT NULL,
checkpoint_ns TEXT NOT NULL DEFAULT '',
checkpoint_id TEXT NOT NULL,
parent_checkpoint_id TEXT,
type TEXT,
checkpoint BLOB,
metadata BLOB,
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id)
);
CREATE TABLE writes (
thread_id TEXT NOT NULL,
checkpoint_ns TEXT NOT NULL DEFAULT '',
checkpoint_id TEXT NOT NULL,
task_id TEXT NOT NULL,
idx INTEGER NOT NULL,
channel TEXT NOT NULL,
type TEXT,
value BLOB,
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
);
"""
def _write_legacy_db(path: Path) -> None:
conn = sqlite3.connect(path)
try:
conn.executescript(LEGACY_SCHEMA)
conn.commit()
finally:
conn.close()
def _columns(conn: sqlite3.Connection, table: str) -> list[str]:
return [row[1] for row in conn.execute(f"PRAGMA table_info({table})")]
def test_fresh_database_has_task_path(tmp_path: Path) -> None:
with SqliteSaver.from_conn_string(str(tmp_path / "fresh.sqlite")) as saver:
saver.setup()
assert "task_path" in _columns(saver.conn, "writes")
def test_legacy_database_gains_task_path(tmp_path: Path) -> None:
db = tmp_path / "legacy.sqlite"
_write_legacy_db(db)
with SqliteSaver.from_conn_string(str(db)) as saver:
saver.setup()
columns = _columns(saver.conn, "writes")
assert "task_path" in columns
# Existing columns are untouched — this is additive, not a table rebuild.
assert columns[:8] == [
"thread_id",
"checkpoint_ns",
"checkpoint_id",
"task_id",
"idx",
"channel",
"type",
"value",
]
def test_setup_is_repeatable_on_migrated_database(tmp_path: Path) -> None:
"""A second `setup()` must not re-issue the `ALTER`.
Sqlite raises `duplicate column name` rather than ignoring it, so an
unguarded `ALTER` would break every reopen of a migrated database.
"""
db = tmp_path / "legacy.sqlite"
_write_legacy_db(db)
with SqliteSaver.from_conn_string(str(db)) as saver:
saver.setup()
saver.is_setup = False
saver.setup()
assert "task_path" in _columns(saver.conn, "writes")
# And again through a fresh connection to the migrated file.
with SqliteSaver.from_conn_string(str(db)) as saver:
saver.setup()
assert "task_path" in _columns(saver.conn, "writes")
def test_legacy_rows_keep_default_and_sort_first(tmp_path: Path) -> None:
"""Rows predating the column read back as `''` and order ahead of paths.
`''` precedes every `task_path_str` output, which puts pre-migration
writes before path-carrying ones within their checkpoint instead of
interleaving them under a rule that never applied to them.
"""
db = tmp_path / "legacy.sqlite"
_write_legacy_db(db)
conn = sqlite3.connect(db)
try:
conn.execute(
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id,"
" idx, channel, type, value) VALUES ('t', '', 'c', 'task', 0, 'ch',"
" 'null', X'')"
)
conn.commit()
finally:
conn.close()
with SqliteSaver.from_conn_string(str(db)) as saver:
saver.setup()
stored = saver.conn.execute("SELECT task_path FROM writes").fetchall()
assert stored == [("",)]
ordered = saver.conn.execute(
"SELECT task_path FROM writes ORDER BY task_path, task_id, idx"
).fetchall()
assert ordered[0] == ("",)
def test_setup_survives_losing_the_migration_race(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
"""`setup()` succeeds when another connection migrates first.
The `pragma_table_info` probe is not a lock. Two connections opening the
same file can both see the column missing, and whichever issues the `ALTER`
second gets `duplicate column name` — `ALTER TABLE ADD COLUMN` has no
`IF NOT EXISTS` form to fall back on, unlike the `CREATE TABLE`s above it.
Stubbing the probe to always report the column missing reproduces exactly
the losing interleaving (probe says absent, another connection adds it,
then we `ALTER`) without depending on thread timing.
"""
db = tmp_path / "legacy.sqlite"
_write_legacy_db(db)
# Winner of the race: migrates the file out from under the saver below.
winner = sqlite3.connect(db)
try:
winner.execute(ADD_WRITES_TASK_PATH_SQL)
winner.commit()
finally:
winner.close()
monkeypatch.setattr(
"langgraph.checkpoint.sqlite.HAS_WRITES_TASK_PATH_SQL", "SELECT 1 WHERE 0"
)
with SqliteSaver.from_conn_string(str(db)) as loser:
loser.setup()
assert "task_path" in _columns(loser.conn, "writes")
def test_put_writes_persists_task_path(tmp_path: Path) -> None:
with SqliteSaver.from_conn_string(str(tmp_path / "fresh.sqlite")) as saver:
config = saver.put(
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
empty_checkpoint(),
{},
{},
)
saver.put_writes(config, [("ch", "v")], "task-1", "~__pregel_pull, node")
stored = saver.conn.execute("SELECT task_id, task_path FROM writes").fetchall()
assert stored == [("task-1", "~__pregel_pull, node")]
@pytest.mark.asyncio
async def test_async_saver_migrates_and_persists_task_path(tmp_path: Path) -> None:
db = tmp_path / "legacy.sqlite"
_write_legacy_db(db)
async with AsyncSqliteSaver.from_conn_string(str(db)) as saver:
await saver.setup()
config = await saver.aput(
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
empty_checkpoint(),
{},
{},
)
await saver.aput_writes(config, [("ch", "v")], "task-1", "~__pregel_pull, node")
# Idempotent for the async saver too.
saver.is_setup = False
await saver.setup()
async with aiosqlite.connect(db) as conn:
async with conn.execute("SELECT task_id, task_path FROM writes") as cur:
assert await cur.fetchall() == [("task-1", "~__pregel_pull, node")]
@@ -161,6 +161,23 @@ class DeltaChannelHistory(TypedDict):
Always present; possibly empty. Already filtered to one channel.
Writes stored at the target checkpoint itself are pending for the
next super-step and are excluded.
Within a single checkpoint, writes are ordered by
`(task_path, task_id, idx)`. This mirrors the order live execution
applied them in: `apply_writes` sorts a super-step's tasks by
`task_path_str(task.path[:3])` before handing their values to
`channel.update`, so a path-ordered replay reproduces the value
`invoke` returned. Ordering by `(task_id, idx)` alone does not —
`task_id` is a hash of the path, so for two or more tasks writing
the same channel in one super-step it permutes the values against
the order the reducer originally saw them in. Reducers are only
required to be batching-invariant, not order-invariant, so that
permutation changes the reconstructed value.
Writes persisted without a `task_path` (a task-less write such as
graph input, or a row written before the saver recorded the column)
sort first within their checkpoint, since `""` precedes every
`task_path_str` output — `task_path_str` prefixes tuples with `~`.
* `seed` — the stored value at the nearest ancestor whose
`channel_values[ch]` is populated. Omitted if the walk reached the
root without finding any stored value (consumer treats absence as
@@ -610,6 +627,14 @@ class BaseCheckpointSaver(Generic[V]):
`PostgresSaver`) override for performance; the return contract is
fixed here.
`PendingWrite` carries no `task_path`, so the default takes each
ancestor's write order straight from `get_tuple`. Savers relying
on it must therefore return `pending_writes` ordered by
`(task_path, task_id, idx)` to satisfy the intra-checkpoint order
documented on `DeltaChannelHistory`; savers that order
`pending_writes` by `(task_id, idx)` alone need to override this
method (as the in-tree savers do) rather than inherit it.
Args:
config: Configuration identifying the target checkpoint.
channels: Channel names to walk for. Empty → empty mapping.
@@ -30,6 +30,23 @@ from langgraph.checkpoint.base import (
logger = logging.getLogger(__name__)
# How `InMemorySaver.writes[thread, ns, checkpoint]` keys and stores one write.
_WriteKey = tuple[str, int] # task ID, write idx
_WriteValue = tuple[str, str, tuple[str, bytes], str] # + channel, value, path
_WriteEntry = tuple[_WriteKey, _WriteValue] # one `dict.items()` pair
def _delta_replay_sort_key(entry: _WriteEntry) -> tuple[str, str, int]:
"""Order one checkpoint's writes as `apply_writes` applied them live.
Live order is `(task_path, task_id, idx)` — see `DeltaChannelHistory`. It
has to be assembled from both halves of the entry: `(task_id, idx)` is the
key, `task_path` is the last element of the value.
"""
(task_id, idx), (_, _, _, task_path) = entry
return (task_path, task_id, idx)
class InMemorySaver(
BaseCheckpointSaver[str], AbstractContextManager, AbstractAsyncContextManager
):
@@ -71,10 +88,7 @@ class InMemorySaver(
dict[str, dict[str, tuple[tuple[str, bytes], tuple[str, bytes], str | None]]],
]
# (thread ID, checkpoint NS, checkpoint ID) -> (task ID, write idx)
writes: defaultdict[
tuple[str, str, str],
dict[tuple[str, int], tuple[str, str, tuple[str, bytes], str]],
]
writes: defaultdict[tuple[str, str, str], dict[_WriteKey, _WriteValue]]
blobs: dict[
tuple[
str, str, str, str | int | float
@@ -200,8 +214,10 @@ class InMemorySaver(
terminated_here.add(ch)
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
for (_task_id, _idx), (tid, ch, serialized, _) in sorted(
step_writes.items(), reverse=True
# Newest-first; the caller reverses to get the public oldest-first
# order.
for (task_id, _idx), (_, ch, serialized, _task_path) in sorted(
step_writes.items(), key=_delta_replay_sort_key, reverse=True
):
if ch not in remaining:
continue
@@ -211,7 +227,7 @@ class InMemorySaver(
):
continue
collected_by_ch[ch].append(
(tid, ch, self.serde.loads_typed(serialized))
(task_id, ch, self.serde.loads_typed(serialized))
)
for ch in terminated_here:
@@ -0,0 +1,184 @@
"""Tests that `DeltaChannel` replay preserves live parallel-write order.
Regression suite for #8382.
`apply_writes` sorts a super-step's tasks by `task_path_str(task.path[:3])`
before handing their values to `channel.update`, so the order a reducer sees is
deterministic and independent of which parallel task finishes first. Replay has
to recover that same order. Ordering a checkpoint's writes by `(task_id, idx)`
does not: `task_id` is a hash of the path, so for two or more tasks writing one
`DeltaChannel` in a single super-step it yields an effectively arbitrary
permutation. Reducers are required to be batching-invariant, not
order-invariant, so the permutation changes the reconstructed value.
Every test runs against the full `async_checkpointer` matrix — memory, sqlite,
and postgres in three pool modes — because each saver reconstructs delta
channels through its own `aget_delta_channel_history` override rather than a
shared code path, and the three stored `task_path` differently before this fix.
"""
from itertools import pairwise
from typing import Annotated, Any
import pytest
from langgraph.checkpoint.base import BaseCheckpointSaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, START, StateGraph
pytestmark = pytest.mark.anyio
# Node names double as the values written. They are listed in sorted order,
# which is also the order live execution applies them: each node is a PULL task
# whose path is `("__pregel_pull", name)`, so sorting paths sorts by name.
FAN_OUT_NAMES = ["a", "b", "c", "d", "e", "f", "g", "h"]
def _append_reducer(current: list, updates: list) -> list:
"""Order-sensitive list accumulation, as in the `DeltaChannel` docstring."""
result = list(current)
for update in updates:
if isinstance(update, list):
result.extend(update)
else:
result.append(update)
return result
def _build_graph(checkpointer: BaseCheckpointSaver, *, sequential: bool = False) -> Any:
"""Compile a `DeltaChannel`-backed `items` graph over `FAN_OUT_NAMES`.
By default every node is wired off `START`, so they all write `items` in one
super-step — the shape #8382 is about. `sequential=True` chains them
instead, giving one writer per super-step as a control.
`snapshot_frequency` is far above the number of updates these tests make, so
no snapshot is ever written and the value has to come from replaying
ancestor writes — the path under test.
"""
class State(TypedDict):
items: Annotated[
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
]
def make_node(label: str) -> Any:
def node(state: State) -> dict:
return {"items": [label]}
return node
builder = StateGraph(State)
for name in FAN_OUT_NAMES:
builder.add_node(name, make_node(name))
if sequential:
for source, target in pairwise([START, *FAN_OUT_NAMES, END]):
builder.add_edge(source, target)
else:
for name in FAN_OUT_NAMES:
builder.add_edge(START, name)
builder.add_edge(name, END)
return builder.compile(checkpointer=checkpointer)
async def test_get_state_matches_live_invoke_order(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""A cold read reports the same order `invoke` returned."""
graph = _build_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
live = (await graph.ainvoke({"items": []}, config))["items"]
replayed = (await graph.aget_state(config)).values["items"]
assert live == FAN_OUT_NAMES
assert replayed == live
async def test_continuing_thread_preserves_committed_prefix(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""A second run appends without reordering the first run's items.
The more serious half of #8382: the reordered replay becomes the base that
later writes build on, so the corruption is persisted rather than confined
to a read.
"""
graph = _build_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
first = (await graph.ainvoke({"items": []}, config))["items"]
second = (await graph.ainvoke({"items": []}, config))["items"]
assert second == first + first
assert (await graph.aget_state(config)).values["items"] == second
async def test_order_stable_across_many_supersteps(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Order holds over a chain spanning several supersteps with no snapshot."""
runs = 5
graph = _build_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
for _ in range(runs):
live = (await graph.ainvoke({"items": []}, config))["items"]
assert live == FAN_OUT_NAMES * runs
assert (await graph.aget_state(config)).values["items"] == live
async def test_state_history_reports_live_order_at_every_step(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Every checkpoint in the history replays in live order.
Guards the walk at intermediate depths, not just from the head. Entries are
checked against the order live execution produced rather than against the
replayed head — comparing replayed values only to each other passes even
when every one of them is permuted the same wrong way.
"""
runs = 3
graph = _build_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
for _ in range(runs):
await graph.ainvoke({"items": []}, config)
live_expected = FAN_OUT_NAMES * runs
seen = [
snapshot.values["items"]
async for snapshot in graph.aget_state_history(config)
if "items" in snapshot.values
]
assert seen, "expected at least one snapshot carrying `items`"
# The deepest entry is the head, so the matrix below covers the full value
# as well as every partial prefix.
assert max(len(values) for values in seen) == len(live_expected)
for values in seen:
assert values == live_expected[: len(values)], (
f"history entry {values} is not the live order "
f"{live_expected[: len(values)]}"
)
async def test_sequential_graph_unaffected(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""One writer per super-step replays correctly with or without the fix.
Control: it localises #8382 to multiple tasks writing one channel in a
single super-step, rather than to delta replay in general. This is the one
test here that passes on main.
"""
graph = _build_graph(async_checkpointer, sequential=True)
config = {"configurable": {"thread_id": "1"}}
live = (await graph.ainvoke({"items": []}, config))["items"]
replayed = (await graph.aget_state(config)).values["items"]
assert live == FAN_OUT_NAMES
assert replayed == live