mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-25 11:05:12 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cdd02084e9 |
+57
@@ -267,6 +267,61 @@ 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"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_history_orders_parallel_writes_by_task_path(
|
||||||
|
saver: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
"""Writes from parallel tasks replay in task_path order, not task_id order."""
|
||||||
|
configs = await build_delta_chain(
|
||||||
|
saver,
|
||||||
|
thread_id=str(uuid4()),
|
||||||
|
channel="ch",
|
||||||
|
snapshots_at_steps=[0],
|
||||||
|
total_steps=3,
|
||||||
|
)
|
||||||
|
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"]]
|
||||||
|
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 (graph input) replay before task writes."""
|
||||||
|
configs = await build_delta_chain(
|
||||||
|
saver,
|
||||||
|
thread_id=str(uuid4()),
|
||||||
|
channel="ch",
|
||||||
|
snapshots_at_steps=[0],
|
||||||
|
total_steps=3,
|
||||||
|
)
|
||||||
|
step_1, head = configs[1], configs[2]
|
||||||
|
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"]]
|
||||||
|
assert values == [1, "from_input", "from_node"], (
|
||||||
|
f"Expected pathless writes first, got {values}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
||||||
test_history_returns_writes_oldest_first,
|
test_history_returns_writes_oldest_first,
|
||||||
test_history_seed_is_nearest_snapshot,
|
test_history_seed_is_nearest_snapshot,
|
||||||
@@ -276,6 +331,8 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
|||||||
test_history_walk_to_root_no_seed,
|
test_history_walk_to_root_no_seed,
|
||||||
test_history_migration_plain_value_as_seed,
|
test_history_migration_plain_value_as_seed,
|
||||||
test_history_seed_ancestor_own_writes_are_replayed,
|
test_history_seed_ancestor_own_writes_are_replayed,
|
||||||
|
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
|
type: str | None
|
||||||
blob: bytes | None
|
blob: bytes | None
|
||||||
task_id: str | None # "w" rows only
|
task_id: str | None # "w" rows only
|
||||||
|
task_path: str | None # "w" rows only
|
||||||
idx: int | None # "w" rows only
|
idx: int | None # "w" rows only
|
||||||
version: str | None # "b" rows only
|
version: str | None # "b" rows only
|
||||||
|
|
||||||
@@ -319,7 +320,7 @@ def _build_delta_stage2_sql(
|
|||||||
branches.append(
|
branches.append(
|
||||||
"SELECT 'w'::text AS _kind, "
|
"SELECT 'w'::text AS _kind, "
|
||||||
"checkpoint_id, channel, "
|
"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 "
|
"FROM checkpoint_writes "
|
||||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||||
"AND checkpoint_id = ANY(%s)"
|
"AND checkpoint_id = ANY(%s)"
|
||||||
@@ -327,7 +328,8 @@ def _build_delta_stage2_sql(
|
|||||||
for _ in channels_with_seed:
|
for _ in channels_with_seed:
|
||||||
branches.append(
|
branches.append(
|
||||||
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
|
"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 "
|
"FROM checkpoint_blobs "
|
||||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||||
"AND version = %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
|
stored value, or when the seed blob is sentinel "empty" — in both cases
|
||||||
the consumer treats absence as "start empty".
|
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[channel][cid] = list of
|
||||||
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
# (type, blob, task_id, idx, task_path)
|
||||||
ch: {} for ch in channels
|
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[(channel, version)] = (type, blob)
|
||||||
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
||||||
|
|
||||||
@@ -506,8 +509,14 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
cid = cast(str, r["checkpoint_id"])
|
cid = cast(str, r["checkpoint_id"])
|
||||||
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
||||||
cast(
|
cast(
|
||||||
"tuple[str, bytes, str, int]",
|
"tuple[str, bytes, str, int, str]",
|
||||||
(r["type"], r["blob"], r["task_id"], r["idx"]),
|
(
|
||||||
|
r["type"],
|
||||||
|
r["blob"],
|
||||||
|
r["task_id"],
|
||||||
|
r["idx"],
|
||||||
|
r["task_path"],
|
||||||
|
),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
else: # kind == "b"
|
else: # kind == "b"
|
||||||
@@ -516,10 +525,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
"tuple[str, bytes]", (r["type"], r["blob"])
|
"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)
|
||||||
for cid_map in writes_by_ch_by_cid.values():
|
for cid_map in writes_by_ch_by_cid.values():
|
||||||
for ws in cid_map.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] = {}
|
result: dict[str, DeltaChannelHistory] = {}
|
||||||
for ch in channels:
|
for ch in channels:
|
||||||
@@ -529,7 +538,9 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
collected: list[PendingWrite] = []
|
collected: list[PendingWrite] = []
|
||||||
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
||||||
for cid in chain_cids:
|
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))
|
val = self.serde.loads_typed((type_tag, write_blob))
|
||||||
collected.append((task_id, ch, val))
|
collected.append((task_id, ch, val))
|
||||||
collected.reverse()
|
collected.reverse()
|
||||||
|
|||||||
@@ -154,6 +154,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||||
checkpoint_id TEXT NOT NULL,
|
checkpoint_id TEXT NOT NULL,
|
||||||
task_id TEXT NOT NULL,
|
task_id TEXT NOT NULL,
|
||||||
|
task_path TEXT NOT NULL DEFAULT '',
|
||||||
idx INTEGER NOT NULL,
|
idx INTEGER NOT NULL,
|
||||||
channel TEXT NOT NULL,
|
channel TEXT NOT NULL,
|
||||||
type TEXT,
|
type TEXT,
|
||||||
@@ -162,6 +163,15 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
);
|
);
|
||||||
"""
|
"""
|
||||||
)
|
)
|
||||||
|
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
|
||||||
|
# created before `task_path` existed and is a no-op on the rest.
|
||||||
|
try:
|
||||||
|
self.conn.execute(
|
||||||
|
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
|
||||||
|
)
|
||||||
|
except sqlite3.OperationalError as e:
|
||||||
|
if "duplicate column name" not in str(e):
|
||||||
|
raise
|
||||||
|
|
||||||
self.is_setup = True
|
self.is_setup = True
|
||||||
|
|
||||||
@@ -460,9 +470,9 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
task_path: Path of the task creating the writes.
|
task_path: Path of the task creating the writes.
|
||||||
"""
|
"""
|
||||||
query = (
|
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)
|
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:
|
with self.cursor() as cur:
|
||||||
cur.executemany(
|
cur.executemany(
|
||||||
@@ -473,6 +483,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
str(config["configurable"]["checkpoint_ns"]),
|
str(config["configurable"]["checkpoint_ns"]),
|
||||||
str(config["configurable"]["checkpoint_id"]),
|
str(config["configurable"]["checkpoint_id"]),
|
||||||
task_id,
|
task_id,
|
||||||
|
task_path,
|
||||||
WRITES_IDX_MAP.get(channel, idx),
|
WRITES_IDX_MAP.get(channel, idx),
|
||||||
channel,
|
channel,
|
||||||
*self.serde.dumps_typed(value),
|
*self.serde.dumps_typed(value),
|
||||||
@@ -568,7 +579,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
)
|
)
|
||||||
cur.execute(stage2_sql, stage2_params)
|
cur.execute(stage2_sql, stage2_params)
|
||||||
stage2_rows = cast(
|
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:
|
else:
|
||||||
stage2_rows = []
|
stage2_rows = []
|
||||||
|
|||||||
@@ -57,7 +57,7 @@ def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
|
|||||||
for n in chain_lens:
|
for n in chain_lens:
|
||||||
cid_placeholders = ",".join("?" * n)
|
cid_placeholders = ",".join("?" * n)
|
||||||
branches.append(
|
branches.append(
|
||||||
"SELECT checkpoint_id, channel, task_id, idx, type, value "
|
"SELECT checkpoint_id, channel, task_id, idx, type, value, task_path "
|
||||||
"FROM writes "
|
"FROM writes "
|
||||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
|
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
|
||||||
f"AND checkpoint_id IN ({cid_placeholders})"
|
f"AND checkpoint_id IN ({cid_placeholders})"
|
||||||
@@ -130,29 +130,31 @@ def build_delta_channels_writes_history(
|
|||||||
chain_by_ch: Mapping[str, list[str]],
|
chain_by_ch: Mapping[str, list[str]],
|
||||||
seed_val_by_ch: Mapping[str, Any],
|
seed_val_by_ch: Mapping[str, Any],
|
||||||
seeded: set[str],
|
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,
|
serde: Any,
|
||||||
) -> dict[str, DeltaChannelHistory]:
|
) -> dict[str, DeltaChannelHistory]:
|
||||||
"""Demux stage-2 rows per channel; produce per-channel histories.
|
"""Demux stage-2 rows per channel; produce per-channel histories.
|
||||||
|
|
||||||
Stage-2 rows are `(checkpoint_id, channel, task_id, idx, type, value)`.
|
Stage-2 rows are
|
||||||
Final write order is oldest→newest globally and `(task_id, idx)` within
|
`(checkpoint_id, channel, task_id, idx, type, value, task_path)`.
|
||||||
a checkpoint, matching the contract on `DeltaChannelHistory.writes`.
|
Final write order is oldest→newest globally and
|
||||||
|
`(task_path, task_id, idx)` within a checkpoint, matching the contract
|
||||||
|
on `DeltaChannelHistory.writes`.
|
||||||
|
|
||||||
`seed` is omitted when the walk reached a true root with no snapshot
|
`seed` is omitted when the walk reached a true root with no snapshot
|
||||||
found (channel never entered `seeded`); consumers treat absence as
|
found (channel never entered `seeded`); consumers treat absence as
|
||||||
"start empty".
|
"start empty".
|
||||||
"""
|
"""
|
||||||
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
writes_by_ch_by_cid: dict[
|
||||||
ch: {} for ch in channels
|
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 in stage2_rows:
|
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(
|
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 cid_map in writes_by_ch_by_cid.values():
|
||||||
for ws in cid_map.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] = {}
|
result: dict[str, DeltaChannelHistory] = {}
|
||||||
for ch in channels:
|
for ch in channels:
|
||||||
@@ -161,7 +163,7 @@ def build_delta_channels_writes_history(
|
|||||||
collected: list[PendingWrite] = []
|
collected: list[PendingWrite] = []
|
||||||
# Chain is newest-first; iterate oldest-first for the public order.
|
# Chain is newest-first; iterate oldest-first for the public order.
|
||||||
for cid in reversed(chain_cids):
|
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(
|
collected.append(
|
||||||
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
|
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -331,6 +331,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||||
checkpoint_id TEXT NOT NULL,
|
checkpoint_id TEXT NOT NULL,
|
||||||
task_id TEXT NOT NULL,
|
task_id TEXT NOT NULL,
|
||||||
|
task_path TEXT NOT NULL DEFAULT '',
|
||||||
idx INTEGER NOT NULL,
|
idx INTEGER NOT NULL,
|
||||||
channel TEXT NOT NULL,
|
channel TEXT NOT NULL,
|
||||||
type TEXT,
|
type TEXT,
|
||||||
@@ -341,6 +342,17 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
):
|
):
|
||||||
await self.conn.commit()
|
await self.conn.commit()
|
||||||
|
|
||||||
|
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
|
||||||
|
# created before `task_path` existed and is a no-op on the rest.
|
||||||
|
try:
|
||||||
|
await self.conn.execute(
|
||||||
|
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
|
||||||
|
)
|
||||||
|
await self.conn.commit()
|
||||||
|
except aiosqlite.OperationalError as e:
|
||||||
|
if "duplicate column name" not in str(e):
|
||||||
|
raise
|
||||||
|
|
||||||
self.is_setup = True
|
self.is_setup = True
|
||||||
|
|
||||||
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
||||||
@@ -576,9 +588,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
task_path: Path of the task creating the writes.
|
task_path: Path of the task creating the writes.
|
||||||
"""
|
"""
|
||||||
query = (
|
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)
|
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()
|
await self.setup()
|
||||||
async with self.lock, self.conn.cursor() as cur:
|
async with self.lock, self.conn.cursor() as cur:
|
||||||
@@ -590,6 +602,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
str(config["configurable"]["checkpoint_ns"]),
|
str(config["configurable"]["checkpoint_ns"]),
|
||||||
str(config["configurable"]["checkpoint_id"]),
|
str(config["configurable"]["checkpoint_id"]),
|
||||||
task_id,
|
task_id,
|
||||||
|
task_path,
|
||||||
WRITES_IDX_MAP.get(channel, idx),
|
WRITES_IDX_MAP.get(channel, idx),
|
||||||
channel,
|
channel,
|
||||||
*self.serde.dumps_typed(value),
|
*self.serde.dumps_typed(value),
|
||||||
@@ -681,7 +694,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
)
|
)
|
||||||
await cur.execute(stage2_sql, stage2_params)
|
await cur.execute(stage2_sql, stage2_params)
|
||||||
stage2_rows = cast(
|
stage2_rows = cast(
|
||||||
"list[tuple[str, str, str, int, str, bytes]]",
|
"list[tuple[str, str, str, int, str, bytes, str]]",
|
||||||
await cur.fetchall(),
|
await cur.fetchall(),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -0,0 +1,87 @@
|
|||||||
|
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.aio import AsyncSqliteSaver
|
||||||
|
|
||||||
|
WRITES_BEFORE_TASK_PATH = """
|
||||||
|
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)
|
||||||
|
);
|
||||||
|
INSERT INTO writes VALUES ('t', '', 'c', 'old-task', 0, 'ch', 'null', X'');
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def legacy_db(tmp_path: Path) -> Path:
|
||||||
|
db = tmp_path / "legacy.sqlite"
|
||||||
|
with sqlite3.connect(db) as conn:
|
||||||
|
conn.executescript(WRITES_BEFORE_TASK_PATH)
|
||||||
|
return db
|
||||||
|
|
||||||
|
|
||||||
|
def test_setup_migrates_legacy_writes_table_repeatably(legacy_db: Path) -> None:
|
||||||
|
for _ in range(2):
|
||||||
|
with SqliteSaver.from_conn_string(str(legacy_db)) as saver:
|
||||||
|
saver.setup()
|
||||||
|
rows = saver.conn.execute(
|
||||||
|
"SELECT task_id, task_path FROM writes"
|
||||||
|
).fetchall()
|
||||||
|
assert rows == [("old-task", "")]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("fresh", [True, False], ids=["fresh", "legacy"])
|
||||||
|
def test_put_writes_persists_task_path(
|
||||||
|
tmp_path: Path, legacy_db: Path, fresh: bool
|
||||||
|
) -> None:
|
||||||
|
db = tmp_path / "fresh.sqlite" if fresh else legacy_db
|
||||||
|
with SqliteSaver.from_conn_string(str(db)) 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_path FROM writes WHERE task_id = 'task-1'"
|
||||||
|
).fetchall()
|
||||||
|
assert stored == [("~__pregel_pull, node",)]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_async_setup_migrates_legacy_writes_table_repeatably(
|
||||||
|
legacy_db: Path,
|
||||||
|
) -> None:
|
||||||
|
for _ in range(2):
|
||||||
|
async with AsyncSqliteSaver.from_conn_string(str(legacy_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"
|
||||||
|
)
|
||||||
|
|
||||||
|
async with aiosqlite.connect(legacy_db) as conn:
|
||||||
|
async with conn.execute(
|
||||||
|
"SELECT DISTINCT task_id, task_path FROM writes ORDER BY task_id"
|
||||||
|
) as cur:
|
||||||
|
assert await cur.fetchall() == [
|
||||||
|
("old-task", ""),
|
||||||
|
("task-1", "~__pregel_pull, node"),
|
||||||
|
]
|
||||||
@@ -162,6 +162,13 @@ class DeltaChannelHistory(TypedDict):
|
|||||||
Always present; possibly empty. Already filtered to one channel.
|
Always present; possibly empty. Already filtered to one channel.
|
||||||
Writes stored at the target checkpoint itself are pending for the
|
Writes stored at the target checkpoint itself are pending for the
|
||||||
next super-step and are excluded.
|
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.
|
||||||
* `seed` — the stored value at the nearest ancestor whose
|
* `seed` — the stored value at the nearest ancestor whose
|
||||||
`channel_values[ch]` is populated. Omitted if the walk reached the
|
`channel_values[ch]` is populated. Omitted if the walk reached the
|
||||||
root without finding any stored value (consumer treats absence as
|
root without finding any stored value (consumer treats absence as
|
||||||
@@ -611,6 +618,11 @@ class BaseCheckpointSaver(Generic[V]):
|
|||||||
`PostgresSaver`) override for performance; the return contract is
|
`PostgresSaver`) override for performance; the return contract is
|
||||||
fixed here.
|
fixed here.
|
||||||
|
|
||||||
|
`PendingWrite` carries no `task_path`, so this default replays each
|
||||||
|
checkpoint's writes in `get_tuple`'s `pending_writes` order. Savers
|
||||||
|
that do not return `pending_writes` ordered by
|
||||||
|
`(task_path, task_id, idx)` must override it.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Configuration identifying the target checkpoint.
|
config: Configuration identifying the target checkpoint.
|
||||||
channels: Channel names to walk for. Empty → empty mapping.
|
channels: Channel names to walk for. Empty → empty mapping.
|
||||||
|
|||||||
@@ -199,8 +199,8 @@ class InMemorySaver(
|
|||||||
terminated_here.add(ch)
|
terminated_here.add(ch)
|
||||||
|
|
||||||
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
|
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
|
||||||
for (_task_id, _idx), (tid, ch, serialized, _) in sorted(
|
for _, (tid, ch, serialized, _) in sorted(
|
||||||
step_writes.items(), reverse=True
|
step_writes.items(), key=lambda kv: (kv[1][3], kv[0]), reverse=True
|
||||||
):
|
):
|
||||||
if ch not in remaining:
|
if ch not in remaining:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -0,0 +1,83 @@
|
|||||||
|
"""`DeltaChannel` replay must apply parallel writes in the order `invoke` did."""
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
# Sorted, because live execution applies PULL tasks in node-name order.
|
||||||
|
FAN_OUT_NAMES = ["a", "b", "c", "d", "e", "f", "g", "h"]
|
||||||
|
|
||||||
|
|
||||||
|
def _append_reducer(current: list, updates: list) -> list:
|
||||||
|
return [*current, *(x for u in updates for x in u)]
|
||||||
|
|
||||||
|
|
||||||
|
def _build_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
|
||||||
|
class State(TypedDict):
|
||||||
|
items: Annotated[
|
||||||
|
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
|
||||||
|
]
|
||||||
|
|
||||||
|
def make_node(label: str) -> Any:
|
||||||
|
return lambda state: {"items": [label]}
|
||||||
|
|
||||||
|
builder = StateGraph(State)
|
||||||
|
for name in FAN_OUT_NAMES:
|
||||||
|
builder.add_node(name, make_node(name))
|
||||||
|
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:
|
||||||
|
graph = _build_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 == FAN_OUT_NAMES
|
||||||
|
assert replayed == live
|
||||||
|
|
||||||
|
|
||||||
|
async def test_continuing_thread_preserves_committed_prefix(
|
||||||
|
async_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
graph = _build_fan_out_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_state_history_reports_live_order_at_every_step(
|
||||||
|
async_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
runs = 3
|
||||||
|
graph = _build_fan_out_graph(async_checkpointer)
|
||||||
|
config = {"configurable": {"thread_id": "1"}}
|
||||||
|
for _ in range(runs):
|
||||||
|
await graph.ainvoke({"items": []}, config)
|
||||||
|
live = FAN_OUT_NAMES * runs
|
||||||
|
|
||||||
|
seen = [
|
||||||
|
s.values["items"]
|
||||||
|
async for s in graph.aget_state_history(config)
|
||||||
|
if "items" in s.values
|
||||||
|
]
|
||||||
|
|
||||||
|
assert max(map(len, seen)) == len(live)
|
||||||
|
for values in seen:
|
||||||
|
assert values == live[: len(values)], f"{values} is not a prefix of {live}"
|
||||||
Reference in New Issue
Block a user