mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
453da3328b |
-57
@@ -267,61 +267,6 @@ 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 = [
|
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
||||||
test_history_returns_writes_oldest_first,
|
test_history_returns_writes_oldest_first,
|
||||||
test_history_seed_is_nearest_snapshot,
|
test_history_seed_is_nearest_snapshot,
|
||||||
@@ -331,8 +276,6 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
|||||||
test_history_walk_to_root_no_seed,
|
test_history_walk_to_root_no_seed,
|
||||||
test_history_migration_plain_value_as_seed,
|
test_history_migration_plain_value_as_seed,
|
||||||
test_history_seed_ancestor_own_writes_are_replayed,
|
test_history_seed_ancestor_own_writes_are_replayed,
|
||||||
test_history_orders_parallel_writes_by_task_path,
|
|
||||||
test_history_orders_pathless_writes_first,
|
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -168,7 +168,6 @@ class _DeltaStage2Row(TypedDict, total=False):
|
|||||||
type: str | None
|
type: str | None
|
||||||
blob: bytes | None
|
blob: bytes | None
|
||||||
task_id: str | None # "w" rows only
|
task_id: str | None # "w" rows only
|
||||||
task_path: str | None # "w" rows only
|
|
||||||
idx: int | None # "w" rows only
|
idx: int | None # "w" rows only
|
||||||
version: str | None # "b" rows only
|
version: str | None # "b" rows only
|
||||||
|
|
||||||
@@ -320,7 +319,7 @@ def _build_delta_stage2_sql(
|
|||||||
branches.append(
|
branches.append(
|
||||||
"SELECT 'w'::text AS _kind, "
|
"SELECT 'w'::text AS _kind, "
|
||||||
"checkpoint_id, channel, "
|
"checkpoint_id, channel, "
|
||||||
"type, blob, task_id, task_path, idx, NULL::text AS version "
|
"type, blob, task_id, idx, NULL::text AS version "
|
||||||
"FROM checkpoint_writes "
|
"FROM checkpoint_writes "
|
||||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||||
"AND checkpoint_id = ANY(%s)"
|
"AND checkpoint_id = ANY(%s)"
|
||||||
@@ -328,8 +327,7 @@ def _build_delta_stage2_sql(
|
|||||||
for _ in channels_with_seed:
|
for _ in channels_with_seed:
|
||||||
branches.append(
|
branches.append(
|
||||||
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
|
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
|
||||||
"type, blob, NULL::text AS task_id, NULL::text AS task_path, "
|
"type, blob, NULL::text AS task_id, NULL::int AS idx, version "
|
||||||
"NULL::int AS idx, version "
|
|
||||||
"FROM checkpoint_blobs "
|
"FROM checkpoint_blobs "
|
||||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||||
"AND version = %s"
|
"AND version = %s"
|
||||||
@@ -494,11 +492,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
stored value, or when the seed blob is sentinel "empty" — in both cases
|
stored value, or when the seed blob is sentinel "empty" — in both cases
|
||||||
the consumer treats absence as "start empty".
|
the consumer treats absence as "start empty".
|
||||||
"""
|
"""
|
||||||
# writes_by_ch_by_cid[channel][cid] = list of
|
# writes_by_ch_by_cid[channel][cid] = list of (type, blob, task_id, idx)
|
||||||
# (type, blob, task_id, idx, task_path)
|
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
||||||
writes_by_ch_by_cid: dict[
|
ch: {} for ch in channels
|
||||||
str, dict[str, list[tuple[str, bytes, str, int, str]]]
|
}
|
||||||
] = {ch: {} for ch in channels}
|
|
||||||
# seed_blob_by_ver[(channel, version)] = (type, blob)
|
# seed_blob_by_ver[(channel, version)] = (type, blob)
|
||||||
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
||||||
|
|
||||||
@@ -509,14 +506,8 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
cid = cast(str, r["checkpoint_id"])
|
cid = cast(str, r["checkpoint_id"])
|
||||||
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
||||||
cast(
|
cast(
|
||||||
"tuple[str, bytes, str, int, str]",
|
"tuple[str, bytes, str, int]",
|
||||||
(
|
(r["type"], r["blob"], r["task_id"], r["idx"]),
|
||||||
r["type"],
|
|
||||||
r["blob"],
|
|
||||||
r["task_id"],
|
|
||||||
r["idx"],
|
|
||||||
r["task_path"],
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
else: # kind == "b"
|
else: # kind == "b"
|
||||||
@@ -525,10 +516,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
"tuple[str, bytes]", (r["type"], r["blob"])
|
"tuple[str, bytes]", (r["type"], r["blob"])
|
||||||
)
|
)
|
||||||
|
|
||||||
# Sort writes per (channel, cid) newest-first by (task_path, task_id, idx)
|
# Sort writes per (channel, cid) newest-first by (task_id, idx)
|
||||||
for cid_map in writes_by_ch_by_cid.values():
|
for cid_map in writes_by_ch_by_cid.values():
|
||||||
for ws in cid_map.values():
|
for ws in cid_map.values():
|
||||||
ws.sort(key=lambda w: (w[4], w[2], w[3]), reverse=True)
|
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
|
||||||
|
|
||||||
result: dict[str, DeltaChannelHistory] = {}
|
result: dict[str, DeltaChannelHistory] = {}
|
||||||
for ch in channels:
|
for ch in channels:
|
||||||
@@ -538,9 +529,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
collected: list[PendingWrite] = []
|
collected: list[PendingWrite] = []
|
||||||
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
||||||
for cid in chain_cids:
|
for cid in chain_cids:
|
||||||
for type_tag, write_blob, task_id, _idx, _path in cid_writes.get(
|
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, []):
|
||||||
cid, []
|
|
||||||
):
|
|
||||||
val = self.serde.loads_typed((type_tag, write_blob))
|
val = self.serde.loads_typed((type_tag, write_blob))
|
||||||
collected.append((task_id, ch, val))
|
collected.append((task_id, ch, val))
|
||||||
collected.reverse()
|
collected.reverse()
|
||||||
|
|||||||
@@ -81,7 +81,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
conn: sqlite3.Connection
|
conn: sqlite3.Connection
|
||||||
is_setup: bool
|
is_setup: bool
|
||||||
_has_task_path: bool = True
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -155,7 +154,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||||
checkpoint_id TEXT NOT NULL,
|
checkpoint_id TEXT NOT NULL,
|
||||||
task_id TEXT NOT NULL,
|
task_id TEXT NOT NULL,
|
||||||
task_path TEXT NOT NULL DEFAULT '',
|
|
||||||
idx INTEGER NOT NULL,
|
idx INTEGER NOT NULL,
|
||||||
channel TEXT NOT NULL,
|
channel TEXT NOT NULL,
|
||||||
type TEXT,
|
type TEXT,
|
||||||
@@ -164,19 +162,6 @@ 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
|
self.is_setup = True
|
||||||
|
|
||||||
@@ -475,9 +460,9 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
task_path: Path of the task creating the writes.
|
task_path: Path of the task creating the writes.
|
||||||
"""
|
"""
|
||||||
query = (
|
query = (
|
||||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
)
|
)
|
||||||
with self.cursor() as cur:
|
with self.cursor() as cur:
|
||||||
cur.executemany(
|
cur.executemany(
|
||||||
@@ -488,7 +473,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
str(config["configurable"]["checkpoint_ns"]),
|
str(config["configurable"]["checkpoint_ns"]),
|
||||||
str(config["configurable"]["checkpoint_id"]),
|
str(config["configurable"]["checkpoint_id"]),
|
||||||
task_id,
|
task_id,
|
||||||
task_path,
|
|
||||||
WRITES_IDX_MAP.get(channel, idx),
|
WRITES_IDX_MAP.get(channel, idx),
|
||||||
channel,
|
channel,
|
||||||
*self.serde.dumps_typed(value),
|
*self.serde.dumps_typed(value),
|
||||||
@@ -574,7 +558,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
||||||
stage2_sql = build_delta_stage2_sql(
|
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],
|
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
||||||
)
|
)
|
||||||
if stage2_sql:
|
if stage2_sql:
|
||||||
@@ -585,7 +568,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
)
|
)
|
||||||
cur.execute(stage2_sql, stage2_params)
|
cur.execute(stage2_sql, stage2_params)
|
||||||
stage2_rows = cast(
|
stage2_rows = cast(
|
||||||
"list[tuple[str, str, str, int, str, bytes, str]]", cur.fetchall()
|
"list[tuple[str, str, str, int, str, bytes]]", cur.fetchall()
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
stage2_rows = []
|
stage2_rows = []
|
||||||
|
|||||||
@@ -39,9 +39,7 @@ DELTA_STAGE1_SQL = (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_delta_stage2_sql(
|
def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
|
||||||
*, chain_lens: Sequence[int], has_task_path: bool = True
|
|
||||||
) -> str:
|
|
||||||
"""Stage-2 per-channel UNION ALL fetching writes from `writes`.
|
"""Stage-2 per-channel UNION ALL fetching writes from `writes`.
|
||||||
|
|
||||||
One branch per channel with a non-empty chain. Each branch inlines its
|
One branch per channel with a non-empty chain. Each branch inlines its
|
||||||
@@ -55,12 +53,11 @@ def build_delta_stage2_sql(
|
|||||||
of a single `channel = ANY(channels)` filter when channels have
|
of a single `channel = ANY(channels)` filter when channels have
|
||||||
different chain depths — same rationale as postgres.
|
different chain depths — same rationale as postgres.
|
||||||
"""
|
"""
|
||||||
task_path = "task_path" if has_task_path else "''"
|
|
||||||
branches: list[str] = []
|
branches: list[str] = []
|
||||||
for n in chain_lens:
|
for n in chain_lens:
|
||||||
cid_placeholders = ",".join("?" * n)
|
cid_placeholders = ",".join("?" * n)
|
||||||
branches.append(
|
branches.append(
|
||||||
f"SELECT checkpoint_id, channel, task_id, idx, type, value, {task_path} "
|
"SELECT checkpoint_id, channel, task_id, idx, type, value "
|
||||||
"FROM writes "
|
"FROM writes "
|
||||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
|
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
|
||||||
f"AND checkpoint_id IN ({cid_placeholders})"
|
f"AND checkpoint_id IN ({cid_placeholders})"
|
||||||
@@ -133,31 +130,29 @@ def build_delta_channels_writes_history(
|
|||||||
chain_by_ch: Mapping[str, list[str]],
|
chain_by_ch: Mapping[str, list[str]],
|
||||||
seed_val_by_ch: Mapping[str, Any],
|
seed_val_by_ch: Mapping[str, Any],
|
||||||
seeded: set[str],
|
seeded: set[str],
|
||||||
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes, str]],
|
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes]],
|
||||||
serde: Any,
|
serde: Any,
|
||||||
) -> dict[str, DeltaChannelHistory]:
|
) -> dict[str, DeltaChannelHistory]:
|
||||||
"""Demux stage-2 rows per channel; produce per-channel histories.
|
"""Demux stage-2 rows per channel; produce per-channel histories.
|
||||||
|
|
||||||
Stage-2 rows are
|
Stage-2 rows are `(checkpoint_id, channel, task_id, idx, type, value)`.
|
||||||
`(checkpoint_id, channel, task_id, idx, type, value, task_path)`.
|
Final write order is oldest→newest globally and `(task_id, idx)` within
|
||||||
Final write order is oldest→newest globally and
|
a checkpoint, matching the contract on `DeltaChannelHistory.writes`.
|
||||||
`(task_path, task_id, idx)` within a checkpoint, matching the contract
|
|
||||||
on `DeltaChannelHistory.writes`.
|
|
||||||
|
|
||||||
`seed` is omitted when the walk reached a true root with no snapshot
|
`seed` is omitted when the walk reached a true root with no snapshot
|
||||||
found (channel never entered `seeded`); consumers treat absence as
|
found (channel never entered `seeded`); consumers treat absence as
|
||||||
"start empty".
|
"start empty".
|
||||||
"""
|
"""
|
||||||
writes_by_ch_by_cid: dict[
|
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
||||||
str, dict[str, list[tuple[str, bytes, str, int, str]]]
|
ch: {} for ch in channels
|
||||||
] = {ch: {} for ch in channels}
|
}
|
||||||
for cid, ch, task_id, idx, type_tag, value_blob, task_path in stage2_rows:
|
for cid, ch, task_id, idx, type_tag, value_blob in stage2_rows:
|
||||||
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
||||||
(type_tag, value_blob, task_id, idx, task_path)
|
(type_tag, value_blob, task_id, idx)
|
||||||
)
|
)
|
||||||
for cid_map in writes_by_ch_by_cid.values():
|
for cid_map in writes_by_ch_by_cid.values():
|
||||||
for ws in cid_map.values():
|
for ws in cid_map.values():
|
||||||
ws.sort(key=lambda w: (w[4], w[2], w[3]))
|
ws.sort(key=lambda w: (w[2], w[3]))
|
||||||
|
|
||||||
result: dict[str, DeltaChannelHistory] = {}
|
result: dict[str, DeltaChannelHistory] = {}
|
||||||
for ch in channels:
|
for ch in channels:
|
||||||
@@ -166,7 +161,7 @@ def build_delta_channels_writes_history(
|
|||||||
collected: list[PendingWrite] = []
|
collected: list[PendingWrite] = []
|
||||||
# Chain is newest-first; iterate oldest-first for the public order.
|
# Chain is newest-first; iterate oldest-first for the public order.
|
||||||
for cid in reversed(chain_cids):
|
for cid in reversed(chain_cids):
|
||||||
for type_tag, value_blob, task_id, _idx, _path in cid_writes.get(cid, []):
|
for type_tag, value_blob, task_id, _idx in cid_writes.get(cid, []):
|
||||||
collected.append(
|
collected.append(
|
||||||
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
|
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -114,7 +114,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
lock: asyncio.Lock
|
lock: asyncio.Lock
|
||||||
is_setup: bool
|
is_setup: bool
|
||||||
_has_task_path: bool = True
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -332,7 +331,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||||
checkpoint_id TEXT NOT NULL,
|
checkpoint_id TEXT NOT NULL,
|
||||||
task_id TEXT NOT NULL,
|
task_id TEXT NOT NULL,
|
||||||
task_path TEXT NOT NULL DEFAULT '',
|
|
||||||
idx INTEGER NOT NULL,
|
idx INTEGER NOT NULL,
|
||||||
channel TEXT NOT NULL,
|
channel TEXT NOT NULL,
|
||||||
type TEXT,
|
type TEXT,
|
||||||
@@ -343,21 +341,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
):
|
):
|
||||||
await self.conn.commit()
|
await self.conn.commit()
|
||||||
|
|
||||||
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
|
|
||||||
# created before `task_path` existed and is a no-op on the rest.
|
|
||||||
try:
|
|
||||||
await self.conn.execute(
|
|
||||||
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
|
|
||||||
)
|
|
||||||
await self.conn.commit()
|
|
||||||
except aiosqlite.OperationalError as e:
|
|
||||||
# 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
|
self.is_setup = True
|
||||||
|
|
||||||
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
||||||
@@ -593,9 +576,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
task_path: Path of the task creating the writes.
|
task_path: Path of the task creating the writes.
|
||||||
"""
|
"""
|
||||||
query = (
|
query = (
|
||||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
)
|
)
|
||||||
await self.setup()
|
await self.setup()
|
||||||
async with self.lock, self.conn.cursor() as cur:
|
async with self.lock, self.conn.cursor() as cur:
|
||||||
@@ -607,7 +590,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
str(config["configurable"]["checkpoint_ns"]),
|
str(config["configurable"]["checkpoint_ns"]),
|
||||||
str(config["configurable"]["checkpoint_id"]),
|
str(config["configurable"]["checkpoint_id"]),
|
||||||
task_id,
|
task_id,
|
||||||
task_path,
|
|
||||||
WRITES_IDX_MAP.get(channel, idx),
|
WRITES_IDX_MAP.get(channel, idx),
|
||||||
channel,
|
channel,
|
||||||
*self.serde.dumps_typed(value),
|
*self.serde.dumps_typed(value),
|
||||||
@@ -689,7 +671,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
||||||
stage2_sql = build_delta_stage2_sql(
|
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],
|
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
||||||
)
|
)
|
||||||
if stage2_sql:
|
if stage2_sql:
|
||||||
@@ -700,7 +681,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
)
|
)
|
||||||
await cur.execute(stage2_sql, stage2_params)
|
await cur.execute(stage2_sql, stage2_params)
|
||||||
stage2_rows = cast(
|
stage2_rows = cast(
|
||||||
"list[tuple[str, str, str, int, str, bytes, str]]",
|
"list[tuple[str, str, str, int, str, bytes]]",
|
||||||
await cur.fetchall(),
|
await cur.fetchall(),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1,128 +0,0 @@
|
|||||||
import sqlite3
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import aiosqlite
|
|
||||||
import pytest
|
|
||||||
from langgraph.checkpoint.base import empty_checkpoint
|
|
||||||
|
|
||||||
from langgraph.checkpoint.sqlite import SqliteSaver
|
|
||||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
|
||||||
|
|
||||||
WRITES_BEFORE_TASK_PATH = """
|
|
||||||
CREATE TABLE writes (
|
|
||||||
thread_id TEXT NOT NULL,
|
|
||||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
|
||||||
checkpoint_id TEXT NOT NULL,
|
|
||||||
task_id TEXT NOT NULL,
|
|
||||||
idx INTEGER NOT NULL,
|
|
||||||
channel TEXT NOT NULL,
|
|
||||||
type TEXT,
|
|
||||||
value BLOB,
|
|
||||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
|
||||||
);
|
|
||||||
INSERT INTO writes VALUES ('t', '', 'c', 'old-task', 0, 'ch', 'null', X'');
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def legacy_db(tmp_path: Path) -> Path:
|
|
||||||
db = tmp_path / "legacy.sqlite"
|
|
||||||
with sqlite3.connect(db) as conn:
|
|
||||||
conn.executescript(WRITES_BEFORE_TASK_PATH)
|
|
||||||
return db
|
|
||||||
|
|
||||||
|
|
||||||
def test_setup_migrates_legacy_writes_table_repeatably(legacy_db: Path) -> None:
|
|
||||||
for _ in range(2):
|
|
||||||
with SqliteSaver.from_conn_string(str(legacy_db)) as saver:
|
|
||||||
saver.setup()
|
|
||||||
rows = saver.conn.execute(
|
|
||||||
"SELECT task_id, task_path FROM writes"
|
|
||||||
).fetchall()
|
|
||||||
assert rows == [("old-task", "")]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("fresh", [True, False], ids=["fresh", "legacy"])
|
|
||||||
def test_put_writes_persists_task_path(
|
|
||||||
tmp_path: Path, legacy_db: Path, fresh: bool
|
|
||||||
) -> None:
|
|
||||||
db = tmp_path / "fresh.sqlite" if fresh else legacy_db
|
|
||||||
with SqliteSaver.from_conn_string(str(db)) as saver:
|
|
||||||
config = saver.put(
|
|
||||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
|
||||||
empty_checkpoint(),
|
|
||||||
{},
|
|
||||||
{},
|
|
||||||
)
|
|
||||||
saver.put_writes(config, [("ch", "v")], "task-1", "~__pregel_pull, node")
|
|
||||||
stored = saver.conn.execute(
|
|
||||||
"SELECT task_path FROM writes WHERE task_id = 'task-1'"
|
|
||||||
).fetchall()
|
|
||||||
assert stored == [("~__pregel_pull, node",)]
|
|
||||||
|
|
||||||
|
|
||||||
async def test_async_setup_migrates_legacy_writes_table_repeatably(
|
|
||||||
legacy_db: Path,
|
|
||||||
) -> None:
|
|
||||||
for _ in range(2):
|
|
||||||
async with AsyncSqliteSaver.from_conn_string(str(legacy_db)) as saver:
|
|
||||||
await saver.setup()
|
|
||||||
config = await saver.aput(
|
|
||||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
|
||||||
empty_checkpoint(),
|
|
||||||
{},
|
|
||||||
{},
|
|
||||||
)
|
|
||||||
await saver.aput_writes(
|
|
||||||
config, [("ch", "v")], "task-1", "~__pregel_pull, node"
|
|
||||||
)
|
|
||||||
|
|
||||||
async with aiosqlite.connect(legacy_db) as conn:
|
|
||||||
async with conn.execute(
|
|
||||||
"SELECT DISTINCT task_id, task_path FROM writes ORDER BY task_id"
|
|
||||||
) as cur:
|
|
||||||
assert await cur.fetchall() == [
|
|
||||||
("old-task", ""),
|
|
||||||
("task-1", "~__pregel_pull, node"),
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
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")]}
|
|
||||||
@@ -162,14 +162,6 @@ class DeltaChannelHistory(TypedDict):
|
|||||||
Always present; possibly empty. Already filtered to one channel.
|
Always present; possibly empty. Already filtered to one channel.
|
||||||
Writes stored at the target checkpoint itself are pending for the
|
Writes stored at the target checkpoint itself are pending for the
|
||||||
next super-step and are excluded.
|
next super-step and are excluded.
|
||||||
|
|
||||||
Within a single checkpoint, writes are ordered by
|
|
||||||
`(task_path, task_id, idx)`, which is the order live execution applies
|
|
||||||
a super-step's task writes in. `task_id` is a hash of the path, so
|
|
||||||
ordering by it permutes parallel tasks writing one channel, and
|
|
||||||
reducers need not be order-invariant. Writes stored without a
|
|
||||||
`task_path` (graph input, `update_state` updates, exit-durability runs,
|
|
||||||
rows predating the column) sort first, by `task_id`.
|
|
||||||
* `seed` — the stored value at the nearest ancestor whose
|
* `seed` — the stored value at the nearest ancestor whose
|
||||||
`channel_values[ch]` is populated. Omitted if the walk reached the
|
`channel_values[ch]` is populated. Omitted if the walk reached the
|
||||||
root without finding any stored value (consumer treats absence as
|
root without finding any stored value (consumer treats absence as
|
||||||
@@ -619,11 +611,6 @@ class BaseCheckpointSaver(Generic[V]):
|
|||||||
`PostgresSaver`) override for performance; the return contract is
|
`PostgresSaver`) override for performance; the return contract is
|
||||||
fixed here.
|
fixed here.
|
||||||
|
|
||||||
`PendingWrite` carries no `task_path`, so this default replays each
|
|
||||||
checkpoint's writes in `get_tuple`'s `pending_writes` order. Savers
|
|
||||||
that do not return `pending_writes` ordered by
|
|
||||||
`(task_path, task_id, idx)` must override it.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Configuration identifying the target checkpoint.
|
config: Configuration identifying the target checkpoint.
|
||||||
channels: Channel names to walk for. Empty → empty mapping.
|
channels: Channel names to walk for. Empty → empty mapping.
|
||||||
|
|||||||
@@ -199,8 +199,8 @@ class InMemorySaver(
|
|||||||
terminated_here.add(ch)
|
terminated_here.add(ch)
|
||||||
|
|
||||||
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
|
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
|
||||||
for _, (tid, ch, serialized, _) in sorted(
|
for (_task_id, _idx), (tid, ch, serialized, _) in sorted(
|
||||||
step_writes.items(), key=lambda kv: (kv[1][3], kv[0]), reverse=True
|
step_writes.items(), reverse=True
|
||||||
):
|
):
|
||||||
if ch not in remaining:
|
if ch not in remaining:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -103,8 +103,6 @@ The CLI uses a `langgraph.json` configuration file with these key settings:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
Git dependencies should use credential-free URLs. The CLI conservatively scans direct `langgraph.json` dependencies, common Python package files, uv project and lock files, and common Node.js package and lock files for HTTP Git URLs with userinfo. This check is not exhaustive: generated Docker builds can copy other files, including nested requirement or constraint files, into image layers without scanning them. For private dependencies, provide short-lived credentials through your build environment's secret-backed Git credential helper. Do not store credentials in copied files such as `langgraph.json` or `pip_config_file`.
|
|
||||||
|
|
||||||
See the [full documentation](https://reference.langchain.com/python/langgraph-cli) for detailed configuration options.
|
See the [full documentation](https://reference.langchain.com/python/langgraph-cli) for detailed configuration options.
|
||||||
|
|
||||||
## Development
|
## Development
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import re
|
|||||||
import shlex
|
import shlex
|
||||||
import textwrap
|
import textwrap
|
||||||
from collections import Counter
|
from collections import Counter
|
||||||
from collections.abc import Iterable
|
|
||||||
from typing import Literal, NamedTuple
|
from typing import Literal, NamedTuple
|
||||||
|
|
||||||
import click
|
import click
|
||||||
@@ -37,10 +36,6 @@ DISALLOWED_BUILD_COMMAND_CHARS = [
|
|||||||
# This blocks background execution (cmd &) while allowing command
|
# This blocks background execution (cmd &) while allowing command
|
||||||
# chaining (cmd1 && cmd2) which is common in build commands.
|
# chaining (cmd1 && cmd2) which is common in build commands.
|
||||||
_SINGLE_AMPERSAND_RE = re.compile(r"(?<!&)&(?:&&)*(?!&)")
|
_SINGLE_AMPERSAND_RE = re.compile(r"(?<!&)&(?:&&)*(?!&)")
|
||||||
_GIT_HTTP_AUTHORITY_RES = (
|
|
||||||
re.compile(r"git\+https?://(?P<authority>[^/\s\"']+)", re.I),
|
|
||||||
re.compile(r"\bgit\s*=\s*[\"']https?://(?P<authority>[^/\s\"']+)", re.I),
|
|
||||||
)
|
|
||||||
_API_VERSION_PATTERN = re.compile(
|
_API_VERSION_PATTERN = re.compile(
|
||||||
r"^(?P<major>\d+)"
|
r"^(?P<major>\d+)"
|
||||||
r"(?:\.(?P<minor>\d+))?"
|
r"(?:\.(?P<minor>\d+))?"
|
||||||
@@ -83,62 +78,6 @@ def has_disallowed_build_command_content(command: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def _has_git_http_url_userinfo(dependency: str) -> bool:
|
|
||||||
"""Check whether a Git HTTP URL contains userinfo."""
|
|
||||||
return any(
|
|
||||||
"@" in match.group("authority")
|
|
||||||
for pattern in _GIT_HTTP_AUTHORITY_RES
|
|
||||||
for match in pattern.finditer(dependency)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_git_http_url_userinfo(
|
|
||||||
values: Iterable[str], *, source: pathlib.Path | None = None
|
|
||||||
) -> None:
|
|
||||||
"""Reject credential-bearing Git HTTP URLs without echoing their values."""
|
|
||||||
if not any(_has_git_http_url_userinfo(value) for value in values):
|
|
||||||
return
|
|
||||||
message = (
|
|
||||||
"Git dependency URLs must not contain credentials or other URL "
|
|
||||||
"userinfo because generated Dockerfiles and image layers can retain "
|
|
||||||
"them. Use a credential-free Git URL and provide short-lived "
|
|
||||||
"credentials through your build environment's secret-backed Git "
|
|
||||||
"credential helper."
|
|
||||||
)
|
|
||||||
if source is not None:
|
|
||||||
message += f" Found in: {source}"
|
|
||||||
raise click.UsageError(message)
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_git_http_url_userinfo_files(paths: Iterable[pathlib.Path]) -> None:
|
|
||||||
"""Reject credential-bearing Git HTTP URLs in dependency files."""
|
|
||||||
for path in paths:
|
|
||||||
path = path.resolve()
|
|
||||||
if not path.is_file():
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
contents = path.read_text(encoding="utf-8", errors="replace")
|
|
||||||
except OSError:
|
|
||||||
raise click.UsageError(
|
|
||||||
f"Could not inspect dependency file for embedded credentials: {path}"
|
|
||||||
) from None
|
|
||||||
_validate_git_http_url_userinfo([contents], source=path)
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_local_dependency_files(config_path: pathlib.Path, config: Config) -> None:
|
|
||||||
"""Validate dependency files copied into a non-uv Python image."""
|
|
||||||
paths: list[pathlib.Path] = []
|
|
||||||
for dependency in config["dependencies"]:
|
|
||||||
if not isinstance(dependency, str) or not dependency.startswith("."):
|
|
||||||
continue
|
|
||||||
root = (config_path.parent / dependency).resolve()
|
|
||||||
paths.extend(
|
|
||||||
root / name
|
|
||||||
for name in ("requirements.txt", "pyproject.toml", "setup.py", "setup.cfg")
|
|
||||||
)
|
|
||||||
_validate_git_http_url_userinfo_files(paths)
|
|
||||||
|
|
||||||
|
|
||||||
MIN_PYTHON_VERSION = "3.11"
|
MIN_PYTHON_VERSION = "3.11"
|
||||||
DEFAULT_PYTHON_VERSION = "3.11"
|
DEFAULT_PYTHON_VERSION = "3.11"
|
||||||
|
|
||||||
@@ -381,9 +320,7 @@ def _get_source_kind(config: Config) -> str | None:
|
|||||||
return kind if isinstance(kind, str) else None
|
return kind if isinstance(kind, str) else None
|
||||||
|
|
||||||
|
|
||||||
def validate_config(
|
def validate_config(config: Config) -> Config:
|
||||||
config: Config, *, source_path: pathlib.Path | None = None
|
|
||||||
) -> Config:
|
|
||||||
"""Validate a configuration dictionary."""
|
"""Validate a configuration dictionary."""
|
||||||
|
|
||||||
graphs = config.get("graphs", {})
|
graphs = config.get("graphs", {})
|
||||||
@@ -478,15 +415,6 @@ def validate_config(
|
|||||||
' "source": {"kind": "uv", "root": ".."}'
|
' "source": {"kind": "uv", "root": ".."}'
|
||||||
)
|
)
|
||||||
|
|
||||||
_validate_git_http_url_userinfo(
|
|
||||||
(
|
|
||||||
dependency
|
|
||||||
for dependency in config["dependencies"]
|
|
||||||
if isinstance(dependency, str)
|
|
||||||
),
|
|
||||||
source=source_path,
|
|
||||||
)
|
|
||||||
|
|
||||||
source = config.get("source")
|
source = config.get("source")
|
||||||
source_kind = _get_source_kind(config)
|
source_kind = _get_source_kind(config)
|
||||||
if source is not None and not isinstance(source, dict):
|
if source is not None and not isinstance(source, dict):
|
||||||
@@ -681,7 +609,7 @@ def validate_config_file(config_path: pathlib.Path) -> Config:
|
|||||||
"""Load and validate a configuration file."""
|
"""Load and validate a configuration file."""
|
||||||
with open(config_path) as f:
|
with open(config_path) as f:
|
||||||
config = json.load(f)
|
config = json.load(f)
|
||||||
validated = validate_config(config, source_path=config_path.resolve())
|
validated = validate_config(config)
|
||||||
# Enforce the package.json doesn't enforce an
|
# Enforce the package.json doesn't enforce an
|
||||||
# incompatible Node.js version
|
# incompatible Node.js version
|
||||||
if validated.get("node_version"):
|
if validated.get("node_version"):
|
||||||
@@ -1352,7 +1280,6 @@ def python_config_to_docker(
|
|||||||
api_version=api_version,
|
api_version=api_version,
|
||||||
build_tools_to_uninstall=build_tools_to_uninstall,
|
build_tools_to_uninstall=build_tools_to_uninstall,
|
||||||
)
|
)
|
||||||
_validate_local_dependency_files(config_path, config)
|
|
||||||
if pip_installer == "auto":
|
if pip_installer == "auto":
|
||||||
if _image_supports_uv(base_image):
|
if _image_supports_uv(base_image):
|
||||||
pip_installer = "uv"
|
pip_installer = "uv"
|
||||||
@@ -1563,18 +1490,7 @@ def node_config_to_docker(
|
|||||||
) -> tuple[str, dict[str, str]]:
|
) -> tuple[str, dict[str, str]]:
|
||||||
# Calculate paths for monorepo support
|
# Calculate paths for monorepo support
|
||||||
install_root = (
|
install_root = (
|
||||||
pathlib.Path(build_context).resolve()
|
pathlib.Path(build_context).resolve() if build_context else config_path.parent
|
||||||
if build_context
|
|
||||||
else config_path.parent.resolve()
|
|
||||||
)
|
|
||||||
config_root = config_path.parent.resolve()
|
|
||||||
dependency_roots = (
|
|
||||||
(install_root, config_root) if install_root != config_root else (install_root,)
|
|
||||||
)
|
|
||||||
_validate_git_http_url_userinfo_files(
|
|
||||||
root / name
|
|
||||||
for root in dependency_roots
|
|
||||||
for name in ("package.json", "package-lock.json", "yarn.lock", "pnpm-lock.yaml")
|
|
||||||
)
|
)
|
||||||
install_cmd = install_command or _get_node_pm_install_cmd(install_root)
|
install_cmd = install_command or _get_node_pm_install_cmd(install_root)
|
||||||
if build_context:
|
if build_context:
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from collections.abc import Callable, Mapping, Sequence
|
|||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from dataclasses import asdict, dataclass, field
|
from dataclasses import asdict, dataclass, field
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
from functools import partial
|
||||||
from typing import Protocol, TypeVar
|
from typing import Protocol, TypeVar
|
||||||
|
|
||||||
import click
|
import click
|
||||||
@@ -1917,7 +1918,8 @@ OPT_AGENT_ID = click.option(
|
|||||||
help="Logical agent ID (requires agent mode enabled for the tenant).",
|
help="Logical agent ID (requires agent mode enabled for the tenant).",
|
||||||
)
|
)
|
||||||
|
|
||||||
OPT_AGENT_ENVIRONMENT = click.option(
|
OPT_AGENT_ENVIRONMENT = partial(
|
||||||
|
click.option,
|
||||||
"--agent-environment",
|
"--agent-environment",
|
||||||
"environment",
|
"environment",
|
||||||
envvar="LANGSMITH_AGENT_ENVIRONMENT",
|
envvar="LANGSMITH_AGENT_ENVIRONMENT",
|
||||||
@@ -2025,7 +2027,9 @@ def _deploy_base_options(
|
|||||||
OPT_HOST_API_KEY,
|
OPT_HOST_API_KEY,
|
||||||
OPT_HOST_DEPLOYMENT_NAME,
|
OPT_HOST_DEPLOYMENT_NAME,
|
||||||
OPT_AGENT_ID,
|
OPT_AGENT_ID,
|
||||||
OPT_AGENT_ENVIRONMENT,
|
OPT_AGENT_ENVIRONMENT()
|
||||||
|
if include_docker_args
|
||||||
|
else OPT_AGENT_ENVIRONMENT(type=str),
|
||||||
click.option(
|
click.option(
|
||||||
"--deployment-id",
|
"--deployment-id",
|
||||||
help=(
|
help=(
|
||||||
@@ -2172,6 +2176,12 @@ def deploy(ctx: click.Context, **_: object):
|
|||||||
# otherwise, we return None here and click will proceed to actually run the subcommand (list or delete)
|
# otherwise, we return None here and click will proceed to actually run the subcommand (list or delete)
|
||||||
if ctx.invoked_subcommand is not None:
|
if ctx.invoked_subcommand is not None:
|
||||||
return
|
return
|
||||||
|
environment_param = next(
|
||||||
|
param for param in _deploy_cmd.params if param.name == "environment"
|
||||||
|
)
|
||||||
|
ctx.params["environment"] = environment_param.type_cast_value(
|
||||||
|
ctx, ctx.params["environment"]
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
ctx.params.get("agent_id") is not None
|
ctx.params.get("agent_id") is not None
|
||||||
or ctx.params.get("environment") is not None
|
or ctx.params.get("environment") is not None
|
||||||
@@ -2371,7 +2381,7 @@ def _deploy_cmd(
|
|||||||
@OPT_HOST_API_KEY
|
@OPT_HOST_API_KEY
|
||||||
@OPT_HOST_URL
|
@OPT_HOST_URL
|
||||||
@OPT_AGENT_ID
|
@OPT_AGENT_ID
|
||||||
@OPT_AGENT_ENVIRONMENT
|
@OPT_AGENT_ENVIRONMENT()
|
||||||
@click.option(
|
@click.option(
|
||||||
"--name-contains",
|
"--name-contains",
|
||||||
default="",
|
default="",
|
||||||
|
|||||||
@@ -650,8 +650,7 @@ class Config(TypedDict, total=False):
|
|||||||
|
|
||||||
pip_config_file: str | None
|
pip_config_file: str | None
|
||||||
"""Optional. Path to a pip config file (e.g., "/etc/pip.conf" or "pip.ini") for controlling
|
"""Optional. Path to a pip config file (e.g., "/etc/pip.conf" or "pip.ini") for controlling
|
||||||
package installation (custom indices, timeouts, etc.). The file is copied into the
|
package installation (custom indices, credentials, etc.).
|
||||||
generated image, so it must not contain credentials or other secrets.
|
|
||||||
|
|
||||||
Only relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.
|
Only relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.
|
||||||
"""
|
"""
|
||||||
@@ -690,9 +689,6 @@ class Config(TypedDict, total=False):
|
|||||||
- "." or "./src" if you have a local Python package
|
- "." or "./src" if you have a local Python package
|
||||||
- str (aka "anthropic") for a PyPI package
|
- str (aka "anthropic") for a PyPI package
|
||||||
- "git+https://github.com/org/repo.git@main" for a Git-based package
|
- "git+https://github.com/org/repo.git@main" for a Git-based package
|
||||||
Git HTTP URLs must not contain userinfo such as a username or token. For private
|
|
||||||
dependencies, provide short-lived credentials through the build environment's
|
|
||||||
secret-backed Git credential helper.
|
|
||||||
Defaults to an empty list, meaning no additional packages installed beyond your base environment.
|
Defaults to an empty list, meaning no additional packages installed beyond your base environment.
|
||||||
|
|
||||||
This field is not supported when `source.kind` is `uv`.
|
This field is not supported when `source.kind` is `uv`.
|
||||||
|
|||||||
@@ -880,7 +880,6 @@ def python_config_to_docker_uv_lock(
|
|||||||
_get_node_pm_install_cmd,
|
_get_node_pm_install_cmd,
|
||||||
_get_pip_cleanup_lines,
|
_get_pip_cleanup_lines,
|
||||||
_image_supports_uv,
|
_image_supports_uv,
|
||||||
_validate_git_http_url_userinfo_files,
|
|
||||||
docker_tag,
|
docker_tag,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -891,20 +890,11 @@ def python_config_to_docker_uv_lock(
|
|||||||
)
|
)
|
||||||
|
|
||||||
config_root = config_path.parent.resolve()
|
config_root = config_path.parent.resolve()
|
||||||
source_root = config["source"].get("root", ".")
|
|
||||||
project_root = (config_root / source_root).resolve()
|
|
||||||
_validate_git_http_url_userinfo_files(
|
|
||||||
[project_root / "pyproject.toml", project_root / "uv.lock"]
|
|
||||||
)
|
|
||||||
|
|
||||||
install_cmd = "uv pip install --system"
|
install_cmd = "uv pip install --system"
|
||||||
_, global_reqs_pip_install, pip_config_file_str = _build_python_install_commands(
|
_, global_reqs_pip_install, pip_config_file_str = _build_python_install_commands(
|
||||||
config, install_cmd
|
config, install_cmd
|
||||||
)
|
)
|
||||||
plan = _plan_uv_lock_workspace(config_path, config)
|
plan = _plan_uv_lock_workspace(config_path, config)
|
||||||
_validate_git_http_url_userinfo_files(
|
|
||||||
package.pyproject_path for package in plan.install_order
|
|
||||||
)
|
|
||||||
|
|
||||||
_update_uv_lock_graph_paths(config_path, config, plan)
|
_update_uv_lock_graph_paths(config_path, config, plan)
|
||||||
for section, key in [
|
for section, key in [
|
||||||
|
|||||||
@@ -28,7 +28,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
@@ -270,7 +270,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
|
|||||||
@@ -28,7 +28,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
@@ -270,7 +270,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
|
|||||||
@@ -255,243 +255,6 @@ def test_validate_config():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"dependency",
|
|
||||||
[
|
|
||||||
"git+https://user:secret-token@github.com/org/private.git@main",
|
|
||||||
"private-package @ git+http://token@github.com/org/private.git",
|
|
||||||
"git+HTTPS://user%40example.com:secret%2Ftoken@github.com/org/private.git",
|
|
||||||
"git+https://${GIT_TOKEN}@github.com/org/private.git",
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_validate_config_rejects_git_http_url_userinfo(dependency: str):
|
|
||||||
with pytest.raises(click.UsageError) as exc_info:
|
|
||||||
validate_config(
|
|
||||||
{
|
|
||||||
"python_version": "3.11",
|
|
||||||
"dependencies": [dependency],
|
|
||||||
"graphs": {"agent": "./agent.py:graph"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
message = str(exc_info.value)
|
|
||||||
assert "must not contain credentials or other URL userinfo" in message
|
|
||||||
assert "secret-token" not in message
|
|
||||||
assert "secret%2Ftoken" not in message
|
|
||||||
|
|
||||||
|
|
||||||
def test_validate_config_file_reports_source_for_git_http_url_userinfo(
|
|
||||||
tmp_path: pathlib.Path,
|
|
||||||
):
|
|
||||||
config_path = tmp_path / "langgraph.json"
|
|
||||||
config_path.write_text(
|
|
||||||
json.dumps(
|
|
||||||
{
|
|
||||||
"python_version": "3.11",
|
|
||||||
"dependencies": ["git+https://secret-token@github.com/org/private.git"],
|
|
||||||
"graphs": {"agent": "./agent.py:graph"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(click.UsageError) as exc_info:
|
|
||||||
validate_config_file(config_path)
|
|
||||||
|
|
||||||
message = str(exc_info.value)
|
|
||||||
assert "secret-token" not in message
|
|
||||||
assert f"Found in: {config_path.resolve()}" in message
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"manifest", ["package.json", "package-lock.json", "yarn.lock", "pnpm-lock.yaml"]
|
|
||||||
)
|
|
||||||
def test_config_to_docker_rejects_git_http_url_userinfo_in_node_files(
|
|
||||||
tmp_path: pathlib.Path, manifest: str
|
|
||||||
):
|
|
||||||
config_path = tmp_path / "langgraph.json"
|
|
||||||
config_path.write_text("{}\n")
|
|
||||||
(tmp_path / "agent.js").write_text("export const graph = {};\n")
|
|
||||||
(tmp_path / "package.json").write_text('{"name":"agent"}\n')
|
|
||||||
(tmp_path / manifest).write_text(
|
|
||||||
'"priv": "git+https://user:secret-token@github.com/org/private.git"\n'
|
|
||||||
)
|
|
||||||
config = validate_config(
|
|
||||||
{
|
|
||||||
"node_version": "20",
|
|
||||||
"graphs": {"agent": "./agent.js:graph"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(click.UsageError) as exc_info:
|
|
||||||
config_to_docker(
|
|
||||||
config_path,
|
|
||||||
config,
|
|
||||||
base_image="langchain/langgraphjs-api",
|
|
||||||
)
|
|
||||||
|
|
||||||
message = str(exc_info.value)
|
|
||||||
assert "must not contain credentials or other URL userinfo" in message
|
|
||||||
assert "secret-token" not in message
|
|
||||||
assert f"Found in: {(tmp_path / manifest).resolve()}" in message
|
|
||||||
|
|
||||||
|
|
||||||
def test_config_to_docker_allows_node_git_urls_without_http_userinfo(
|
|
||||||
tmp_path: pathlib.Path,
|
|
||||||
):
|
|
||||||
config_path = tmp_path / "langgraph.json"
|
|
||||||
config_path.write_text("{}\n")
|
|
||||||
(tmp_path / "agent.js").write_text("export const graph = {};\n")
|
|
||||||
(tmp_path / "package.json").write_text(
|
|
||||||
'{"dependencies":{"public":"git+https://github.com/org/public.git"}}\n'
|
|
||||||
)
|
|
||||||
config = validate_config(
|
|
||||||
{
|
|
||||||
"node_version": "20",
|
|
||||||
"graphs": {"agent": "./agent.js:graph"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
docker, _ = config_to_docker(
|
|
||||||
config_path,
|
|
||||||
config,
|
|
||||||
base_image="langchain/langgraphjs-api",
|
|
||||||
)
|
|
||||||
|
|
||||||
assert f"ADD . /deps/{tmp_path.name}" in docker
|
|
||||||
|
|
||||||
|
|
||||||
def test_config_to_docker_rejects_git_http_url_userinfo_in_node_workspace(
|
|
||||||
tmp_path: pathlib.Path,
|
|
||||||
):
|
|
||||||
config_root = tmp_path / "apps" / "agent"
|
|
||||||
config_root.mkdir(parents=True)
|
|
||||||
config_path = config_root / "langgraph.json"
|
|
||||||
config_path.write_text("{}\n")
|
|
||||||
(config_root / "agent.js").write_text("export const graph = {};\n")
|
|
||||||
(config_root / "package.json").write_text(
|
|
||||||
'{"dependencies":{"priv":"git+https://secret-token@github.com/org/private.git"}}\n'
|
|
||||||
)
|
|
||||||
(tmp_path / "package.json").write_text('{"name":"workspace"}\n')
|
|
||||||
config = validate_config(
|
|
||||||
{
|
|
||||||
"node_version": "20",
|
|
||||||
"graphs": {"agent": "./agent.js:graph"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(click.UsageError) as exc_info:
|
|
||||||
config_to_docker(
|
|
||||||
config_path,
|
|
||||||
config,
|
|
||||||
base_image="langchain/langgraphjs-api",
|
|
||||||
build_context=str(tmp_path),
|
|
||||||
)
|
|
||||||
|
|
||||||
message = str(exc_info.value)
|
|
||||||
assert "secret-token" not in message
|
|
||||||
assert f"Found in: {(config_root / 'package.json').resolve()}" in message
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"dependency",
|
|
||||||
[
|
|
||||||
"git+https://github.com/org/public.git@main",
|
|
||||||
"private-package @ git+https://github.com/org/private.git@main",
|
|
||||||
"git+ssh://git@github.com/org/private.git@main",
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_validate_config_allows_git_urls_without_http_userinfo(dependency: str):
|
|
||||||
config = validate_config(
|
|
||||||
{
|
|
||||||
"python_version": "3.11",
|
|
||||||
"dependencies": [dependency],
|
|
||||||
"graphs": {"agent": "./agent.py:graph"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
assert config["dependencies"] == [dependency]
|
|
||||||
|
|
||||||
|
|
||||||
def test_config_to_docker_rejects_git_http_url_userinfo_in_requirements(
|
|
||||||
tmp_path: pathlib.Path,
|
|
||||||
):
|
|
||||||
config_path = tmp_path / "langgraph.json"
|
|
||||||
config_path.write_text("{}\n")
|
|
||||||
(tmp_path / "agent.py").write_text("graph = object()\n")
|
|
||||||
(tmp_path / "requirements.txt").write_text(
|
|
||||||
"private @ git+https://secret-token@github.com/org/private.git\n"
|
|
||||||
)
|
|
||||||
config = validate_config(
|
|
||||||
{
|
|
||||||
"python_version": "3.11",
|
|
||||||
"dependencies": ["."],
|
|
||||||
"graphs": {"agent": "./agent.py:graph"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(click.UsageError) as exc_info:
|
|
||||||
config_to_docker(
|
|
||||||
config_path,
|
|
||||||
config,
|
|
||||||
base_image="langchain/langgraph-api:0.2.47",
|
|
||||||
)
|
|
||||||
|
|
||||||
message = str(exc_info.value)
|
|
||||||
assert "must not contain credentials or other URL userinfo" in message
|
|
||||||
assert "secret-token" not in message
|
|
||||||
assert f"Found in: {(tmp_path / 'requirements.txt').resolve()}" in message
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("manifest", ["pyproject.toml", "uv.lock"])
|
|
||||||
def test_config_to_docker_rejects_git_http_url_userinfo_in_uv_files(
|
|
||||||
tmp_path: pathlib.Path, manifest: str
|
|
||||||
):
|
|
||||||
config_path = tmp_path / "langgraph.json"
|
|
||||||
config_path.write_text("{}\n")
|
|
||||||
(tmp_path / "src").mkdir()
|
|
||||||
(tmp_path / "src" / "agent.py").write_text("graph = object()\n")
|
|
||||||
pyproject = textwrap.dedent(
|
|
||||||
"""
|
|
||||||
[project]
|
|
||||||
name = "agent"
|
|
||||||
version = "0.1.0"
|
|
||||||
dependencies = ["private"]
|
|
||||||
|
|
||||||
[tool.uv.sources]
|
|
||||||
private = { git = "https://github.com/org/private.git" }
|
|
||||||
"""
|
|
||||||
).strip()
|
|
||||||
uv_lock = "# uv lock file\n"
|
|
||||||
if manifest == "pyproject.toml":
|
|
||||||
pyproject = pyproject.replace(
|
|
||||||
"https://github.com", "https://secret-token@github.com"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
uv_lock += (
|
|
||||||
'source = { git = "https://secret-token@github.com/org/private.git" }\n'
|
|
||||||
)
|
|
||||||
(tmp_path / "pyproject.toml").write_text(pyproject + "\n")
|
|
||||||
(tmp_path / "uv.lock").write_text(uv_lock)
|
|
||||||
config = validate_config(
|
|
||||||
{
|
|
||||||
"python_version": "3.11",
|
|
||||||
"graphs": {"agent": "./src/agent.py:graph"},
|
|
||||||
"source": {"kind": "uv"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(click.UsageError) as exc_info:
|
|
||||||
config_to_docker(
|
|
||||||
config_path,
|
|
||||||
config,
|
|
||||||
base_image="langchain/langgraph-api:0.2.47",
|
|
||||||
)
|
|
||||||
|
|
||||||
message = str(exc_info.value)
|
|
||||||
assert "must not contain credentials or other URL userinfo" in message
|
|
||||||
assert "secret-token" not in message
|
|
||||||
|
|
||||||
|
|
||||||
def test_validate_config_image_distro():
|
def test_validate_config_image_distro():
|
||||||
"""Test validation of image_distro field."""
|
"""Test validation of image_distro field."""
|
||||||
# Valid image_distro values should work
|
# Valid image_distro values should work
|
||||||
|
|||||||
@@ -1,113 +0,0 @@
|
|||||||
"""`DeltaChannel` replay must apply parallel writes in the order `invoke` did."""
|
|
||||||
|
|
||||||
from typing import Annotated, Any
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
|
||||||
from typing_extensions import TypedDict
|
|
||||||
|
|
||||||
from langgraph.channels.delta import DeltaChannel
|
|
||||||
from langgraph.graph import END, START, StateGraph
|
|
||||||
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_continuing_thread_preserves_committed_prefix(
|
|
||||||
async_checkpointer: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
graph = _build_fan_out_graph(async_checkpointer)
|
|
||||||
config = {"configurable": {"thread_id": "1"}}
|
|
||||||
|
|
||||||
first = (await graph.ainvoke({"items": []}, config))["items"]
|
|
||||||
second = (await graph.ainvoke({"items": []}, config))["items"]
|
|
||||||
|
|
||||||
assert second == first + first
|
|
||||||
assert (await graph.aget_state(config)).values["items"] == second
|
|
||||||
|
|
||||||
|
|
||||||
async def test_state_history_reports_live_order_at_every_step(
|
|
||||||
async_checkpointer: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
runs = 3
|
|
||||||
graph = _build_fan_out_graph(async_checkpointer)
|
|
||||||
config = {"configurable": {"thread_id": "1"}}
|
|
||||||
for _ in range(runs):
|
|
||||||
await graph.ainvoke({"items": []}, config)
|
|
||||||
live = FAN_OUT_NAMES * runs
|
|
||||||
|
|
||||||
seen = [
|
|
||||||
s.values["items"]
|
|
||||||
async for s in graph.aget_state_history(config)
|
|
||||||
if "items" in s.values
|
|
||||||
]
|
|
||||||
|
|
||||||
assert max(map(len, seen)) == len(live)
|
|
||||||
for values in seen:
|
|
||||||
assert values == live[: len(values)], f"{values} is not a prefix of {live}"
|
|
||||||
Reference in New Issue
Block a user