mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 21:15:11 +02:00
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:
+1
-1
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user