fix(checkpoint-sqlite): read a read-only database that predates task_path

setup() now runs an ALTER that a read-only database refuses, so one
created before the column could no longer be opened at all. Setup
treats that as "no column" and the delta query selects '' instead,
which is what those rows read back as anyway.

Also narrow the documented replay order to what it covers (task
writes; writes stored without a path sort first by task id), give the
conformance suite a valid UUID for its "sorts last" task id, and cover
Send fan-out in the parallel-order tests.
This commit is contained in:
Elior Nataf Lackritz
2026-09-29 12:40:49 -04:00
parent 35c3609b65
commit 02c4bc992b
7 changed files with 97 additions and 10 deletions
@@ -270,7 +270,7 @@ async def test_history_seed_ancestor_own_writes_are_replayed(
# Every uuid4 `build_delta_chain` tags its own writes with sorts between these
# two, so task_id order is fixed and always disagrees with task_path order.
TASK_ID_SORTS_FIRST = "00000000-0000-0000-0000-000000000000"
TASK_ID_SORTS_LAST = "zzzzzzzz-0000-0000-0000-000000000000"
TASK_ID_SORTS_LAST = "ffffffff-ffff-ffff-ffff-ffffffffffff"
async def test_history_orders_parallel_writes_by_task_path(
@@ -81,6 +81,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
conn: sqlite3.Connection
is_setup: bool
_has_task_path: bool = True
def __init__(
self,
@@ -170,7 +171,11 @@ class SqliteSaver(BaseCheckpointSaver[str]):
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
)
except sqlite3.OperationalError as e:
if "duplicate column name" not in str(e):
# A read-only database from before the column can still be read;
# its rows would all read back as '' anyway.
if "readonly database" in str(e):
self._has_task_path = False
elif "duplicate column name" not in str(e):
raise
self.is_setup = True
@@ -569,6 +574,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
stage2_sql = build_delta_stage2_sql(
has_task_path=self._has_task_path,
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
)
if stage2_sql:
@@ -39,7 +39,9 @@ DELTA_STAGE1_SQL = (
)
def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
def build_delta_stage2_sql(
*, chain_lens: Sequence[int], has_task_path: bool = True
) -> str:
"""Stage-2 per-channel UNION ALL fetching writes from `writes`.
One branch per channel with a non-empty chain. Each branch inlines its
@@ -53,11 +55,12 @@ def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
of a single `channel = ANY(channels)` filter when channels have
different chain depths — same rationale as postgres.
"""
task_path = "task_path" if has_task_path else "''"
branches: list[str] = []
for n in chain_lens:
cid_placeholders = ",".join("?" * n)
branches.append(
"SELECT checkpoint_id, channel, task_id, idx, type, value, task_path "
f"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})"
@@ -114,6 +114,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
lock: asyncio.Lock
is_setup: bool
_has_task_path: bool = True
def __init__(
self,
@@ -350,7 +351,11 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
)
await self.conn.commit()
except aiosqlite.OperationalError as e:
if "duplicate column name" not in str(e):
# A read-only database from before the column can still be read;
# its rows would all read back as '' anyway.
if "readonly database" in str(e):
self._has_task_path = False
elif "duplicate column name" not in str(e):
raise
self.is_setup = True
@@ -684,6 +689,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
stage2_sql = build_delta_stage2_sql(
has_task_path=self._has_task_path,
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
)
if stage2_sql:
@@ -85,3 +85,44 @@ async def test_async_setup_migrates_legacy_writes_table_repeatably(
("old-task", ""),
("task-1", "~__pregel_pull, node"),
]
def _legacy_database_with_history(db: Path) -> dict:
root = empty_checkpoint()
root["channel_values"] = {"ch": "seed"}
root["channel_versions"] = {"ch": 1}
with SqliteSaver.from_conn_string(str(db)) as saver:
root_config = saver.put(
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
root,
{},
{"ch": 1},
)
saver.put_writes(root_config, [("ch", "write")], "task", "~__pregel_pull, n")
child = saver.put(root_config, empty_checkpoint(), {}, {})
saver.conn.execute("ALTER TABLE writes DROP COLUMN task_path")
saver.conn.commit()
return child
def test_read_only_legacy_database_still_reads_delta_history(tmp_path: Path) -> None:
db = tmp_path / "legacy.sqlite"
child = _legacy_database_with_history(db)
saver = SqliteSaver(sqlite3.connect(f"file:{db}?mode=ro", uri=True))
got = saver.get_delta_channel_history(config=child, channels=["ch"])
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
async def test_async_read_only_legacy_database_still_reads_delta_history(
tmp_path: Path,
) -> None:
db = tmp_path / "legacy.sqlite"
child = _legacy_database_with_history(db)
async with aiosqlite.connect(f"file:{db}?mode=ro", uri=True) as conn:
saver = AsyncSqliteSaver(conn)
got = await saver.aget_delta_channel_history(config=child, channels=["ch"])
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
@@ -164,11 +164,12 @@ class DeltaChannelHistory(TypedDict):
next super-step and are excluded.
Within a single checkpoint, writes are ordered by
`(task_path, task_id, idx)`: the order `apply_writes` applied them in
live. `task_id` is a hash of the path, so ordering by it permutes
parallel tasks writing one channel, and reducers need not be
order-invariant. Writes stored without a `task_path` (graph input, or
rows predating the column) sort first.
`(task_path, task_id, idx)`, which is the order live execution applies
a super-step's task writes in. `task_id` is a hash of the path, so
ordering by it permutes parallel tasks writing one channel, and
reducers need not be order-invariant. Writes stored without a
`task_path` (graph input, `update_state` updates, exit-durability runs,
rows predating the column) sort first, by `task_id`.
* `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
@@ -8,11 +8,13 @@ from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, START, StateGraph
from langgraph.types import Send
pytestmark = pytest.mark.anyio
# Sorted, because live execution applies PULL tasks in node-name order.
FAN_OUT_NAMES = ["a", "b", "c", "d", "e", "f", "g", "h"]
SEND_ARGS = [f"send-{i:02d}" for i in range(12)]
def _append_reducer(current: list, updates: list) -> list:
@@ -36,6 +38,34 @@ def _build_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
return builder.compile(checkpointer=checkpointer)
def _build_send_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
class State(TypedDict):
items: Annotated[
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
]
builder = StateGraph(State)
builder.add_node("worker", lambda arg: {"items": [arg]})
builder.add_conditional_edges(
START, lambda state: [Send("worker", n) for n in SEND_ARGS]
)
builder.add_edge("worker", END)
return builder.compile(checkpointer=checkpointer)
async def test_get_state_matches_live_send_order(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_send_fan_out_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 == SEND_ARGS
assert replayed == live
async def test_get_state_matches_live_invoke_order(
async_checkpointer: BaseCheckpointSaver,
) -> None: