mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-08 09:25:08 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b791c1f1d5 | ||
|
|
87f1c8eb9a | ||
|
|
736a18eaab | ||
|
|
ec4f8535f8 | ||
|
|
ce14d33687 | ||
|
|
708aaa4ebc | ||
|
|
39c523eb0a | ||
|
|
3a1796ecf0 | ||
|
|
2b93e8b79f |
+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 = "ffffffff-ffff-ffff-ffff-ffffffffffff"
|
||||
|
||||
|
||||
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 = [
|
||||
test_history_returns_writes_oldest_first,
|
||||
test_history_seed_is_nearest_snapshot,
|
||||
@@ -276,6 +331,8 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
||||
test_history_walk_to_root_no_seed,
|
||||
test_history_migration_plain_value_as_seed,
|
||||
test_history_seed_ancestor_own_writes_are_replayed,
|
||||
test_history_orders_parallel_writes_by_task_path,
|
||||
test_history_orders_pathless_writes_first,
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -176,6 +176,35 @@ async def test_get_tuple_pending_writes(saver: BaseCheckpointSaver) -> None:
|
||||
)
|
||||
|
||||
|
||||
async def test_get_tuple_pending_writes_in_writes_sort_key_order(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""pending_writes come back in writes_sort_key order, not put or task_id order."""
|
||||
config = generate_config(str(uuid4()))
|
||||
stored = await saver.aput(config, generate_checkpoint(), generate_metadata(), {})
|
||||
await saver.aput_writes(
|
||||
stored, [("ch", "b")], "00000000-0000-0000-0000-000000000000", "~pull, 02"
|
||||
)
|
||||
await saver.aput_writes(
|
||||
stored,
|
||||
[("ch", "a1"), ("ch", "a2")],
|
||||
"ffffffff-ffff-ffff-ffff-ffffffffffff",
|
||||
"~pull, 01",
|
||||
)
|
||||
await saver.aput_writes(
|
||||
stored, [("ch", "input")], "88888888-8888-8888-8888-888888888888"
|
||||
)
|
||||
|
||||
tup = await saver.aget_tuple(stored)
|
||||
assert tup is not None
|
||||
values = [w[2] for w in tup.pending_writes or []]
|
||||
assert values == ["input", "a1", "a2", "b"], (
|
||||
f"Expected writes_sort_key order ['input', 'a1', 'a2', 'b'], got {values}. "
|
||||
"Put order gives ['b', 'a1', 'a2', 'input'], task_id order gives "
|
||||
"['b', 'input', 'a1', 'a2']."
|
||||
)
|
||||
|
||||
|
||||
async def test_get_tuple_respects_namespace(saver: BaseCheckpointSaver) -> None:
|
||||
"""checkpoint_ns filtering."""
|
||||
tid = str(uuid4())
|
||||
@@ -223,6 +252,7 @@ ALL_GET_TUPLE_TESTS = [
|
||||
test_get_tuple_metadata,
|
||||
test_get_tuple_parent_config,
|
||||
test_get_tuple_pending_writes,
|
||||
test_get_tuple_pending_writes_in_writes_sort_key_order,
|
||||
test_get_tuple_respects_namespace,
|
||||
test_get_tuple_nonexistent_checkpoint_id,
|
||||
]
|
||||
|
||||
Generated
+1
-1
@@ -279,7 +279,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -687,5 +687,23 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
self.adelete_thread(thread_id), self.loop
|
||||
).result()
|
||||
|
||||
def get_delta_channel_history(
|
||||
self, *, config: RunnableConfig, channels: Sequence[str]
|
||||
) -> Mapping[str, DeltaChannelHistory]:
|
||||
"""Sync bridge to `aget_delta_channel_history`, guarded like `get_tuple`."""
|
||||
try:
|
||||
if asyncio.get_running_loop() is self.loop:
|
||||
raise asyncio.InvalidStateError(
|
||||
"Synchronous calls to AsyncPostgresSaver are only allowed from a "
|
||||
"different thread. From the main thread, use the async interface. "
|
||||
"For example, use `await checkpointer.aget_delta_channel_history(...)`."
|
||||
)
|
||||
except RuntimeError:
|
||||
pass
|
||||
return asyncio.run_coroutine_threadsafe(
|
||||
self.aget_delta_channel_history(config=config, channels=channels),
|
||||
self.loop,
|
||||
).result()
|
||||
|
||||
|
||||
__all__ = ["AsyncPostgresSaver", "AsyncShallowPostgresSaver", "Conn"]
|
||||
|
||||
@@ -14,6 +14,7 @@ from langgraph.checkpoint.base import (
|
||||
DeltaChannelHistory,
|
||||
PendingWrite,
|
||||
get_checkpoint_id,
|
||||
writes_sort_key,
|
||||
)
|
||||
from langgraph.checkpoint.serde.types import TASKS
|
||||
from psycopg.types.json import Jsonb
|
||||
@@ -109,7 +110,7 @@ select
|
||||
) as channel_values,
|
||||
(
|
||||
select
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob, convert_to(cw.task_path, 'UTF8'), cw.idx::text::bytea])
|
||||
from checkpoint_writes cw
|
||||
where cw.thread_id = checkpoints.thread_id
|
||||
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
||||
@@ -168,6 +169,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
|
||||
|
||||
@@ -297,7 +299,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)"
|
||||
@@ -305,7 +307,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"
|
||||
@@ -473,10 +476,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]] = {}
|
||||
|
||||
@@ -487,8 +491,14 @@ 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"],
|
||||
r["task_path"],
|
||||
),
|
||||
)
|
||||
)
|
||||
else: # kind == "b"
|
||||
@@ -497,10 +507,10 @@ 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
|
||||
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: writes_sort_key(w[4], w[2], w[3]), reverse=True)
|
||||
|
||||
result: dict[str, DeltaChannelHistory] = {}
|
||||
for ch in channels:
|
||||
@@ -510,7 +520,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()
|
||||
@@ -553,20 +565,15 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
]
|
||||
|
||||
def _load_writes(
|
||||
self, writes: list[tuple[bytes, bytes, bytes, bytes]]
|
||||
self, writes: list[tuple[bytes, bytes, bytes, bytes, bytes, bytes]] | None
|
||||
) -> list[tuple[str, str, Any]]:
|
||||
return (
|
||||
[
|
||||
(
|
||||
tid.decode(),
|
||||
channel.decode(),
|
||||
self.serde.loads_typed((t.decode(), v)),
|
||||
)
|
||||
for tid, channel, t, v in writes
|
||||
]
|
||||
if writes
|
||||
else []
|
||||
)
|
||||
return [
|
||||
(tid.decode(), channel.decode(), self.serde.loads_typed((t.decode(), v)))
|
||||
for tid, channel, t, v, _, _ in sorted(
|
||||
writes or [],
|
||||
key=lambda w: writes_sort_key(w[4].decode(), w[0].decode(), int(w[5])),
|
||||
)
|
||||
]
|
||||
|
||||
def _dump_writes(
|
||||
self,
|
||||
|
||||
@@ -97,7 +97,7 @@ select
|
||||
) as channel_values,
|
||||
(
|
||||
select
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob, convert_to(cw.task_path, 'UTF8'), cw.idx::text::bytea])
|
||||
from checkpoint_writes cw
|
||||
where cw.thread_id = checkpoints.thread_id
|
||||
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
||||
|
||||
@@ -12,7 +12,7 @@ readme = "README.md"
|
||||
license = "MIT"
|
||||
license-files = ['LICENSE']
|
||||
dependencies = [
|
||||
"langgraph-checkpoint>=4.1.0,<5.0.0",
|
||||
"langgraph-checkpoint>=4.3.0,<5.0.0",
|
||||
"orjson>=3.11.5",
|
||||
"psycopg>=3.2.0",
|
||||
"psycopg-pool>=3.2.0",
|
||||
|
||||
Generated
+1
-1
@@ -278,7 +278,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -29,7 +29,11 @@ from langgraph.checkpoint.sqlite._delta import (
|
||||
build_delta_stage2_sql,
|
||||
step_walk_with_row,
|
||||
)
|
||||
from langgraph.checkpoint.sqlite.utils import search_where
|
||||
from langgraph.checkpoint.sqlite.utils import (
|
||||
load_pending_writes,
|
||||
pending_writes_sql,
|
||||
search_where,
|
||||
)
|
||||
|
||||
_AIO_ERROR_MSG = (
|
||||
"The SqliteSaver does not support async methods. "
|
||||
@@ -81,6 +85,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
conn: sqlite3.Connection
|
||||
is_setup: bool
|
||||
_has_task_path: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -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,19 @@ 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:
|
||||
# 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
|
||||
|
||||
@@ -260,7 +279,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
}
|
||||
# find any pending writes
|
||||
cur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(
|
||||
str(config["configurable"]["thread_id"]),
|
||||
checkpoint_ns,
|
||||
@@ -286,10 +305,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
for task_id, channel, type, value in cur
|
||||
],
|
||||
load_pending_writes(cur, self.serde),
|
||||
)
|
||||
|
||||
def list(
|
||||
@@ -351,7 +367,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
metadata,
|
||||
) in cur:
|
||||
wcur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(thread_id, checkpoint_ns, checkpoint_id),
|
||||
)
|
||||
yield CheckpointTuple(
|
||||
@@ -378,10 +394,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
for task_id, channel, type, value in wcur
|
||||
],
|
||||
load_pending_writes(wcur, self.serde),
|
||||
)
|
||||
|
||||
def put(
|
||||
@@ -460,9 +473,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 +486,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),
|
||||
@@ -559,6 +573,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:
|
||||
@@ -569,7 +584,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 = []
|
||||
|
||||
@@ -24,7 +24,11 @@ from __future__ import annotations
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite
|
||||
from langgraph.checkpoint.base import (
|
||||
DeltaChannelHistory,
|
||||
PendingWrite,
|
||||
writes_sort_key,
|
||||
)
|
||||
|
||||
# Stage 1 streams target, then its ancestors nearest-first, by following
|
||||
# `parent_checkpoint_id` rather than id order: ids are only monotonic within
|
||||
@@ -56,7 +60,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
|
||||
@@ -70,11 +76,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 "
|
||||
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})"
|
||||
@@ -141,29 +148,31 @@ 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 `writes_sort_key`
|
||||
within a checkpoint, matching the contract on
|
||||
`DeltaChannelHistory.writes`.
|
||||
|
||||
`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: writes_sort_key(w[4], w[2], w[3]))
|
||||
|
||||
result: dict[str, DeltaChannelHistory] = {}
|
||||
for ch in channels:
|
||||
@@ -172,7 +181,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)))
|
||||
)
|
||||
|
||||
@@ -30,7 +30,11 @@ from langgraph.checkpoint.sqlite._delta import (
|
||||
build_delta_stage2_sql,
|
||||
step_walk_with_row,
|
||||
)
|
||||
from langgraph.checkpoint.sqlite.utils import search_where
|
||||
from langgraph.checkpoint.sqlite.utils import (
|
||||
load_pending_writes,
|
||||
pending_writes_sql,
|
||||
search_where,
|
||||
)
|
||||
|
||||
T = TypeVar("T", bound=Callable)
|
||||
|
||||
@@ -114,6 +118,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
lock: asyncio.Lock
|
||||
is_setup: bool
|
||||
_has_task_path: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -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,21 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
):
|
||||
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:
|
||||
# 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
|
||||
|
||||
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
||||
@@ -395,7 +416,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
}
|
||||
# find any pending writes
|
||||
await cur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(
|
||||
str(config["configurable"]["thread_id"]),
|
||||
checkpoint_ns,
|
||||
@@ -421,10 +442,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
async for task_id, channel, type, value in cur
|
||||
],
|
||||
load_pending_writes(await cur.fetchall(), self.serde),
|
||||
)
|
||||
|
||||
async def alist(
|
||||
@@ -473,7 +491,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
metadata,
|
||||
) in cur:
|
||||
await wcur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(thread_id, checkpoint_ns, checkpoint_id),
|
||||
)
|
||||
yield CheckpointTuple(
|
||||
@@ -500,10 +518,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
async for task_id, channel, type, value in wcur
|
||||
],
|
||||
load_pending_writes(await wcur.fetchall(), self.serde),
|
||||
)
|
||||
|
||||
async def aput(
|
||||
@@ -576,9 +591,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 +605,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),
|
||||
@@ -671,6 +687,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:
|
||||
@@ -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:
|
||||
|
||||
@@ -2,11 +2,12 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Iterable, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.base import get_checkpoint_id
|
||||
from langgraph.checkpoint.base import PendingWrite, get_checkpoint_id, writes_sort_key
|
||||
from langgraph.checkpoint.serde.base import SerializerProtocol
|
||||
|
||||
_FILTER_PATTERN = re.compile(r"^[a-zA-Z0-9_.-]+$")
|
||||
|
||||
@@ -114,3 +115,23 @@ def search_where(
|
||||
param_values.append(get_checkpoint_id(before))
|
||||
|
||||
return ("WHERE " + " AND ".join(wheres) if wheres else "", param_values)
|
||||
|
||||
|
||||
def pending_writes_sql(has_task_path: bool) -> str:
|
||||
task_path = "task_path" if has_task_path else "''"
|
||||
return (
|
||||
f"SELECT task_id, channel, type, value, {task_path}, idx FROM writes "
|
||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ?"
|
||||
)
|
||||
|
||||
|
||||
def load_pending_writes(
|
||||
rows: Iterable[Any], serde: SerializerProtocol
|
||||
) -> list[PendingWrite]:
|
||||
"""Deserialize `pending_writes_sql` rows in `writes_sort_key` order."""
|
||||
return [
|
||||
(task_id, channel, serde.loads_typed((type_, value)))
|
||||
for task_id, channel, type_, value, _, _ in sorted(
|
||||
rows, key=lambda r: writes_sort_key(r[4], r[0], r[5])
|
||||
)
|
||||
]
|
||||
|
||||
@@ -12,7 +12,7 @@ readme = "README.md"
|
||||
license = "MIT"
|
||||
license-files = ['LICENSE']
|
||||
dependencies = [
|
||||
"langgraph-checkpoint>=4.1.0,<5.0.0",
|
||||
"langgraph-checkpoint>=4.3.0,<5.0.0",
|
||||
"aiosqlite>=0.20",
|
||||
"sqlite-vec>=0.1.6",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
import sqlite3
|
||||
from collections.abc import Iterator
|
||||
from contextlib import closing
|
||||
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"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def busy_db(tmp_path: Path) -> Iterator[Path]:
|
||||
db = tmp_path / "busy.sqlite"
|
||||
with SqliteSaver.from_conn_string(str(db)) as saver:
|
||||
saver.setup()
|
||||
with closing(sqlite3.connect(db, isolation_level=None)) as writer:
|
||||
writer.execute("BEGIN IMMEDIATE")
|
||||
yield db
|
||||
writer.execute("ROLLBACK")
|
||||
|
||||
|
||||
def test_setup_does_not_wait_on_another_writer(busy_db: Path) -> None:
|
||||
with closing(sqlite3.connect(busy_db, timeout=0)) as conn:
|
||||
SqliteSaver(conn).setup()
|
||||
|
||||
|
||||
async def test_async_setup_does_not_wait_on_another_writer(busy_db: Path) -> None:
|
||||
async with aiosqlite.connect(busy_db, timeout=0) as conn:
|
||||
await AsyncSqliteSaver(conn).setup()
|
||||
|
||||
|
||||
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")]}
|
||||
Generated
+1
-1
@@ -285,7 +285,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -162,6 +162,14 @@ 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 `writes_sort_key`,
|
||||
`(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` and `Command` updates, 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
|
||||
@@ -174,6 +182,20 @@ class DeltaChannelHistory(TypedDict):
|
||||
seed: NotRequired[Any]
|
||||
|
||||
|
||||
def writes_sort_key(
|
||||
task_path: str, task_id: str = "", idx: int = 0
|
||||
) -> tuple[str, str, int]:
|
||||
"""Sort key for the writes of one super-step.
|
||||
|
||||
Live execution applies a super-step's tasks in this order, so a saver
|
||||
that replays stored writes, as `get_delta_channel_history` does, must
|
||||
sort them by it too, or an order-sensitive reducer rebuilds a different
|
||||
value than the run produced. `task_path` is the string passed to
|
||||
`put_writes`.
|
||||
"""
|
||||
return (task_path, task_id, idx)
|
||||
|
||||
|
||||
class BaseCheckpointSaver(Generic[V]):
|
||||
"""Base class for creating a graph checkpointer.
|
||||
|
||||
@@ -244,7 +266,8 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
config: Configuration specifying which checkpoint to retrieve.
|
||||
|
||||
Returns:
|
||||
The requested checkpoint tuple, or `None` if not found.
|
||||
The requested checkpoint tuple, or `None` if not found. Its
|
||||
`pending_writes` must be in `writes_sort_key` order.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Implement this method in your custom checkpoint saver.
|
||||
@@ -434,7 +457,8 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
config: Configuration specifying which checkpoint to retrieve.
|
||||
|
||||
Returns:
|
||||
The requested checkpoint tuple, or `None` if not found.
|
||||
The requested checkpoint tuple, or `None` if not found. Its
|
||||
`pending_writes` must be in `writes_sort_key` order.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Implement this method in your custom checkpoint saver.
|
||||
@@ -611,6 +635,10 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
`PostgresSaver`) override for performance; the return contract is
|
||||
fixed here.
|
||||
|
||||
`PendingWrite` carries no `task_path`, so this default replays each
|
||||
checkpoint's writes in `get_tuple`'s `pending_writes` order, which
|
||||
`get_tuple` must return in `writes_sort_key` order.
|
||||
|
||||
Args:
|
||||
config: Configuration identifying the target checkpoint.
|
||||
channels: Channel names to walk for. Empty → empty mapping.
|
||||
|
||||
@@ -25,6 +25,7 @@ from langgraph.checkpoint.base import (
|
||||
SerializerProtocol,
|
||||
get_checkpoint_id,
|
||||
get_checkpoint_metadata,
|
||||
writes_sort_key,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -139,6 +140,15 @@ class InMemorySaver(
|
||||
result[k] = self.serde.loads_typed(vv)
|
||||
return result
|
||||
|
||||
def _ordered_writes(
|
||||
self, thread_id: str, checkpoint_ns: str, checkpoint_id: str
|
||||
) -> list[tuple[str, str, tuple[str, bytes], str]]:
|
||||
stored = self.writes.get((thread_id, checkpoint_ns, checkpoint_id), {})
|
||||
return [
|
||||
stored[k]
|
||||
for k in sorted(stored, key=lambda k: writes_sort_key(stored[k][3], *k))
|
||||
]
|
||||
|
||||
def get_delta_channel_history(
|
||||
self, *, config: RunnableConfig, channels: Sequence[str]
|
||||
) -> Mapping[str, DeltaChannelHistory]:
|
||||
@@ -198,9 +208,8 @@ class InMemorySaver(
|
||||
blob_value_by_ch[ch] = self.serde.loads_typed(blob_entry)
|
||||
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
|
||||
for tid, ch, serialized, _ in reversed(
|
||||
self._ordered_writes(thread_id, checkpoint_ns, cp_id)
|
||||
):
|
||||
if ch not in remaining:
|
||||
continue
|
||||
@@ -246,7 +255,7 @@ class InMemorySaver(
|
||||
if checkpoint_id := get_checkpoint_id(config):
|
||||
if saved := self.storage[thread_id][checkpoint_ns].get(checkpoint_id):
|
||||
checkpoint, metadata, parent_checkpoint_id = saved
|
||||
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
|
||||
writes = self._ordered_writes(thread_id, checkpoint_ns, checkpoint_id)
|
||||
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
|
||||
return CheckpointTuple(
|
||||
config=config,
|
||||
@@ -276,7 +285,7 @@ class InMemorySaver(
|
||||
if checkpoints := self.storage[thread_id][checkpoint_ns]:
|
||||
checkpoint_id = max(checkpoints.keys())
|
||||
checkpoint, metadata, parent_checkpoint_id = checkpoints[checkpoint_id]
|
||||
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
|
||||
writes = self._ordered_writes(thread_id, checkpoint_ns, checkpoint_id)
|
||||
checkpoint_ = self.serde.loads_typed(checkpoint)
|
||||
return CheckpointTuple(
|
||||
config={
|
||||
@@ -379,9 +388,9 @@ class InMemorySaver(
|
||||
elif limit is not None:
|
||||
limit -= 1
|
||||
|
||||
writes = self.writes[
|
||||
(thread_id, checkpoint_ns, checkpoint_id)
|
||||
].values()
|
||||
writes = self._ordered_writes(
|
||||
thread_id, checkpoint_ns, checkpoint_id
|
||||
)
|
||||
|
||||
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
description = "Library with base interfaces for LangGraph checkpoint savers."
|
||||
authors = []
|
||||
requires-python = ">=3.10"
|
||||
|
||||
Generated
+1
-1
@@ -301,7 +301,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.4.32"
|
||||
__version__ = "0.4.33"
|
||||
|
||||
@@ -32,7 +32,7 @@ from langgraph_cli.host_backend import (
|
||||
HostBackendError,
|
||||
SourceName,
|
||||
)
|
||||
from langgraph_cli.image_reference import ImageReference
|
||||
from langgraph_cli.image_reference import DIGEST_SEPARATOR, ImageReference
|
||||
from langgraph_cli.progress import Progress
|
||||
from langgraph_cli.util import warn_non_wolfi_distro
|
||||
|
||||
@@ -624,6 +624,41 @@ def format_deployments_table(deployments: Sequence[dict[str, object]]) -> str:
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _extract_listener_namespaces(listener: dict[str, object]) -> str:
|
||||
compute_config = listener.get("compute_config")
|
||||
namespaces = (
|
||||
compute_config.get("k8s_namespaces")
|
||||
if isinstance(compute_config, dict)
|
||||
else None
|
||||
)
|
||||
if isinstance(namespaces, list) and namespaces:
|
||||
return ", ".join(str(namespace) for namespace in namespaces)
|
||||
return "-"
|
||||
|
||||
|
||||
def format_listeners_table(listeners: Sequence[dict[str, object]]) -> str:
|
||||
headers = ("Listener ID", "Compute ID", "Namespaces")
|
||||
rows = [
|
||||
(
|
||||
str(listener.get("id", "-") or "-"),
|
||||
str(listener.get("compute_id", "-") or "-"),
|
||||
_extract_listener_namespaces(listener),
|
||||
)
|
||||
for listener in listeners
|
||||
]
|
||||
widths = [
|
||||
max(len(headers[index]), *(len(row[index]) for row in rows))
|
||||
for index in range(len(headers))
|
||||
]
|
||||
|
||||
def format_row(row: Sequence[str]) -> str:
|
||||
return " ".join(value.ljust(widths[index]) for index, value in enumerate(row))
|
||||
|
||||
lines = [format_row(headers), format_row(tuple("-" * width for width in widths))]
|
||||
lines.extend(format_row(row) for row in rows)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def format_revisions_table(revisions: Sequence[dict[str, object]]) -> str:
|
||||
headers = ("Revision ID", "Status", "Created At")
|
||||
latest_deployed_seen = False
|
||||
@@ -1544,9 +1579,9 @@ def _ensure_customer_registry_source(existing: ExistingDeployment) -> None:
|
||||
if existing.source != _CUSTOMER_REGISTRY_SOURCE:
|
||||
raise click.UsageError(
|
||||
f"Deployment {existing.id} was not created from an external image "
|
||||
"and cannot be updated with --push-to. Run without --push-to to keep "
|
||||
"its current build mode, or use a different --name to create a new "
|
||||
"deployment."
|
||||
"and cannot be updated with --push-to or --image-uri. Run without "
|
||||
"either flag to keep its current build mode, or use a different "
|
||||
"--name to create a new deployment."
|
||||
)
|
||||
|
||||
|
||||
@@ -1599,9 +1634,22 @@ class RemoteBuildSource:
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CustomerRegistrySource:
|
||||
class BuildAndPush:
|
||||
reference: ImageReference
|
||||
prebuilt_image: str | None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PublishedImage:
|
||||
image_uri: str
|
||||
|
||||
|
||||
ImageSource = BuildAndPush | PublishedImage
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CustomerRegistrySource:
|
||||
image: ImageSource
|
||||
requested_placement: RequestedPlacement
|
||||
|
||||
def run(self, ctx: DeployContext) -> DeployOutcome:
|
||||
@@ -1686,16 +1734,22 @@ class CustomerRegistrySource:
|
||||
)
|
||||
|
||||
def _publish(self, ctx: DeployContext, step: int) -> tuple[str, int]:
|
||||
image = str(self.reference)
|
||||
if isinstance(self.image, PublishedImage):
|
||||
return self.image.image_uri, step
|
||||
image = str(self.image.reference)
|
||||
with Runner() as runner:
|
||||
if self.prebuilt_image:
|
||||
_log_deploy_step(step, f"Validating image {self.prebuilt_image}")
|
||||
if self.image.prebuilt_image:
|
||||
_log_deploy_step(step, f"Validating image {self.image.prebuilt_image}")
|
||||
_validate_prebuilt_image(
|
||||
runner, self.prebuilt_image, verbose=ctx.verbose
|
||||
runner, self.image.prebuilt_image, verbose=ctx.verbose
|
||||
)
|
||||
runner.run(
|
||||
subp_exec(
|
||||
"docker", "tag", self.prebuilt_image, image, verbose=ctx.verbose
|
||||
"docker",
|
||||
"tag",
|
||||
self.image.prebuilt_image,
|
||||
image,
|
||||
verbose=ctx.verbose,
|
||||
)
|
||||
)
|
||||
else:
|
||||
@@ -1733,20 +1787,35 @@ def _push_reference(push_to: str, tag: str | None) -> ImageReference:
|
||||
return reference.with_tag(normalize_image_tag(tag or _DEFAULT_IMAGE_TAG))
|
||||
|
||||
|
||||
def _validate_image_uri(image_uri: str) -> str:
|
||||
value = image_uri.strip()
|
||||
if not value:
|
||||
raise click.UsageError("--image-uri must not be empty.")
|
||||
if DIGEST_SEPARATOR not in value:
|
||||
raise click.UsageError(
|
||||
"--image-uri must pin a digest, e.g. "
|
||||
f"repository{DIGEST_SEPARATOR}<sha256 hex>. Kubernetes can cache "
|
||||
"images by tag, so redeploying a mutable tag may silently keep "
|
||||
"running the previous image."
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def _select_source(
|
||||
*,
|
||||
push_to: str | None,
|
||||
image: str | None,
|
||||
image_uri: str | None,
|
||||
image_name: str | None,
|
||||
tag: str | None,
|
||||
remote_build_flag: bool | None,
|
||||
placement: RequestedPlacement,
|
||||
selector: DeploymentSelector,
|
||||
) -> DeploymentSource:
|
||||
if push_to is None and placement.requested:
|
||||
if push_to is None and image_uri is None and placement.requested:
|
||||
raise click.UsageError(
|
||||
"--listener-id and --k8s-namespace only apply when creating a "
|
||||
"deployment with --push-to."
|
||||
"deployment with --push-to or --image-uri."
|
||||
)
|
||||
if placement.requested and isinstance(selector, ById):
|
||||
raise click.UsageError(
|
||||
@@ -1754,6 +1823,19 @@ def _select_source(
|
||||
"they cannot be set for an existing --deployment-id. Drop them, or "
|
||||
"create a new deployment with --name."
|
||||
)
|
||||
if image_uri is not None:
|
||||
if push_to is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --push-to.")
|
||||
if image is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --image.")
|
||||
if tag is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --tag.")
|
||||
if remote_build_flag is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --remote.")
|
||||
return CustomerRegistrySource(
|
||||
image=PublishedImage(_validate_image_uri(image_uri)),
|
||||
requested_placement=placement,
|
||||
)
|
||||
if push_to is not None:
|
||||
if remote_build_flag is True:
|
||||
raise click.UsageError("--push-to cannot be combined with --remote.")
|
||||
@@ -1761,8 +1843,7 @@ def _select_source(
|
||||
if image is None:
|
||||
_require_local_docker()
|
||||
return CustomerRegistrySource(
|
||||
reference=reference,
|
||||
prebuilt_image=image,
|
||||
image=BuildAndPush(reference=reference, prebuilt_image=image),
|
||||
requested_placement=placement,
|
||||
)
|
||||
if image and remote_build_flag is True:
|
||||
@@ -2075,19 +2156,30 @@ def _deploy_base_options(
|
||||
"Give the tag here or with --tag (default: latest)."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
"--image-uri",
|
||||
help=(
|
||||
"Deploy an image that's already in a registry you manage, "
|
||||
"without building, retagging, or pushing anything. For "
|
||||
"self-hosted and hybrid LangSmith. Give the full reference, "
|
||||
"e.g. 123456789.dkr.ecr.us-east-1.amazonaws.com/agents/"
|
||||
"my-agent:v1.2.3 or ...@sha256:<digest>. Cannot be combined "
|
||||
"with --push-to, --image, --tag, or --remote."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
"--listener-id",
|
||||
help=(
|
||||
"Listener that will run the deployment, for workspaces that "
|
||||
"deploy through a listener in your own cluster. Only used when "
|
||||
"creating a deployment with --push-to."
|
||||
"creating a deployment with --push-to or --image-uri."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
"--k8s-namespace",
|
||||
help=(
|
||||
"Kubernetes namespace the listener deploys into. Only used when "
|
||||
"creating a deployment with --push-to."
|
||||
"creating a deployment with --push-to or --image-uri."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
@@ -2200,6 +2292,7 @@ def _deploy_cmd(
|
||||
image_name: str | None,
|
||||
image: str | None,
|
||||
push_to: str | None,
|
||||
image_uri: str | None,
|
||||
listener_id: str | None,
|
||||
k8s_namespace: str | None,
|
||||
tag: str | None,
|
||||
@@ -2273,6 +2366,7 @@ def _deploy_cmd(
|
||||
source = _select_source(
|
||||
push_to=push_to,
|
||||
image=image,
|
||||
image_uri=image_uri,
|
||||
image_name=image_name,
|
||||
tag=tag,
|
||||
remote_build_flag=remote_build_flag,
|
||||
@@ -2408,6 +2502,41 @@ def deploy_list(
|
||||
click.echo(format_deployments_table(deployments))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# deploy listeners
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@deploy.group(
|
||||
"listeners",
|
||||
cls=NestedHelpGroup,
|
||||
help="[Beta] Inspect listeners available to this workspace.",
|
||||
)
|
||||
def deploy_listeners() -> None:
|
||||
pass
|
||||
|
||||
|
||||
@OPT_HOST_API_KEY
|
||||
@OPT_HOST_URL
|
||||
@deploy_listeners.command(
|
||||
"list",
|
||||
help=(
|
||||
"[Beta] List listeners available to this workspace.\n\n"
|
||||
"Pass a listener's id to `langgraph deploy --push-to ... "
|
||||
"--listener-id <id>` to deploy through it."
|
||||
),
|
||||
)
|
||||
def deploy_listeners_list(api_key: str | None, host_url: str | None) -> None:
|
||||
client = _create_host_backend_client(host_url, api_key)
|
||||
listeners = _call_host_backend_with_optional_tenant(
|
||||
client, lambda c: c.list_listeners()
|
||||
)
|
||||
if not listeners:
|
||||
click.echo("No listeners found for this workspace.")
|
||||
return
|
||||
click.echo(format_listeners_table(listeners))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# deploy revisions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -453,6 +453,88 @@ def test_deploy_list_command_no_results(monkeypatch) -> None:
|
||||
assert result.output.strip() == "No deployments found."
|
||||
|
||||
|
||||
def test_deploy_listeners_list_command(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
captured["host_url"] = host_url
|
||||
captured["api_key"] = api_key
|
||||
captured["tenant_id"] = tenant_id or ""
|
||||
|
||||
def list_listeners(self):
|
||||
return [
|
||||
{
|
||||
"id": "listener-1",
|
||||
"compute_id": "prod-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents"]},
|
||||
},
|
||||
{
|
||||
"id": "listener-2",
|
||||
"compute_id": "multi-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents", "agents-staging"]},
|
||||
},
|
||||
]
|
||||
|
||||
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"listeners",
|
||||
"list",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert captured == {
|
||||
"host_url": "https://api.example.com",
|
||||
"api_key": "test-key",
|
||||
"tenant_id": "",
|
||||
}
|
||||
assert "Listener ID" in result.output
|
||||
assert "Compute ID" in result.output
|
||||
assert "Namespaces" in result.output
|
||||
assert "listener-1" in result.output
|
||||
assert "prod-cluster" in result.output
|
||||
assert "agents, agents-staging" in result.output
|
||||
|
||||
|
||||
def test_deploy_listeners_list_command_no_results(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
pass
|
||||
|
||||
def list_listeners(self):
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"listeners",
|
||||
"list",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert result.output.strip() == "No listeners found for this workspace."
|
||||
|
||||
|
||||
def test_deploy_revisions_list_command(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
@@ -695,6 +695,106 @@ def test_push_to_with_deployment_id_fetches_the_deployment_once(
|
||||
]
|
||||
|
||||
|
||||
def test_image_uri_creates_an_external_deployment_without_any_docker_work(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.timeline == [LIST_DEPLOYMENTS, CREATE_DEPLOYMENT]
|
||||
assert deploy_project.control_plane.bodies[CREATE_DEPLOYMENT][
|
||||
"source_revision_config"
|
||||
] == {"image_uri": EXTERNAL_DIGEST}
|
||||
|
||||
|
||||
def test_image_uri_updates_an_existing_external_deployment_without_any_docker_work(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
deploy_project.control_plane.existing_deployments = [
|
||||
{"id": "dep-ext", "name": "my-app", "source": "external_docker"}
|
||||
]
|
||||
|
||||
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.timeline == [LIST_DEPLOYMENTS, _patch("dep-ext")]
|
||||
assert deploy_project.control_plane.bodies[_patch("dep-ext")] == {
|
||||
"source_revision_config": {"image_uri": EXTERNAL_DIGEST},
|
||||
"secrets": [],
|
||||
"tracked_packages": TRACKED_PACKAGES,
|
||||
}
|
||||
|
||||
|
||||
def test_image_uri_places_a_new_deployment_on_the_only_listener(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
deploy_project.control_plane.listeners = [LISTENER]
|
||||
|
||||
result = deploy_project.run(
|
||||
"--image-uri", EXTERNAL_DIGEST, host_url=CLOUD_CONTROL_PLANE_URL
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.control_plane.bodies[CREATE_DEPLOYMENT]["source_config"] == {
|
||||
"resource_spec": {},
|
||||
"listener_id": LISTENER_ID,
|
||||
"listener_config": {"k8s_namespace": "agents"},
|
||||
}
|
||||
|
||||
|
||||
def test_image_uri_rejects_a_non_external_deployment_before_any_docker_work(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
deploy_project.control_plane.existing_deployments = [
|
||||
{"id": "dep-cli", "name": "my-app", "source": "internal_docker"}
|
||||
]
|
||||
|
||||
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
|
||||
|
||||
assert result.exit_code != 0
|
||||
assert "cannot be updated with --push-to or --image-uri" in result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("args", "message"),
|
||||
[
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--push-to", PUSH_REPOSITORY),
|
||||
"--image-uri cannot be combined with --push-to.",
|
||||
id="with_push_to",
|
||||
),
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--image", "local/app:dev"),
|
||||
"--image-uri cannot be combined with --image.",
|
||||
id="with_image",
|
||||
),
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--tag", "v2"),
|
||||
"--image-uri cannot be combined with --tag.",
|
||||
id="with_tag",
|
||||
),
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--remote"),
|
||||
"--image-uri cannot be combined with --remote.",
|
||||
id="with_remote",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_image_uri_conflicting_flags_are_rejected_before_any_docker_work(
|
||||
deploy_project: DeployProject, args: tuple[str, ...], message: str
|
||||
) -> None:
|
||||
result = deploy_project.run(*args)
|
||||
|
||||
assert result.exit_code != 0
|
||||
assert message in result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.timeline == []
|
||||
|
||||
|
||||
def test_invalid_tag_fails_before_any_control_plane_call(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
|
||||
@@ -13,6 +13,7 @@ import pytest
|
||||
|
||||
import langgraph_cli.deploy as deploy_mod
|
||||
from langgraph_cli.deploy import (
|
||||
BuildAndPush,
|
||||
ById,
|
||||
ByName,
|
||||
CustomerRegistrySource,
|
||||
@@ -21,6 +22,7 @@ from langgraph_cli.deploy import (
|
||||
Listener,
|
||||
ManagedRegistrySource,
|
||||
OnListener,
|
||||
PublishedImage,
|
||||
RemoteBuildSource,
|
||||
RequestedPlacement,
|
||||
Unplaced,
|
||||
@@ -614,6 +616,7 @@ class TestSelectSource:
|
||||
OPTIONS = {
|
||||
"push_to": None,
|
||||
"image": None,
|
||||
"image_uri": None,
|
||||
"image_name": None,
|
||||
"tag": None,
|
||||
"remote_build_flag": None,
|
||||
@@ -629,8 +632,10 @@ class TestSelectSource:
|
||||
{"push_to": REPOSITORY},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="push_to_selects_the_external_source_with_the_default_tag",
|
||||
@@ -639,8 +644,10 @@ class TestSelectSource:
|
||||
{"push_to": f"{REPOSITORY}:v2"},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "v2"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "v2"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="push_to_keeps_a_tag_given_in_the_reference",
|
||||
@@ -649,8 +656,10 @@ class TestSelectSource:
|
||||
{"push_to": REPOSITORY, "tag": "v3"},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "v3"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "v3"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="tag_flag_composes_with_push_to",
|
||||
@@ -659,8 +668,10 @@ class TestSelectSource:
|
||||
{"push_to": REPOSITORY, "image": "app:dev"},
|
||||
False,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image="app:dev",
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image="app:dev",
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="prebuilt_image_is_retagged_for_push_to_without_docker_checks",
|
||||
@@ -672,12 +683,44 @@ class TestSelectSource:
|
||||
},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement("listener-1", "agents"),
|
||||
),
|
||||
id="push_to_carries_the_requested_placement",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": f"{REPOSITORY}@sha256:abc123"},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="image_uri_selects_the_published_image_source_by_digest",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": f" {REPOSITORY}@sha256:abc123 "},
|
||||
False,
|
||||
CustomerRegistrySource(
|
||||
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="image_uri_needs_no_local_docker_and_is_trimmed",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"image_uri": f"{REPOSITORY}@sha256:abc123",
|
||||
"placement": RequestedPlacement("listener-1", "agents"),
|
||||
},
|
||||
False,
|
||||
CustomerRegistrySource(
|
||||
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
|
||||
requested_placement=RequestedPlacement("listener-1", "agents"),
|
||||
),
|
||||
id="image_uri_carries_the_requested_placement",
|
||||
),
|
||||
pytest.param(
|
||||
{"remote_build_flag": True},
|
||||
True,
|
||||
@@ -763,6 +806,46 @@ class TestSelectSource:
|
||||
"only apply when creating a deployment with --push-to",
|
||||
id="namespace_without_push_to",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "push_to": REPOSITORY},
|
||||
"--image-uri cannot be combined with --push-to.",
|
||||
id="image_uri_with_push_to",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "image": "app:dev"},
|
||||
"--image-uri cannot be combined with --image.",
|
||||
id="image_uri_with_image",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "tag": "v2"},
|
||||
"--image-uri cannot be combined with --tag.",
|
||||
id="image_uri_with_tag",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "remote_build_flag": True},
|
||||
"--image-uri cannot be combined with --remote.",
|
||||
id="image_uri_with_remote",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": ""},
|
||||
"--image-uri must not be empty.",
|
||||
id="image_uri_empty",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": " "},
|
||||
"--image-uri must not be empty.",
|
||||
id="image_uri_blank",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": f"{REPOSITORY}:v1.2.3"},
|
||||
"--image-uri must pin a digest",
|
||||
id="image_uri_with_a_mutable_tag",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY},
|
||||
"--image-uri must pin a digest",
|
||||
id="image_uri_without_any_tag_or_digest",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_conflicting_flags_are_rejected(self, monkeypatch, flags, message):
|
||||
@@ -1164,6 +1247,7 @@ def test_a_deployment_id_with_listener_flags_is_refused_without_probing_docker(
|
||||
_select_source(
|
||||
push_to="registry.example.com/app",
|
||||
image=None,
|
||||
image_uri=None,
|
||||
image_name=None,
|
||||
tag=None,
|
||||
remote_build_flag=None,
|
||||
|
||||
@@ -3,6 +3,7 @@ from unittest.mock import patch
|
||||
from langgraph_cli.deploy import (
|
||||
_extract_deployment_url,
|
||||
format_deployments_table,
|
||||
format_listeners_table,
|
||||
format_revisions_table,
|
||||
)
|
||||
from langgraph_cli.util import clean_empty_lines, warn_non_wolfi_distro
|
||||
@@ -255,3 +256,34 @@ def test_format_revisions_table():
|
||||
assert "rev-456" in output
|
||||
assert "rev-789" in output
|
||||
assert "REPLACED" in output
|
||||
|
||||
|
||||
def test_format_listeners_table():
|
||||
output = format_listeners_table(
|
||||
[
|
||||
{
|
||||
"id": "listener-1",
|
||||
"compute_id": "prod-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents"]},
|
||||
},
|
||||
{
|
||||
"id": "listener-2",
|
||||
"compute_id": "multi-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents", "agents-staging"]},
|
||||
},
|
||||
{
|
||||
"id": "listener-3",
|
||||
"compute_id": "broken-cluster",
|
||||
},
|
||||
]
|
||||
)
|
||||
assert "Listener ID" in output
|
||||
assert "Compute ID" in output
|
||||
assert "Namespaces" in output
|
||||
assert "listener-1" in output
|
||||
assert "prod-cluster" in output
|
||||
assert "agents" in output
|
||||
assert "listener-2" in output
|
||||
assert "agents, agents-staging" in output
|
||||
assert "listener-3" in output
|
||||
assert "broken-cluster" in output
|
||||
|
||||
@@ -27,6 +27,7 @@ from langgraph.checkpoint.base import (
|
||||
Checkpoint,
|
||||
PendingWrite,
|
||||
V,
|
||||
writes_sort_key,
|
||||
)
|
||||
from langgraph.store.base import BaseStore
|
||||
from xxhash import xxh3_128_hexdigest
|
||||
@@ -251,9 +252,7 @@ def apply_writes(
|
||||
Set of channels that were updated in this step.
|
||||
"""
|
||||
# sort tasks on path, to ensure deterministic order for update application
|
||||
# any path parts after the 3rd are ignored for sorting
|
||||
# (we use them for eg. task ids which aren't good for sorting)
|
||||
tasks = sorted(tasks, key=lambda t: task_path_str(t.path[:3]))
|
||||
tasks = sorted(tasks, key=lambda t: writes_sort_key(task_path_str(t.path)))
|
||||
# if no task has triggers this is applying writes from the null task only
|
||||
# so we don't do anything other than update the channels written to
|
||||
bump_step = any(t.triggers for t in tasks)
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
@@ -49,6 +50,15 @@ def empty_checkpoint() -> Checkpoint:
|
||||
)
|
||||
|
||||
|
||||
def put_writes_accepts_task_path(put_writes: Callable[..., Any]) -> bool:
|
||||
"""Whether a saver's `put_writes` or `aput_writes` takes `task_path`.
|
||||
|
||||
Savers written before the parameter existed don't, so it is passed only
|
||||
when this is true.
|
||||
"""
|
||||
return signature(put_writes).parameters.get("task_path") is not None
|
||||
|
||||
|
||||
def exit_delta_task_id(step: int, task_id: str) -> str:
|
||||
"""Synthetic task id for exit-mode DeltaChannel writes.
|
||||
|
||||
@@ -60,6 +70,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]],
|
||||
|
||||
@@ -12,7 +12,6 @@ from contextlib import (
|
||||
ExitStack,
|
||||
)
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
Any,
|
||||
@@ -106,7 +105,9 @@ from langgraph.pregel._checkpoint import (
|
||||
delta_channels_to_snapshot,
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
exit_delta_late_task_id,
|
||||
exit_delta_task_id,
|
||||
put_writes_accepts_task_path,
|
||||
)
|
||||
from langgraph.pregel._executor import (
|
||||
AsyncBackgroundExecutor,
|
||||
@@ -190,6 +191,7 @@ class PregelLoop:
|
||||
Callable[
|
||||
[
|
||||
concurrent.futures.Future | None,
|
||||
Sequence[Any],
|
||||
RunnableConfig,
|
||||
Checkpoint,
|
||||
str,
|
||||
@@ -203,11 +205,13 @@ class PregelLoop:
|
||||
submit: Submit
|
||||
channels: Mapping[str, BaseChannel]
|
||||
# Futures from `checkpointer.put_writes` calls that produced delta-channel
|
||||
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
|
||||
# local `futs` then reset to `[]` and wait/gather) before putting the
|
||||
# next checkpoint, so a checkpoint never becomes durable before the
|
||||
# writes that produced it. Initialised to `[]` in both sync and async
|
||||
# `__enter__`; stays `None` only when no checkpointer.
|
||||
# writes. `_put_checkpoint` hands this list to the save it submits, which
|
||||
# waits for them first, so a checkpoint never becomes durable before the
|
||||
# writes that produced it. If a write or the previous save failed, the
|
||||
# save fails too: a DeltaChannel is rebuilt from its writes along the
|
||||
# parent chain, so a checkpoint saved past either gap reads back short
|
||||
# for good. Initialised to `[]` in both sync and async `__enter__`;
|
||||
# stays `None` only when no checkpointer.
|
||||
_delta_write_futs: list[Any] | None = None
|
||||
|
||||
# Same pattern as `_delta_write_futs` but for error-handler writes.
|
||||
@@ -221,10 +225,18 @@ 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 (task_id, channel) pairs whose delta writes are stored on the loaded
|
||||
# checkpoint, which a resume addressed by `checkpoint_id` can rerun; this
|
||||
# run's `Command` delta writes, kept apart from the NULL_TASK_ID writes
|
||||
# loaded with the checkpoint; and the checkpoint's own superstep, the
|
||||
# first one this run ticks.
|
||||
_stored_delta_writes: set[tuple[str, str]]
|
||||
_exit_command_writes: list[tuple[str, Any]]
|
||||
_exit_first_step: int | None = None
|
||||
|
||||
# Delta channels that must snapshot at the next checkpoint, whatever their
|
||||
# cadence counters say:
|
||||
@@ -737,9 +749,25 @@ class PregelLoop:
|
||||
)
|
||||
# capture delta-channel writes for exit-mode accumulator before clearing
|
||||
if self._exit_delta_writes is not None:
|
||||
# On the first tick the pending writes still hold the ones loaded
|
||||
# with the checkpoint, which are already stored on it.
|
||||
first = self._exit_first_step is None
|
||||
if first:
|
||||
self._exit_first_step = self.step
|
||||
self._exit_delta_writes.extend(
|
||||
(self.step, NULL_TASK_ID, "", ch, v)
|
||||
for ch, v in self._exit_command_writes
|
||||
)
|
||||
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 not isinstance(self.specs.get(ch), DeltaChannel):
|
||||
continue
|
||||
if first and (
|
||||
tid == NULL_TASK_ID or (tid, ch) in self._stored_delta_writes
|
||||
):
|
||||
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))
|
||||
# clear pending writes
|
||||
self.checkpoint_pending_writes.clear()
|
||||
# only replay (re-execute) done tasks on the first tick
|
||||
@@ -860,6 +888,12 @@ class PregelLoop:
|
||||
def _first(
|
||||
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
|
||||
) -> set[str] | None:
|
||||
self._stored_delta_writes = {
|
||||
(tid, ch)
|
||||
for tid, ch, _ in self.checkpoint_pending_writes
|
||||
if tid != NULL_TASK_ID and isinstance(self.specs.get(ch), DeltaChannel)
|
||||
}
|
||||
self._exit_command_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)
|
||||
@@ -977,6 +1011,12 @@ class PregelLoop:
|
||||
carried.extend((tid, c, v) for c, v in ws)
|
||||
else:
|
||||
self.put_writes(tid, ws)
|
||||
if self._exit_delta_writes is not None and tid == NULL_TASK_ID:
|
||||
self._exit_command_writes.extend(
|
||||
(c, v)
|
||||
for c, v in ws
|
||||
if isinstance(self.specs.get(c), DeltaChannel)
|
||||
)
|
||||
self._delta_channels_forced_snapshot.update(
|
||||
delta_channels_with_pending_writes(self.specs, carried)
|
||||
)
|
||||
@@ -1074,7 +1114,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":
|
||||
@@ -1260,12 +1302,17 @@ class PregelLoop:
|
||||
)
|
||||
self.checkpoint_previous_versions = channel_versions
|
||||
|
||||
# Take this checkpoint's writes now: saves run in the background
|
||||
# and can start out of order, so a save that took them itself
|
||||
# could get another checkpoint's writes.
|
||||
delta_write_futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
# save it, without blocking
|
||||
# if there's a previous checkpoint save in progress, wait for it
|
||||
# ensuring checkpointers receive checkpoints in order
|
||||
self._put_checkpoint_fut = self.submit(
|
||||
self._checkpointer_put_after_previous,
|
||||
getattr(self, "_put_checkpoint_fut", None),
|
||||
delta_write_futs,
|
||||
self.checkpoint_config,
|
||||
copy_checkpoint(self.checkpoint),
|
||||
self.checkpoint_metadata,
|
||||
@@ -1309,9 +1356,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
|
||||
@@ -1337,6 +1382,7 @@ class PregelLoop:
|
||||
self._put_checkpoint_fut = self.submit(
|
||||
self._checkpointer_put_after_previous,
|
||||
getattr(self, "_put_checkpoint_fut", None),
|
||||
(),
|
||||
stub_put_config,
|
||||
stub_cp,
|
||||
{"step": -2},
|
||||
@@ -1346,11 +1392,21 @@ 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 keeps its real task paths, so it
|
||||
# interleaves with the writes a resume loaded from it. Its task ids stay
|
||||
# synthetic: under the real id, a run whose final checkpoint fails to
|
||||
# save would leave the resumed task looking done to the next resume.
|
||||
# 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 = (exit_delta_task_id(step, 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,
|
||||
{
|
||||
@@ -1360,22 +1416,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)
|
||||
@@ -1584,8 +1639,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
self.checkpointer_get_next_version = checkpointer.get_next_version
|
||||
self.checkpointer_put_writes = checkpointer.put_writes
|
||||
self.checkpointer_put_writes_accepts_task_path = (
|
||||
signature(checkpointer.put_writes).parameters.get("task_path")
|
||||
is not None
|
||||
put_writes_accepts_task_path(checkpointer.put_writes)
|
||||
)
|
||||
else:
|
||||
self.checkpointer_get_next_version = increment
|
||||
@@ -1596,21 +1650,19 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
def _checkpointer_put_after_previous(
|
||||
self,
|
||||
prev: concurrent.futures.Future | None,
|
||||
delta_write_futs: Sequence[concurrent.futures.Future],
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
if self._delta_write_futs:
|
||||
futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
concurrent.futures.wait(futs)
|
||||
try:
|
||||
if prev is not None:
|
||||
prev.result()
|
||||
finally:
|
||||
cast(BaseCheckpointSaver, self.checkpointer).put(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
for fut in delta_write_futs:
|
||||
fut.result()
|
||||
if prev is not None:
|
||||
prev.result()
|
||||
cast(BaseCheckpointSaver, self.checkpointer).put(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
|
||||
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
|
||||
if self.cache is None:
|
||||
@@ -1840,8 +1892,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
self.checkpointer_get_next_version = checkpointer.get_next_version
|
||||
self.checkpointer_put_writes = checkpointer.aput_writes
|
||||
self.checkpointer_put_writes_accepts_task_path = (
|
||||
signature(checkpointer.aput_writes).parameters.get("task_path")
|
||||
is not None
|
||||
put_writes_accepts_task_path(checkpointer.aput_writes)
|
||||
)
|
||||
else:
|
||||
self.checkpointer_get_next_version = increment
|
||||
@@ -1852,23 +1903,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
async def _checkpointer_put_after_previous(
|
||||
self,
|
||||
prev: asyncio.Task | None,
|
||||
delta_write_futs: Sequence[asyncio.Future],
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
# Drain DeltaChannel write futures before committing the checkpoint so
|
||||
# ancestor walks never see a checkpoint without its backing writes.
|
||||
if self._delta_write_futs:
|
||||
futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
await asyncio.gather(*futs)
|
||||
try:
|
||||
if prev is not None:
|
||||
await prev
|
||||
finally:
|
||||
await cast(BaseCheckpointSaver, self.checkpointer).aput(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
if delta_write_futs:
|
||||
await asyncio.gather(*delta_write_futs)
|
||||
if prev is not None:
|
||||
await prev
|
||||
await cast(BaseCheckpointSaver, self.checkpointer).aput(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
|
||||
async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:
|
||||
if self.cache is None:
|
||||
|
||||
@@ -126,6 +126,7 @@ from langgraph.pregel._algo import (
|
||||
apply_writes,
|
||||
local_read,
|
||||
prepare_next_tasks,
|
||||
task_path_str,
|
||||
)
|
||||
from langgraph.pregel._call import identifier
|
||||
from langgraph.pregel._checkpoint import (
|
||||
@@ -139,6 +140,7 @@ from langgraph.pregel._checkpoint import (
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
get_updated_channels_from_tasks,
|
||||
put_writes_accepts_task_path,
|
||||
versions_seen_without_bumps,
|
||||
)
|
||||
from langgraph.pregel._draw import draw_graph
|
||||
@@ -1983,13 +1985,13 @@ class Pregel(
|
||||
run_tasks: list[PregelTaskWrites] = []
|
||||
run_task_ids: list[str] = []
|
||||
|
||||
for as_node, values, provided_task_id in valid_updates:
|
||||
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
|
||||
# create task to run all writers of the chosen node
|
||||
writers = self.nodes[as_node].flat_writers
|
||||
if not writers:
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
|
||||
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
|
||||
# get the task ids that were prepared for this node
|
||||
# if a task id was provided in the StateUpdate, we use it
|
||||
# otherwise, we use the next available task id
|
||||
@@ -1997,7 +1999,7 @@ class Pregel(
|
||||
task_id = provided_task_id or (
|
||||
prepared_task_ids.popleft()
|
||||
if prepared_task_ids
|
||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
||||
else _update_task_id(checkpoint["id"], i)
|
||||
)
|
||||
run_tasks.append(task)
|
||||
run_task_ids.append(task_id)
|
||||
@@ -2050,7 +2052,10 @@ class Pregel(
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
checkpoint_config,
|
||||
channel_writes,
|
||||
task_id,
|
||||
**_task_path_kwarg(checkpointer.put_writes, task),
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
@@ -2471,13 +2476,13 @@ class Pregel(
|
||||
run_tasks: list[PregelTaskWrites] = []
|
||||
run_task_ids: list[str] = []
|
||||
|
||||
for as_node, values, provided_task_id in valid_updates:
|
||||
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
|
||||
# create task to run all writers of the chosen node
|
||||
writers = self.nodes[as_node].flat_writers
|
||||
if not writers:
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
|
||||
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
|
||||
# get the task ids that were prepared for this node
|
||||
# if a task id was provided in the StateUpdate, we use it
|
||||
# otherwise, we use the next available task id
|
||||
@@ -2485,7 +2490,7 @@ class Pregel(
|
||||
task_id = provided_task_id or (
|
||||
prepared_task_ids.popleft()
|
||||
if prepared_task_ids
|
||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
||||
else _update_task_id(checkpoint["id"], i)
|
||||
)
|
||||
run_tasks.append(task)
|
||||
run_task_ids.append(task_id)
|
||||
@@ -2538,7 +2543,10 @@ class Pregel(
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
await checkpointer.aput_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
checkpoint_config,
|
||||
channel_writes,
|
||||
task_id,
|
||||
**_task_path_kwarg(checkpointer.aput_writes, task),
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
@@ -3706,16 +3714,15 @@ class Pregel(
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
version: Literal["v1", "v2", "v3"] = "v2",
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
control: RunControl | None = None,
|
||||
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""Stream events from this graph.
|
||||
|
||||
For `version="v1"` / `"v2"`, yields `StreamEvent` dicts (see
|
||||
`Runnable.stream_events`). For `version="v3"`, returns a
|
||||
For `version="v1"` / `"v2"`, delegates to
|
||||
`Runnable.(a)stream_events`; synchronous v1/v2 event streaming is
|
||||
not implemented in langchain-core, so use `astream_events` for
|
||||
those versions. For `version="v3"`, returns a
|
||||
`GraphRunStream` whose typed projections the caller drives by
|
||||
iterating — no background thread.
|
||||
|
||||
@@ -3745,20 +3752,27 @@ class Pregel(
|
||||
config: Optional runnable config.
|
||||
version: Streaming-event schema version. `"v3"` selects the
|
||||
content-block-centric streaming protocol.
|
||||
interrupt_before: Nodes to interrupt before, if any. Only
|
||||
used for `version="v3"`.
|
||||
interrupt_after: Nodes to interrupt after, if any. Only
|
||||
used for `version="v3"`.
|
||||
interrupt_before: Nodes to interrupt before, if any.
|
||||
Honored on every version that can run; type-checked
|
||||
only on the `version="v3"` overloads.
|
||||
interrupt_after: Nodes to interrupt after, if any. Honored
|
||||
on every version that can run; type-checked only on
|
||||
the `version="v3"` overloads.
|
||||
control: Optional run control used to request cooperative
|
||||
drain. Only used for `version="v3"`.
|
||||
drain. Honored on every version that can run;
|
||||
type-checked only on the `version="v3"` overloads.
|
||||
transformers: Extra transformer classes or configured
|
||||
factories appended after compile-time
|
||||
`stream_transformers`. Factories are called as
|
||||
`factory(scope)` so they can propagate to subgraph
|
||||
scopes. Only used for `version="v3"`.
|
||||
**kwargs: For `version="v1"`/`"v2"`, forwarded to
|
||||
`Runnable.stream_events`. For `version="v3"`, forwarded
|
||||
to the underlying `stream(...)` call (e.g. `context`,
|
||||
**kwargs: For `version="v1"`/`"v2"` on `astream_events`,
|
||||
forwarded to `Runnable.astream_events`, which passes
|
||||
them through to `astream` — so execution kwargs such as
|
||||
`context`, `durability`, `interrupt_before`,
|
||||
`interrupt_after` and `control` are honored on every
|
||||
version that can run. For `version="v3"`, forwarded to the
|
||||
underlying `stream(...)` call (e.g. `context`,
|
||||
`durability`, `output_keys`, `print_mode`, `debug`).
|
||||
`stream_mode` and `subgraphs` are not accepted under
|
||||
`version="v3"` and raise `TypeError` if supplied; v3
|
||||
@@ -3766,16 +3780,15 @@ class Pregel(
|
||||
|
||||
Returns:
|
||||
For `version="v3"`, a `GraphRunStream` the caller iterates
|
||||
to drive the run. Otherwise an `Iterator[StreamEvent]`.
|
||||
to drive the run. For `version="v1"`/`"v2"`,
|
||||
`astream_events` yields `StreamEvent` dicts; the synchronous
|
||||
v1/v2 path is not implemented in langchain-core.
|
||||
"""
|
||||
if version == "v3":
|
||||
_reject_v3_invariant_kwargs(kwargs)
|
||||
return self._pregel_stream_v3(
|
||||
input,
|
||||
config,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
control=control,
|
||||
transformers=transformers,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -3811,9 +3824,6 @@ class Pregel(
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
version: Literal["v1", "v2", "v3"] = "v2",
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
control: RunControl | None = None,
|
||||
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[StreamEvent] | Awaitable[Any]:
|
||||
@@ -3836,9 +3846,6 @@ class Pregel(
|
||||
return self._apregel_stream_v3(
|
||||
input,
|
||||
config,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
control=control,
|
||||
transformers=transformers,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -4237,6 +4244,27 @@ class Pregel(
|
||||
await self.cache.aclear(namespaces)
|
||||
|
||||
|
||||
def _task_path_kwarg(put_writes: Callable[..., Any], task: PregelTaskWrites) -> dict:
|
||||
"""Pass the task's path to savers whose `put_writes` takes one.
|
||||
|
||||
Savers that replay a checkpoint's writes in task path order then give back
|
||||
updates applied together in the order they were given.
|
||||
"""
|
||||
if not put_writes_accepts_task_path(put_writes):
|
||||
return {}
|
||||
return {"task_path": task_path_str(task.path)}
|
||||
|
||||
|
||||
def _update_task_id(checkpoint_id: str, i: int) -> str:
|
||||
"""Task id for the `i`th update of a superstep that has no task to reuse.
|
||||
|
||||
Savers keep one write per `(task_id, idx)`, so updates sharing an id lose
|
||||
all but the first one's writes, which a `DeltaChannel` replays from. The
|
||||
first update keeps the id a lone update has always had.
|
||||
"""
|
||||
return str(uuid5(UUID(checkpoint_id), INTERRUPT if i == 0 else f"{INTERRUPT}:{i}"))
|
||||
|
||||
|
||||
def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str]]:
|
||||
"""Index from a trigger to nodes that depend on it."""
|
||||
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
|
||||
|
||||
@@ -25,7 +25,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"langchain-core>=1.4.7,<2",
|
||||
"langgraph-checkpoint>=4.1.0,<5.0.0",
|
||||
"langgraph-checkpoint>=4.3.0,<5.0.0",
|
||||
"langgraph-sdk>=0.4.6,<0.5.0",
|
||||
"langgraph-prebuilt>=1.1.0,<1.2.0",
|
||||
"xxhash>=3.5.0",
|
||||
|
||||
@@ -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,197 @@ 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_with_a_command_update_replays_its_write_once(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("done", lambda state: _both("done"))
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
|
||||
graph.invoke(
|
||||
Command(resume="yes", update=_both("cmd")), config, durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "cmd", "ask", "done"]
|
||||
|
||||
|
||||
class _FlagState(_ResumeState, total=False):
|
||||
extra: Annotated[list, DeltaChannel(_append)]
|
||||
flag: bool
|
||||
|
||||
|
||||
def test_addressed_resume_keeps_a_rerun_tasks_write_to_a_new_channel(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
def done(state: _FlagState) -> dict:
|
||||
return {**_both("done"), **({"extra": ["new"]} if state.get("flag") else {})}
|
||||
|
||||
builder = StateGraph(_FlagState)
|
||||
builder.add_node("done", done)
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
head = graph.get_state(config).config
|
||||
|
||||
live = graph.invoke(
|
||||
Command(resume="yes", update={"flag": True}), head, durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert live["extra"] == state.values["extra"] == ["new"]
|
||||
assert state.values["log"] == state.values["plain"]
|
||||
|
||||
|
||||
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"]
|
||||
|
||||
|
||||
class _FailingPutSaver(InMemorySaver):
|
||||
fail = False
|
||||
|
||||
def put(
|
||||
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
|
||||
) -> Any:
|
||||
if self.fail:
|
||||
raise RuntimeError("final checkpoint lost")
|
||||
return super().put(config, checkpoint, metadata, new_versions)
|
||||
|
||||
|
||||
def test_exit_resume_retried_after_its_final_checkpoint_fails_reruns_the_resumed_task() -> (
|
||||
None
|
||||
):
|
||||
saver = _FailingPutSaver()
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("done", lambda state: _both("done"))
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability="exit")
|
||||
saver.fail = True
|
||||
with pytest.raises(RuntimeError, match="final checkpoint lost"):
|
||||
graph.invoke(Command(resume="yes"), config, durability="exit")
|
||||
saver.fail = False
|
||||
|
||||
graph.invoke(Command(resume="yes"), config, durability="exit")
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"]
|
||||
assert sorted(state.values["plain"]) == ["ask", "done", "in"]
|
||||
|
||||
@@ -740,20 +740,6 @@ def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
|
||||
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"durability",
|
||||
[
|
||||
"sync",
|
||||
"async",
|
||||
pytest.param(
|
||||
"exit",
|
||||
marks=pytest.mark.xfail(
|
||||
reason="exit durability stores a resumed run's loaded writes twice",
|
||||
strict=True,
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
"""`DeltaChannel` replay must apply parallel writes in the order `invoke` did."""
|
||||
|
||||
import asyncio
|
||||
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
|
||||
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:
|
||||
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)
|
||||
|
||||
|
||||
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:
|
||||
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_sync_get_state_on_async_saver_matches_live_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 asyncio.to_thread(graph.get_state, config)).values["items"]
|
||||
|
||||
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}"
|
||||
@@ -20,6 +20,7 @@ from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import 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
|
||||
@@ -27,16 +28,17 @@ from typing_extensions import TypedDict
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
from langgraph.types import StateUpdate
|
||||
from langgraph.types import StateSnapshot, StateUpdate
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _build_graph(
|
||||
checkpointer: InMemorySaver,
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
*,
|
||||
two_nodes: bool = False,
|
||||
snapshot_frequency: int = 1000,
|
||||
interrupt_before: list[str] | None = None,
|
||||
) -> Any:
|
||||
"""Compile a minimal DeltaChannel-backed `messages` graph.
|
||||
|
||||
@@ -63,7 +65,7 @@ def _build_graph(
|
||||
builder.set_finish_point("assistant")
|
||||
else:
|
||||
builder.set_finish_point("model")
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
return builder.compile(checkpointer=checkpointer, interrupt_before=interrupt_before)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -304,15 +306,14 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
||||
that each call `put_writes`. Guards the regression where moving
|
||||
`put_writes` outside the per-task loop would persist only the last
|
||||
task's writes.
|
||||
|
||||
Explicit `task_id`s are required to disambiguate writes belonging to
|
||||
different `StateUpdate`s targeting the same node — otherwise both share
|
||||
the deterministic interrupt-derived id and collide in the saver.
|
||||
"""
|
||||
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": "bulk-multi-task"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
base = saver.get_tuple(config)
|
||||
assert base is not None
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
@@ -332,13 +333,155 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
||||
],
|
||||
)
|
||||
|
||||
stored = saver.get_tuple(base.config)
|
||||
assert stored is not None
|
||||
assert {task_id for task_id, _, _ in stored.pending_writes or []} == {
|
||||
"task-1",
|
||||
"task-2",
|
||||
}, "explicit task ids must key the stored writes"
|
||||
state = graph.get_state(config)
|
||||
contents = [m.content for m in state.values["messages"]]
|
||||
ids = [m.id for m in state.values["messages"]]
|
||||
assert sorted(contents) == ["first", "second"], (
|
||||
assert sorted(contents) == ["first", "hi", "second"], (
|
||||
f"both updates' writes must persist; got {contents}"
|
||||
)
|
||||
assert sorted(ids) == ["m1", "m2"]
|
||||
assert sorted(ids) == ["hi", "m1", "m2"]
|
||||
|
||||
|
||||
def _update(content: str, as_node: str) -> StateUpdate:
|
||||
return StateUpdate(
|
||||
values={"messages": [HumanMessage(content=content, id=content)]},
|
||||
as_node=as_node,
|
||||
)
|
||||
|
||||
|
||||
def _contents(state: StateSnapshot) -> list[str]:
|
||||
return [m.content for m in state.values["messages"]]
|
||||
|
||||
|
||||
def test_bulk_update_state_keeps_every_update_without_task_ids(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_graph(sync_checkpointer, two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
[
|
||||
[
|
||||
_update("first", "model"),
|
||||
_update("second", "model"),
|
||||
_update("third", "assistant"),
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
contents = _contents(graph.get_state(config))
|
||||
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||
f"every update's writes must persist; got {contents}"
|
||||
)
|
||||
|
||||
|
||||
async def test_abulk_update_state_keeps_every_update_without_task_ids(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_graph(async_checkpointer, two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
await graph.abulk_update_state(
|
||||
config,
|
||||
[
|
||||
[
|
||||
_update("first", "model"),
|
||||
_update("second", "model"),
|
||||
_update("third", "assistant"),
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
contents = _contents(await graph.aget_state(config))
|
||||
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||
f"every update's writes must persist; got {contents}"
|
||||
)
|
||||
|
||||
|
||||
def test_bulk_update_state_keeps_every_update_next_to_a_pending_task(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_graph(
|
||||
sync_checkpointer, two_nodes=True, interrupt_before=["assistant"]
|
||||
)
|
||||
config = {"configurable": {"thread_id": "bulk-pending-task"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
assert graph.get_state(config).next == ("assistant",)
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
[
|
||||
[
|
||||
_update("first", "assistant"),
|
||||
_update("second", "model"),
|
||||
_update("third", "model"),
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
contents = _contents(graph.get_state(config))
|
||||
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||
f"every update's writes must persist; got {contents}"
|
||||
)
|
||||
|
||||
|
||||
class _TaskPathOrderSaver(InMemorySaver):
|
||||
"""Replays each checkpoint's writes by `(task_path, task_id, idx)`."""
|
||||
|
||||
def get_tuple(self, config: Any) -> Any:
|
||||
tup = super().get_tuple(config)
|
||||
if tup is None or not tup.pending_writes:
|
||||
return tup
|
||||
conf = tup.config["configurable"]
|
||||
stored = self.writes[
|
||||
(conf["thread_id"], conf["checkpoint_ns"], conf["checkpoint_id"])
|
||||
]
|
||||
rows = sorted(
|
||||
zip(stored.items(), tup.pending_writes),
|
||||
key=lambda row: (row[0][1][3], *row[0][0]),
|
||||
)
|
||||
return tup._replace(pending_writes=[write for _, write in rows])
|
||||
|
||||
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
|
||||
aget_delta_channel_history = BaseCheckpointSaver.aget_delta_channel_history
|
||||
|
||||
|
||||
GIVEN = ["u1", "u2", "u3", "u4", "u5", "u6"]
|
||||
|
||||
|
||||
def _updates_in_given_order() -> list[list[StateUpdate]]:
|
||||
return [
|
||||
[_update(c, "assistant" if i % 2 else "model") for i, c in enumerate(GIVEN)]
|
||||
]
|
||||
|
||||
|
||||
def test_bulk_update_state_replays_updates_in_the_order_given() -> None:
|
||||
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-order"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
graph.bulk_update_state(config, _updates_in_given_order())
|
||||
|
||||
assert _contents(graph.get_state(config)) == ["hi", *GIVEN]
|
||||
|
||||
|
||||
async def test_abulk_update_state_replays_updates_in_the_order_given() -> None:
|
||||
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-order"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
await graph.abulk_update_state(config, _updates_in_given_order())
|
||||
|
||||
assert _contents(await graph.aget_state(config)) == ["hi", *GIVEN]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
"""A checkpoint must never be saved without the `DeltaChannel` writes it reads."""
|
||||
|
||||
import operator
|
||||
import threading
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.types import Durability
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
INPUT = {"log": [], "plain": []}
|
||||
FINAL = {"log": ["a", "b", "c"], "plain": ["a", "b", "c"]}
|
||||
|
||||
# Exit mode saves nothing before the failed write, so its retry starts over.
|
||||
RETRIES = [
|
||||
pytest.param("sync", None, id="sync"),
|
||||
pytest.param("async", None, id="async"),
|
||||
pytest.param("exit", INPUT, id="exit"),
|
||||
]
|
||||
|
||||
|
||||
def _append(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
log: Annotated[list, DeltaChannel(_append)]
|
||||
plain: Annotated[list, operator.add]
|
||||
|
||||
|
||||
class _FailsTheWriteOfBOnce(InMemorySaver):
|
||||
failed = False
|
||||
|
||||
def _fail_once(self, writes: Any) -> None:
|
||||
if not self.failed and ("log", ["b"]) in writes:
|
||||
self.failed = True
|
||||
raise ConnectionError("b's write was not saved")
|
||||
|
||||
def put_writes(
|
||||
self, config: Any, writes: Any, task_id: str, task_path: str = ""
|
||||
) -> None:
|
||||
self._fail_once(writes)
|
||||
super().put_writes(config, writes, task_id, task_path)
|
||||
|
||||
async def aput_writes(
|
||||
self, config: Any, writes: Any, task_id: str, task_path: str = ""
|
||||
) -> None:
|
||||
self._fail_once(writes)
|
||||
await super().aput_writes(config, writes, task_id, task_path)
|
||||
|
||||
|
||||
def _a_then_b_then_c(saver: InMemorySaver) -> Any:
|
||||
builder = StateGraph(_State)
|
||||
for name in "abc":
|
||||
builder.add_node(
|
||||
name, lambda state, name=name: {"log": [name], "plain": [name]}
|
||||
)
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("b", "c")
|
||||
return builder.compile(checkpointer=saver)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
|
||||
def test_a_failed_delta_write_is_rerun_not_lost(
|
||||
durability: Durability, retry_input: dict | None
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
graph.invoke(INPUT, config, durability=durability)
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
graph.invoke(retry_input, config, durability=durability)
|
||||
assert graph.get_state(config).values == FINAL
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
|
||||
async def test_a_failed_delta_write_is_rerun_not_lost_async(
|
||||
durability: Durability, retry_input: dict | None
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
await graph.ainvoke(INPUT, config, durability=durability)
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
await graph.ainvoke(retry_input, config, durability=durability)
|
||||
assert (await graph.aget_state(config)).values == FINAL
|
||||
|
||||
|
||||
def test_a_delta_graph_finishes_on_a_single_background_thread() -> None:
|
||||
graph = _a_then_b_then_c(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}, "max_concurrency": 1}
|
||||
result: dict = {}
|
||||
run = threading.Thread(
|
||||
target=lambda: result.update(graph.invoke(INPUT, config, durability="async")),
|
||||
daemon=True,
|
||||
)
|
||||
|
||||
run.start()
|
||||
run.join(timeout=10)
|
||||
|
||||
assert not run.is_alive(), "invoke hung"
|
||||
assert result == FINAL
|
||||
@@ -8,20 +8,33 @@ caller kwargs to the inner ``(a)stream`` call but rejects ``stream_mode`` and
|
||||
``subgraphs`` since v3 owns them (``stream_mode`` is built from the
|
||||
transformer mux; ``subgraphs`` is forced True so nested namespaces flow
|
||||
through scoped muxes).
|
||||
|
||||
A second regression is pinned here: #7677 (first released in 1.2.0a3)
|
||||
declared `interrupt_before` / `interrupt_after` / `control` as named
|
||||
parameters on the `Pregel.stream_events` / `astream_events` dispatchers
|
||||
but forwarded them only on the v3 branch, silently dropping them on v1/v2
|
||||
(where they had reached `(a)stream` through `**kwargs` before).
|
||||
`TestAstreamKwargsForwardedOnEveryVersion` and friends pin that they
|
||||
reach `(a)stream` on every version, with the v1/v2 passthrough semantics
|
||||
restored (exactly what the caller passed, including an explicit `None`)
|
||||
and v3's explicit-default semantics preserved.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.constants import END, START
|
||||
from langgraph.errors import GraphDrained
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.runtime import Runtime
|
||||
from langgraph.runtime import RunControl, Runtime
|
||||
|
||||
NEEDS_CONTEXTVARS = pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
@@ -106,3 +119,274 @@ class TestKwargForwardingAsync:
|
||||
version="v3",
|
||||
subgraphs=False,
|
||||
)
|
||||
|
||||
|
||||
_KWARG_NAMES = ("control", "interrupt_before", "interrupt_after")
|
||||
|
||||
|
||||
def _build_two_step_graph(
|
||||
first: Callable[[_State], dict[str, Any]] | None = None,
|
||||
) -> Any:
|
||||
"""A `first -> second` graph with a checkpointer, for interrupt/drain tests."""
|
||||
|
||||
def default_first(state: _State) -> dict[str, Any]:
|
||||
return {"message": state["message"] + " first"}
|
||||
|
||||
def second(state: _State) -> dict[str, Any]:
|
||||
return {"message": state["message"] + " second"}
|
||||
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("first", first or default_first)
|
||||
builder.add_node("second", second)
|
||||
builder.add_edge(START, "first")
|
||||
builder.add_edge("first", "second")
|
||||
builder.add_edge("second", END)
|
||||
return builder.compile(checkpointer=InMemorySaver())
|
||||
|
||||
|
||||
async def _drive_astream_events(
|
||||
graph: Any, config: dict[str, Any], version: str, **kwargs: Any
|
||||
) -> None:
|
||||
"""Consume an astream_events run for `version` to completion."""
|
||||
if version == "v3":
|
||||
run = await graph.astream_events(
|
||||
{"message": "hi"}, config, version="v3", **kwargs
|
||||
)
|
||||
await run.output()
|
||||
else:
|
||||
async for _ in graph.astream_events(
|
||||
{"message": "hi"}, config, version=version, **kwargs
|
||||
):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
|
||||
@pytest.mark.parametrize("version", ["v1", "v2", "v3"])
|
||||
class TestAstreamKwargsForwardedOnEveryVersion:
|
||||
"""`interrupt_before`/`interrupt_after`/`control` reach `astream` on every
|
||||
version.
|
||||
|
||||
Regression test for #7677 (first released in 1.2.0a3): the dispatchers
|
||||
captured these parameters as named arguments but forwarded them only on
|
||||
the v3 branch, silently dropping them on v1/v2.
|
||||
"""
|
||||
|
||||
async def test_interrupt_before(self, version: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "ib"}}
|
||||
await _drive_astream_events(graph, config, version, interrupt_before=["second"])
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
async def test_interrupt_after(self, version: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "ia"}}
|
||||
await _drive_astream_events(graph, config, version, interrupt_after=["first"])
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
async def test_pre_drained_control(self, version: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "drain"}}
|
||||
control = RunControl()
|
||||
control.request_drain("sigterm")
|
||||
with pytest.raises(GraphDrained, match="sigterm"):
|
||||
await _drive_astream_events(graph, config, version, control=control)
|
||||
|
||||
|
||||
class TestStreamEventsV3SyncInterrupts:
|
||||
"""Sync v3 static interrupts reach `stream()` after the kwargs rewire."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("kwarg", "node"),
|
||||
[("interrupt_before", "second"), ("interrupt_after", "first")],
|
||||
)
|
||||
def test_static_interrupt(self, kwarg: str, node: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "sync"}}
|
||||
run = graph.stream_events(
|
||||
{"message": "hi"}, config, version="v3", **{kwarg: [node]}
|
||||
)
|
||||
list(run.values)
|
||||
state = graph.get_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
|
||||
@pytest.mark.parametrize("version", ["v1", "v2", "v3"])
|
||||
class TestAstreamMidRunDrain:
|
||||
"""A drain requested from inside a node propagates out of `astream_events`.
|
||||
|
||||
This is the graceful-shutdown scenario: `request_drain()` called while
|
||||
the run is in flight (e.g. from a signal handler), with v1/v2 running
|
||||
inside core's event-stream task. The caller's own `RunControl` is used
|
||||
(the drain reason proves identity), `GraphDrained` propagates to the
|
||||
consumer, and the checkpoint keeps the pending step.
|
||||
"""
|
||||
|
||||
async def test_drain_requested_inside_first_node(self, version: str) -> None:
|
||||
control = RunControl()
|
||||
|
||||
def first(state: _State) -> dict[str, Any]:
|
||||
control.request_drain("sigterm-mid")
|
||||
return {"message": state["message"] + " first"}
|
||||
|
||||
graph = _build_two_step_graph(first)
|
||||
config = {"configurable": {"thread_id": "midrun"}}
|
||||
|
||||
with pytest.raises(GraphDrained, match="sigterm-mid"):
|
||||
await _drive_astream_events(graph, config, version, control=control)
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
|
||||
def _record_astream_kwargs(graph: Any) -> list[dict[str, Any]]:
|
||||
"""Patch `graph.astream` to record the kwargs each call receives."""
|
||||
received: list[dict[str, Any]] = []
|
||||
original = graph.astream
|
||||
|
||||
async def recording_astream(
|
||||
input: Any, config: Any = None, **kwargs: Any
|
||||
) -> AsyncIterator[Any]:
|
||||
received.append(kwargs)
|
||||
async for chunk in original(input, config, **kwargs):
|
||||
yield chunk
|
||||
|
||||
graph.astream = recording_astream # type: ignore[method-assign]
|
||||
return received
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
|
||||
@pytest.mark.parametrize("version", ["v1", "v2"])
|
||||
class TestAstreamV1V2KwargsPassthrough:
|
||||
"""v1/v2 forward to `astream` exactly what the caller passed.
|
||||
|
||||
Pre-#7677 semantics: an explicit `None` is forwarded as `None`, and an
|
||||
omitted argument is not forwarded at all (so an override's own default
|
||||
would apply).
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_passed_values_reach_astream(self, version: str, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
value: Any = RunControl() if name == "control" else ["second"]
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "rec"}}, version, **{name: value}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] == value
|
||||
if name == "control":
|
||||
assert received[0][name] is value
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_explicit_none_is_forwarded(self, version: str, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "rec-none"}}, version, **{name: None}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] is None
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_omitted_values_are_absent(self, version: str, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "rec-omit"}}, version
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert name not in received[0]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
class TestAstreamV3KwargsDefaults:
|
||||
"""v3 keeps its since-inception explicit-default semantics (#7519).
|
||||
|
||||
Omitted `interrupt_before`/`interrupt_after`/`control` are supplied to
|
||||
`astream` as `None`; passed values are forwarded as-is.
|
||||
"""
|
||||
|
||||
async def test_omitted_values_arrive_as_none(self) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "v3-rec"}}, "v3"
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0]["control"] is None
|
||||
assert received[0]["interrupt_before"] is None
|
||||
assert received[0]["interrupt_after"] is None
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_passed_values_reach_astream(self, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
value: Any = RunControl() if name == "control" else ["second"]
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "v3-rec-2"}}, "v3", **{name: value}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] == value
|
||||
if name == "control":
|
||||
assert received[0][name] is value
|
||||
|
||||
|
||||
def _record_stream_kwargs(graph: Any) -> list[dict[str, Any]]:
|
||||
"""Patch `graph.stream` to record the kwargs each call receives."""
|
||||
received: list[dict[str, Any]] = []
|
||||
original = graph.stream
|
||||
|
||||
def recording_stream(input: Any, config: Any = None, **kwargs: Any) -> Any:
|
||||
received.append(kwargs)
|
||||
yield from original(input, config, **kwargs)
|
||||
|
||||
graph.stream = recording_stream # type: ignore[method-assign]
|
||||
return received
|
||||
|
||||
|
||||
def _drive_stream_events_v3(graph: Any, config: dict[str, Any], **kwargs: Any) -> None:
|
||||
"""Consume a sync v3 stream_events run to completion."""
|
||||
run = graph.stream_events({"message": "hi"}, config, version="v3", **kwargs)
|
||||
list(run.values)
|
||||
|
||||
|
||||
class TestStreamEventsV3SyncKwargsDefaults:
|
||||
"""Sync v3 keeps its since-inception explicit-default semantics (#7519).
|
||||
|
||||
Mirror of `TestAstreamV3KwargsDefaults`: omitted
|
||||
`interrupt_before`/`interrupt_after`/`control` are supplied to `stream`
|
||||
as `None`; passed values are forwarded as-is. Pins the sync helper
|
||||
against a kwargs-only "simplification" that would change subclass
|
||||
default handling.
|
||||
"""
|
||||
|
||||
def test_omitted_values_arrive_as_none(self) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_stream_kwargs(graph)
|
||||
_drive_stream_events_v3(graph, {"configurable": {"thread_id": "s-rec"}})
|
||||
assert len(received) == 1
|
||||
assert received[0]["control"] is None
|
||||
assert received[0]["interrupt_before"] is None
|
||||
assert received[0]["interrupt_after"] is None
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
def test_passed_values_reach_stream(self, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_stream_kwargs(graph)
|
||||
value: Any = RunControl() if name == "control" else ["second"]
|
||||
_drive_stream_events_v3(
|
||||
graph, {"configurable": {"thread_id": "s-rec-2"}}, **{name: value}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] == value
|
||||
if name == "control":
|
||||
assert received[0][name] is value
|
||||
|
||||
Generated
+23
-8
@@ -1217,6 +1217,20 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/38/64/285f20a31679bf547b75602702f7800e74dbabae36ef324f716c02804753/jupyter-1.1.1-py2.py3-none-any.whl", hash = "sha256:7a59533c22af65439b24bbe60373a4e95af8f16ac65a6c00820ad378e3f7cc83", size = 2657, upload-time = "2024-08-30T07:15:47.045Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jupyter-builder"
|
||||
version = "1.2.3"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "jupyter-core" },
|
||||
{ name = "tomli", marker = "python_full_version < '3.11'" },
|
||||
{ name = "traitlets" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/75/3e/56f593e6a664cd14d441724654ce8243f3cac7d3e45867cf084749c8bc75/jupyter_builder-1.2.3.tar.gz", hash = "sha256:01aba6794eb9b19e0e29ae21137ca60ba4135c70347d4b9f664a586822b8c809", size = 1024218, upload-time = "2026-09-04T18:54:09.411Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/84/aa/be79e87c50698673f196d0633fc1664607da2289f9688d1de7a1c9db9e71/jupyter_builder-1.2.3-py3-none-any.whl", hash = "sha256:c5ea5a7190c2a7b082494abade98eece1b2b5bd5dbd7d610606cbcd10a1d08b3", size = 947541, upload-time = "2026-09-04T18:54:07.644Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jupyter-client"
|
||||
version = "8.8.0"
|
||||
@@ -1342,28 +1356,28 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "jupyterlab"
|
||||
version = "4.5.10"
|
||||
version = "4.6.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "async-lru" },
|
||||
{ name = "httpx" },
|
||||
{ name = "ipykernel" },
|
||||
{ name = "jinja2" },
|
||||
{ name = "jupyter-builder" },
|
||||
{ name = "jupyter-core" },
|
||||
{ name = "jupyter-lsp" },
|
||||
{ name = "jupyter-server" },
|
||||
{ name = "jupyterlab-server" },
|
||||
{ name = "notebook-shim" },
|
||||
{ name = "packaging" },
|
||||
{ name = "setuptools" },
|
||||
{ name = "tomli", marker = "python_full_version < '3.11'" },
|
||||
{ name = "tornado" },
|
||||
{ name = "traitlets" },
|
||||
{ name = "typing-extensions", marker = "python_full_version < '3.12'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/74/24/621aa20ec0d2fe72f52095bda0fc1be7738ac21aabe4129ff623140d5cdf/jupyterlab-4.5.10.tar.gz", hash = "sha256:77e8d80b78be59b2eaba2154562e21caa6e79c2f1281d6f486584f7144ee2f47", size = 23998879, upload-time = "2026-07-21T12:43:27.324Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/33/8d/995cc142f6083346b35e7d3eadc5a6717ee89ce64ea53577010ac493bc3c/jupyterlab-4.6.4.tar.gz", hash = "sha256:404f49b081819378524886c9db66dba57a5565981eff885830df1baba3a17df5", size = 28335647, upload-time = "2026-09-21T15:25:18.779Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/9f/c9/940f95f17ee4e413ad252bf8d4f2ee9a341f18cfeda87775fef3d7847321/jupyterlab-4.5.10-py3-none-any.whl", hash = "sha256:5967ca61e692e67a2f30b5a2b901c941dc6ce56c0b0e357bc6d34fed5ec095f6", size = 12452502, upload-time = "2026-07-21T12:43:23.542Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/3b/e4/072bc0d3c6d45d414062a06af7ee81c02a2c12d769c358ec9919f9995ffd/jupyterlab-4.6.4-py3-none-any.whl", hash = "sha256:15b13f991d3985129c797eb84d9949eeb8b6615e14b444868e642411f2c418b2", size = 17171596, upload-time = "2026-09-21T15:25:14.123Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -1623,7 +1637,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -2146,18 +2160,19 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "notebook"
|
||||
version = "7.5.7"
|
||||
version = "7.6.3"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "jupyter-builder" },
|
||||
{ name = "jupyter-server" },
|
||||
{ name = "jupyterlab" },
|
||||
{ name = "jupyterlab-server" },
|
||||
{ name = "notebook-shim" },
|
||||
{ name = "tornado" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/3e/c4/f71f8716f2903e9e817a47f534b9fd84831e155e2acb32c26691c8e06243/notebook-7.5.7.tar.gz", hash = "sha256:d6d59288a25303b25e1dcb71e9b017ec3a785f7d92f38b9bc288ca1970d5b0a8", size = 14171612, upload-time = "2026-06-04T18:33:45.224Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/42/ac/aedb759dc683dcd129f6c75271bab5eb23020b0f7969163c726343718e1e/notebook-7.6.3.tar.gz", hash = "sha256:e2c08e469c0ae20bb0b3214f0ab77e79653317a2f8e5b34c10361c66874a5b50", size = 5501618, upload-time = "2026-09-21T18:07:24.261Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/e1/4d/b3347f7073a377273531efe4ffc738fc910e93718fd2838c7ebf6736c6af/notebook-7.5.7-py3-none-any.whl", hash = "sha256:1f95f79d117e47d20b5555b5c85a397d2cfecf136978aaab767cf0314b09165b", size = 14583767, upload-time = "2026-06-04T18:33:40.987Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/a1/f7/907b98438cf00bdc4e296570058f29a347214bb42338a48e3b7eae58ae2a/notebook-7.6.3-py3-none-any.whl", hash = "sha256:ad7e0eb765fba836cd4a2ab0c7a3a26cde1d91665fbf6f533b6ae7b2de6d88d2", size = 5548452, upload-time = "2026-09-21T18:07:21.563Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
Generated
+1
-1
@@ -370,7 +370,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -22,18 +22,23 @@ RUN pip install --no-cache-dir \
|
||||
"langchain>=1.3.0" \
|
||||
"deepagents>=0.6.2"
|
||||
|
||||
# Swap the published langgraph *core* for this monorepo's local copy, so the
|
||||
# server executes the langgraph under test rather than the latest release.
|
||||
# Swap the published langgraph *core* and *checkpoint* for this monorepo's
|
||||
# local copies, so the server executes the langgraph under test rather than the
|
||||
# latest release.
|
||||
# We keep the rest of the base image (langgraph-api, runtime, Go core server)
|
||||
# on `latest` — that still surfaces upstream regressions — but the core the
|
||||
# server runs is now the PR's, which is what lets this suite catch core
|
||||
# regressions (e.g. ensure_config / runtime changes) before they're published.
|
||||
# `--no-deps` keeps the base image's already-compatible
|
||||
# checkpoint/prebuilt/sdk; we only replace core. The local source comes from
|
||||
# the `langgraph_src` additional build context (see docker-compose.yml).
|
||||
# Checkpoint goes with core because core can need a checkpoint API from the
|
||||
# same release. `--no-deps` keeps the base image's already-compatible
|
||||
# prebuilt/sdk. The local sources come from the `checkpoint_src` and
|
||||
# `langgraph_src` additional build contexts (see docker-compose.yml).
|
||||
COPY --from=checkpoint_src pyproject.toml README.md LICENSE /opt/checkpoint-src/
|
||||
COPY --from=checkpoint_src langgraph /opt/checkpoint-src/langgraph/
|
||||
COPY --from=langgraph_src pyproject.toml README.md LICENSE /opt/langgraph-src/
|
||||
COPY --from=langgraph_src langgraph /opt/langgraph-src/langgraph/
|
||||
RUN pip install --no-cache-dir --force-reinstall --no-deps /opt/langgraph-src
|
||||
RUN pip install --no-cache-dir --force-reinstall --no-deps \
|
||||
/opt/checkpoint-src /opt/langgraph-src
|
||||
|
||||
# Project graphs + registration config.
|
||||
COPY graph/ /app/graph/
|
||||
|
||||
@@ -38,16 +38,17 @@ services:
|
||||
# see ./Dockerfile. The base image bundles langgraph-api +
|
||||
# langgraph_runtime_postgres + langgraph_license + the Go core-server;
|
||||
# on top we add graph deps (deepagents), the graph files, and this
|
||||
# monorepo's local langgraph *core* (so the server runs the code under
|
||||
# test, not the released langgraph).
|
||||
# monorepo's local langgraph *core* and *checkpoint* (so the server runs
|
||||
# the code under test, not the released langgraph).
|
||||
build:
|
||||
context: .
|
||||
dockerfile: Dockerfile
|
||||
additional_contexts:
|
||||
# The monorepo's local langgraph core (libs/langgraph), installed over
|
||||
# the base image so the server runs the langgraph under test. See the
|
||||
# Dockerfile for why.
|
||||
# The monorepo's local langgraph core (libs/langgraph) and checkpoint
|
||||
# (libs/checkpoint), installed over the base image so the server runs
|
||||
# the langgraph under test. See the Dockerfile for why.
|
||||
langgraph_src: ../../langgraph
|
||||
checkpoint_src: ../../checkpoint
|
||||
image: langgraph-v3-integration-api:local
|
||||
depends_on:
|
||||
postgres:
|
||||
|
||||
Generated
+1
-1
@@ -385,7 +385,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Reference in New Issue
Block a user