feat(middleware): upgrade deepagents 0.5.7 → 0.6.2 (#231)
* feat(middleware): add CodeInterpreterMiddleware with project-specific configuration chore(config): increase checkpoint retention limit for runaway conversations fix(tests): update database schema references from 'blob' to 'value' chore(deps): update deepagents dependency to include quickjs support * feat(deepagents): update to version 0.6.1 and add optional dependencies for quickjs * feat(sessions): improve error handling for message deltas and update Overwrite type check * Enhance PruningCheckpointer with DeltaChannel Awareness - Introduced a new pruning strategy in `_prune_after_put` to preserve the `_DeltaSnapshot` chain during checkpoint pruning. - Implemented methods to fetch recent checkpoint IDs and walk to snapshot ancestors, ensuring that necessary checkpoints are retained. - Updated SQL queries to handle checkpoint and write deletions more efficiently. - Added comprehensive tests for DeltaChannel-aware pruning, ensuring that the pruning logic correctly handles various checkpoint scenarios, including those with and without snapshot seeds. - Refactored `_load_checkpoint_messages` to utilize the new saver interface, improving message reconstruction from checkpoints. * feat(tests): add migration sweep test to preserve snapshot ancestor * feat(sessions): enhance checkpoint retrieval to prevent transcript leakage in multi-agent scenarios * feat(middleware): enhance CodeInterpreterMiddleware with configurable timeout and result character limit feat(config): add CodeInterpreterMiddleware tuning parameters to EvoScientistConfig feat(sessions): implement inline message delta reducer for improved message handling * feat(dependencies): update deepagents version to 0.6.2 in pyproject.toml and uv.lock
This commit is contained in:
+625
-12
@@ -145,12 +145,18 @@ class TestThreadFunctions(unittest.TestCase):
|
||||
idx INTEGER NOT NULL,
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT,
|
||||
blob BLOB,
|
||||
value BLOB,
|
||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
||||
)
|
||||
""")
|
||||
|
||||
# Insert test checkpoints
|
||||
# Insert test checkpoints. ``type`` + ``checkpoint`` are
|
||||
# populated with a serialized empty-state blob so upstream
|
||||
# ``aget_tuple`` (used by message reconstruction) can
|
||||
# deserialize them — production checkpoints always have
|
||||
# these set; bare-metadata rows are a test fiction.
|
||||
serde = JsonPlusSerializer()
|
||||
empty_ck_type, empty_ck_blob = serde.dumps_typed({"channel_values": {}})
|
||||
for i, tid in enumerate(["abc12345", "abc12399", "def00001"]):
|
||||
meta = json.dumps(
|
||||
{
|
||||
@@ -161,8 +167,8 @@ class TestThreadFunctions(unittest.TestCase):
|
||||
}
|
||||
)
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) VALUES (?, '', ?, ?)",
|
||||
(tid, f"cp_{i}", meta),
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, type, checkpoint, metadata) VALUES (?, '', ?, ?, ?, ?)",
|
||||
(tid, f"cp_{i}", empty_ck_type, empty_ck_blob, meta),
|
||||
)
|
||||
|
||||
# Insert a non-EvoScientist checkpoint (should be filtered)
|
||||
@@ -173,8 +179,8 @@ class TestThreadFunctions(unittest.TestCase):
|
||||
}
|
||||
)
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) VALUES (?, '', ?, ?)",
|
||||
("zzz99999", "cp_other", other_meta),
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, type, checkpoint, metadata) VALUES (?, '', ?, ?, ?, ?)",
|
||||
("zzz99999", "cp_other", empty_ck_type, empty_ck_blob, other_meta),
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
@@ -364,6 +370,254 @@ class TestThreadFunctions(unittest.TestCase):
|
||||
finally:
|
||||
_run(_cleanup())
|
||||
|
||||
def test_get_thread_messages_reconstructs_multi_delta_chain(self):
|
||||
"""3-checkpoint chain with ``_DeltaSnapshot`` seed + pending writes.
|
||||
|
||||
Exercises the upstream ``aget_delta_channel_history`` walk: the
|
||||
latest checkpoint has no materialized seed, so the walk must climb
|
||||
back through an intermediate delta-only ancestor to a snapshot
|
||||
further back, then accumulate writes oldest→newest on top.
|
||||
"""
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
async def _insert():
|
||||
import aiosqlite
|
||||
|
||||
serde = JsonPlusSerializer()
|
||||
|
||||
seed_messages = [
|
||||
HumanMessage(content="m1", id="m1"),
|
||||
AIMessage(content="m2", id="m2"),
|
||||
]
|
||||
cp1_type, cp1_blob = serde.dumps_typed(
|
||||
{"channel_values": {"messages": _DeltaSnapshot(value=seed_messages)}}
|
||||
)
|
||||
cp_empty_type, cp_empty_blob = serde.dumps_typed({"channel_values": {}})
|
||||
|
||||
w2_type, w2_blob = serde.dumps_typed([HumanMessage(content="m3", id="m3")])
|
||||
w3_type, w3_blob = serde.dumps_typed([AIMessage(content="m4", id="m4")])
|
||||
|
||||
meta = json.dumps(
|
||||
{
|
||||
"agent_name": AGENT_NAME,
|
||||
"updated_at": "2025-01-26T10:00:00+00:00",
|
||||
}
|
||||
)
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
for cid, parent, ck_type, ck_blob in [
|
||||
("cp_chain_1", None, cp1_type, cp1_blob),
|
||||
("cp_chain_2", "cp_chain_1", cp_empty_type, cp_empty_blob),
|
||||
("cp_chain_3", "cp_chain_2", cp_empty_type, cp_empty_blob),
|
||||
]:
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, parent_checkpoint_id, type, checkpoint, metadata) "
|
||||
"VALUES (?, '', ?, ?, ?, ?, ?)",
|
||||
("chain12345", cid, parent, ck_type, ck_blob, meta),
|
||||
)
|
||||
for cid, wtype, wblob in [
|
||||
("cp_chain_2", w2_type, w2_blob),
|
||||
("cp_chain_3", w3_type, w3_blob),
|
||||
]:
|
||||
await conn.execute(
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) "
|
||||
"VALUES (?, '', ?, ?, ?, ?, ?, ?)",
|
||||
("chain12345", cid, "task0", 0, "messages", wtype, wblob),
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
async def _cleanup():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
await conn.execute(
|
||||
"DELETE FROM checkpoints WHERE thread_id = ?",
|
||||
("chain12345",),
|
||||
)
|
||||
await conn.execute(
|
||||
"DELETE FROM writes WHERE thread_id = ?",
|
||||
("chain12345",),
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
async def _assert_walk_branch_active():
|
||||
"""Confirm cp_chain_3 has no materialized ``messages`` seed.
|
||||
|
||||
Without this guard, a future refactor that ends up writing a
|
||||
``_DeltaSnapshot`` at every checkpoint would silently move
|
||||
this test onto the hybrid's "target seed" branch — the
|
||||
reconstruction would still match, but the ancestor walk
|
||||
being tested here would never run. The assertion locks in
|
||||
which branch the test exercises.
|
||||
"""
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
async with conn.execute(
|
||||
"SELECT type, checkpoint FROM checkpoints "
|
||||
"WHERE thread_id = ? AND checkpoint_id = ?",
|
||||
("chain12345", "cp_chain_3"),
|
||||
) as cur:
|
||||
row = await cur.fetchone()
|
||||
assert row is not None
|
||||
ck = JsonPlusSerializer().loads_typed((row[0], row[1]))
|
||||
assert "messages" not in (ck.get("channel_values") or {}), (
|
||||
"cp_chain_3 must NOT carry a messages seed — this test "
|
||||
"exercises the ancestor-walk branch, not the target-seed "
|
||||
"shortcut."
|
||||
)
|
||||
|
||||
_run(_insert())
|
||||
try:
|
||||
_run(_assert_walk_branch_active())
|
||||
messages = _run(get_thread_messages("chain12345"))
|
||||
assert [m.content for m in messages] == ["m1", "m2", "m3", "m4"]
|
||||
assert isinstance(messages[0], HumanMessage)
|
||||
assert isinstance(messages[1], AIMessage)
|
||||
assert isinstance(messages[2], HumanMessage)
|
||||
assert isinstance(messages[3], AIMessage)
|
||||
finally:
|
||||
_run(_cleanup())
|
||||
|
||||
def test_get_thread_messages_handles_overwrite_bare_message(self):
|
||||
"""``Overwrite(value=<bare BaseMessage>)`` wraps to a single-element list.
|
||||
|
||||
The ``Overwrite`` reset branch in ``_load_checkpoint_messages``
|
||||
has three sub-cases: list value (most common), ``None`` (clears
|
||||
state), and a bare ``BaseMessage`` (rare but valid). The
|
||||
last case has no other test coverage — this guards against a
|
||||
refactor that silently drops the ``[inner]`` wrapping fallback.
|
||||
"""
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from langgraph.types import Overwrite
|
||||
|
||||
async def _insert():
|
||||
import aiosqlite
|
||||
|
||||
serde = JsonPlusSerializer()
|
||||
seed_messages = [
|
||||
HumanMessage(content="m1", id="m1"),
|
||||
AIMessage(content="m2", id="m2"),
|
||||
]
|
||||
ck_type, ck_blob = serde.dumps_typed(
|
||||
{"channel_values": {"messages": _DeltaSnapshot(value=seed_messages)}}
|
||||
)
|
||||
# Bare message (NOT wrapped in a list) — the rare third case
|
||||
# the Overwrite branch handles.
|
||||
ow = Overwrite(value=HumanMessage(content="replaced", id="repl"))
|
||||
w_type, w_blob = serde.dumps_typed(ow)
|
||||
meta = json.dumps({"agent_name": AGENT_NAME})
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, "
|
||||
"parent_checkpoint_id, type, checkpoint, metadata) "
|
||||
"VALUES (?, '', ?, NULL, ?, ?, ?)",
|
||||
("ow_bare01", "cp_bare", ck_type, ck_blob, meta),
|
||||
)
|
||||
await conn.execute(
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) "
|
||||
"VALUES (?, '', ?, 'task0', 0, 'messages', ?, ?)",
|
||||
("ow_bare01", "cp_bare", w_type, w_blob),
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
async def _cleanup():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
await conn.execute(
|
||||
"DELETE FROM checkpoints WHERE thread_id = ?", ("ow_bare01",)
|
||||
)
|
||||
await conn.execute(
|
||||
"DELETE FROM writes WHERE thread_id = ?", ("ow_bare01",)
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
_run(_insert())
|
||||
try:
|
||||
messages = _run(get_thread_messages("ow_bare01"))
|
||||
# Overwrite replaced the seed completely; bare message wrapped
|
||||
# in a 1-element list.
|
||||
assert len(messages) == 1
|
||||
assert isinstance(messages[0], HumanMessage)
|
||||
assert messages[0].content == "replaced"
|
||||
assert messages[0].id == "repl"
|
||||
finally:
|
||||
_run(_cleanup())
|
||||
|
||||
def test_get_thread_messages_ignores_colliding_other_agent(self):
|
||||
"""Multi-agent DB with thread_id collision: must surface only ours.
|
||||
|
||||
Without the agent_name filter on the head-checkpoint lookup,
|
||||
``saver.aget_tuple()`` returns the latest by ``checkpoint_id``
|
||||
alone — so if a third-party agent's checkpoint for the same
|
||||
``thread_id`` happens to have a higher id (lexicographically),
|
||||
we'd leak its transcript into ``/resume``. Pinning the head to
|
||||
the latest EvoScientist-agent row prevents that.
|
||||
"""
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
async def _insert():
|
||||
import aiosqlite
|
||||
|
||||
serde = JsonPlusSerializer()
|
||||
evo_messages = [
|
||||
HumanMessage(content="ours_1", id="o1"),
|
||||
AIMessage(content="ours_2", id="o2"),
|
||||
]
|
||||
other_messages = [
|
||||
HumanMessage(content="theirs_1", id="t1"),
|
||||
AIMessage(content="theirs_2", id="t2"),
|
||||
HumanMessage(content="theirs_3", id="t3"),
|
||||
]
|
||||
evo_type, evo_blob = serde.dumps_typed(
|
||||
{"channel_values": {"messages": _DeltaSnapshot(value=evo_messages)}}
|
||||
)
|
||||
other_type, other_blob = serde.dumps_typed(
|
||||
{"channel_values": {"messages": _DeltaSnapshot(value=other_messages)}}
|
||||
)
|
||||
evo_meta = json.dumps({"agent_name": AGENT_NAME})
|
||||
other_meta = json.dumps({"agent_name": "ThirdPartyAgent"})
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
# EvoScientist's checkpoint id is LEXICOGRAPHICALLY
|
||||
# SMALLER than the third-party agent's, so a naive
|
||||
# "latest by checkpoint_id" lookup would pick the wrong
|
||||
# one.
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, "
|
||||
"parent_checkpoint_id, type, checkpoint, metadata) "
|
||||
"VALUES (?, '', 'aaa_evo', NULL, ?, ?, ?)",
|
||||
("collide01", evo_type, evo_blob, evo_meta),
|
||||
)
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, "
|
||||
"parent_checkpoint_id, type, checkpoint, metadata) "
|
||||
"VALUES (?, '', 'zzz_other', NULL, ?, ?, ?)",
|
||||
("collide01", other_type, other_blob, other_meta),
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
async def _cleanup():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
await conn.execute(
|
||||
"DELETE FROM checkpoints WHERE thread_id = ?", ("collide01",)
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
_run(_insert())
|
||||
try:
|
||||
messages = _run(get_thread_messages("collide01"))
|
||||
assert [m.content for m in messages] == ["ours_1", "ours_2"]
|
||||
# Defense-in-depth: explicitly forbid leakage of the other
|
||||
# agent's content.
|
||||
for msg in messages:
|
||||
assert not msg.content.startswith("theirs_")
|
||||
finally:
|
||||
_run(_cleanup())
|
||||
|
||||
# -- Agent isolation: OtherAgent data should never be visible --
|
||||
|
||||
def test_thread_exists_ignores_other_agent(self):
|
||||
@@ -403,7 +657,7 @@ class TestThreadFunctions(unittest.TestCase):
|
||||
(shared_tid, "cp_evo_shared", evo_meta),
|
||||
)
|
||||
await conn.execute(
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob) "
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) "
|
||||
"VALUES (?, '', ?, 't1', 0, 'ch', 'str', X'AA')",
|
||||
(shared_tid, "cp_evo_shared"),
|
||||
)
|
||||
@@ -420,7 +674,7 @@ class TestThreadFunctions(unittest.TestCase):
|
||||
(shared_tid, "cp_other_shared", other_meta),
|
||||
)
|
||||
await conn.execute(
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob) "
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) "
|
||||
"VALUES (?, '', ?, 't2', 0, 'ch', 'str', X'BB')",
|
||||
(shared_tid, "cp_other_shared"),
|
||||
)
|
||||
@@ -869,6 +1123,264 @@ class TestPruningCheckpointer(unittest.TestCase):
|
||||
assert survivors == [ids[4], ids[3]]
|
||||
|
||||
|
||||
class TestPruningCheckpointerDeltaChannel(unittest.TestCase):
|
||||
"""Tests for DeltaChannel-aware pruning.
|
||||
|
||||
The naive ``keep_latest`` pruner can sever the ``_DeltaSnapshot``
|
||||
chain that ``messages`` reconstruction depends on. These tests
|
||||
exercise the walk-to-snapshot-ancestor extension that preserves the
|
||||
chain head between each kept anchor and the nearest snapshot.
|
||||
|
||||
Checkpoints are inserted directly via SQL with explicit
|
||||
``parent_checkpoint_id`` to give precise control over the chain
|
||||
structure (the chained-``aput`` pattern in
|
||||
``TestPruningCheckpointer`` doesn't let us pick which checkpoint
|
||||
materializes a snapshot).
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self._tmpdir = tempfile.mkdtemp()
|
||||
self._db_path = os.path.join(self._tmpdir, "delta_prune.db")
|
||||
|
||||
def tearDown(self):
|
||||
try:
|
||||
os.unlink(self._db_path)
|
||||
os.rmdir(self._tmpdir)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
async def _insert_chain(conn, thread_id, specs):
|
||||
"""Insert a chain of checkpoints in oldest→newest order.
|
||||
|
||||
``specs`` is a list of ``(cid, parent_cid, msgs_seed)`` tuples:
|
||||
- ``cid``: checkpoint_id
|
||||
- ``parent_cid``: parent_checkpoint_id (``None`` for chain root)
|
||||
- ``msgs_seed``: value stored at ``channel_values["messages"]``
|
||||
— pass ``None`` for delta-only (no seed), a ``list`` for
|
||||
plain seed, or a ``_DeltaSnapshot`` for wrapped seed.
|
||||
"""
|
||||
serde = JsonPlusSerializer()
|
||||
meta = json.dumps({"agent_name": AGENT_NAME})
|
||||
for cid, parent_cid, msgs in specs:
|
||||
cv = {"messages": msgs} if msgs is not None else {}
|
||||
ck_type, ck_blob = serde.dumps_typed({"channel_values": cv})
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, "
|
||||
"parent_checkpoint_id, type, checkpoint, metadata) "
|
||||
"VALUES (?, '', ?, ?, ?, ?, ?)",
|
||||
(thread_id, cid, parent_cid, ck_type, ck_blob, meta),
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
@staticmethod
|
||||
async def _surviving_ids(conn, thread_id):
|
||||
async with conn.execute(
|
||||
"SELECT checkpoint_id FROM checkpoints WHERE thread_id = ? "
|
||||
"ORDER BY checkpoint_id ASC",
|
||||
(thread_id,),
|
||||
) as cur:
|
||||
return [r[0] for r in await cur.fetchall()]
|
||||
|
||||
def test_preserves_snapshot_ancestor(self):
|
||||
"""Snapshot lives outside the anchor window → walk reaches and stops."""
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
from EvoScientist.sessions import PruningCheckpointer
|
||||
|
||||
tid = "tdsa_001"
|
||||
|
||||
async def _go():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
saver = PruningCheckpointer(conn, keep_per_ns=5)
|
||||
await saver.setup()
|
||||
# 10 checkpoints; snapshot at cp_003 (outside the anchor
|
||||
# window of cp_006..cp_010). Walk from cp_005 → cp_004 →
|
||||
# cp_003 (seed found, stop). Survivors: cp_003..cp_010
|
||||
# (8). Pruned: cp_001, cp_002.
|
||||
specs = [
|
||||
("cp_001", None, None),
|
||||
("cp_002", "cp_001", None),
|
||||
("cp_003", "cp_002", _DeltaSnapshot(value=[])),
|
||||
*[(f"cp_{i:03d}", f"cp_{i - 1:03d}", None) for i in range(4, 11)],
|
||||
]
|
||||
await self._insert_chain(conn, tid, specs)
|
||||
await saver._prune_after_put(tid, "")
|
||||
await conn.commit()
|
||||
return await self._surviving_ids(conn, tid)
|
||||
|
||||
survivors = _run(_go())
|
||||
assert survivors == [f"cp_{i:03d}" for i in range(3, 11)]
|
||||
|
||||
def test_preserves_full_chain_when_no_snapshot(self):
|
||||
"""No snapshot anywhere → walk reaches root, preserves everything."""
|
||||
from EvoScientist.sessions import PruningCheckpointer
|
||||
|
||||
tid = "tdsa_002"
|
||||
|
||||
async def _go():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
saver = PruningCheckpointer(conn, keep_per_ns=5)
|
||||
await saver.setup()
|
||||
# 10 delta-only checkpoints, no seed anywhere. Walk
|
||||
# exhausts to root (cp_001's parent is None → break).
|
||||
# All 10 must survive — the alternative is silent
|
||||
# truncation, which is the bug Fix #B prevents.
|
||||
specs = [("cp_001", None, None)] + [
|
||||
(f"cp_{i:03d}", f"cp_{i - 1:03d}", None) for i in range(2, 11)
|
||||
]
|
||||
await self._insert_chain(conn, tid, specs)
|
||||
await saver._prune_after_put(tid, "")
|
||||
await conn.commit()
|
||||
return await self._surviving_ids(conn, tid)
|
||||
|
||||
survivors = _run(_go())
|
||||
assert survivors == [f"cp_{i:03d}" for i in range(1, 11)]
|
||||
|
||||
def test_plain_list_seed_also_terminates_walk(self):
|
||||
"""Pre-DeltaChannel format (plain list in channel_values) also counts as seed."""
|
||||
from EvoScientist.sessions import PruningCheckpointer
|
||||
|
||||
tid = "tdsa_003"
|
||||
|
||||
async def _go():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
saver = PruningCheckpointer(conn, keep_per_ns=3)
|
||||
await saver.setup()
|
||||
# 6 checkpoints; cp_002 has plain-list seed (legacy
|
||||
# format). Walk from cp_003 → cp_002 (seed) → stop.
|
||||
# Survivors: cp_002..cp_006 (5). Pruned: cp_001.
|
||||
specs = [
|
||||
("cp_001", None, None),
|
||||
("cp_002", "cp_001", []), # plain list seed
|
||||
("cp_003", "cp_002", None),
|
||||
("cp_004", "cp_003", None),
|
||||
("cp_005", "cp_004", None),
|
||||
("cp_006", "cp_005", None),
|
||||
]
|
||||
await self._insert_chain(conn, tid, specs)
|
||||
await saver._prune_after_put(tid, "")
|
||||
await conn.commit()
|
||||
return await self._surviving_ids(conn, tid)
|
||||
|
||||
survivors = _run(_go())
|
||||
assert survivors == ["cp_002", "cp_003", "cp_004", "cp_005", "cp_006"]
|
||||
|
||||
def test_chain_break_stops_walk_cleanly(self):
|
||||
"""Missing ancestor row breaks the chain; walk stops without raising."""
|
||||
from EvoScientist.sessions import PruningCheckpointer
|
||||
|
||||
tid = "tdsa_004"
|
||||
|
||||
async def _go():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
saver = PruningCheckpointer(conn, keep_per_ns=2)
|
||||
await saver.setup()
|
||||
# Insert cp_001..cp_005, then DELETE cp_002 to break
|
||||
# the chain. Walk from cp_003 → tries cp_002 →
|
||||
# _fetch_checkpoint_blob returns None → break with
|
||||
# nothing added (because we add the cursor's id only
|
||||
# AFTER fetching its blob succeeds).
|
||||
specs = [
|
||||
("cp_001", None, None),
|
||||
("cp_002", "cp_001", None),
|
||||
("cp_003", "cp_002", None),
|
||||
("cp_004", "cp_003", None),
|
||||
("cp_005", "cp_004", None),
|
||||
]
|
||||
await self._insert_chain(conn, tid, specs)
|
||||
await conn.execute(
|
||||
"DELETE FROM checkpoints WHERE thread_id = ? AND checkpoint_id = ?",
|
||||
(tid, "cp_002"),
|
||||
)
|
||||
await conn.commit()
|
||||
await saver._prune_after_put(tid, "")
|
||||
await conn.commit()
|
||||
return await self._surviving_ids(conn, tid)
|
||||
|
||||
survivors = _run(_go())
|
||||
# anchors = cp_004, cp_005. Walk visits cp_003 (preserved),
|
||||
# then cp_002 → None → break. cp_001 pruned. cp_002 already
|
||||
# absent. Survivors: cp_003, cp_004, cp_005.
|
||||
assert survivors == ["cp_003", "cp_004", "cp_005"]
|
||||
|
||||
def test_deserialization_failure_safe_side_over_preserves(self):
|
||||
"""Corrupt blob mid-walk: pruner preserves what it visited so far."""
|
||||
from EvoScientist.sessions import PruningCheckpointer
|
||||
|
||||
tid = "tdsa_005"
|
||||
|
||||
async def _go():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
saver = PruningCheckpointer(conn, keep_per_ns=2)
|
||||
await saver.setup()
|
||||
# cp_001..cp_005, all delta-only. Then overwrite cp_003
|
||||
# with a corrupt blob. Walk from cp_003: fetch blob
|
||||
# succeeds (returns garbage bytes), add cp_003 to
|
||||
# extra_preserve, deserialize FAILS → break with cp_003
|
||||
# already preserved.
|
||||
specs = [
|
||||
("cp_001", None, None),
|
||||
("cp_002", "cp_001", None),
|
||||
("cp_003", "cp_002", None),
|
||||
("cp_004", "cp_003", None),
|
||||
("cp_005", "cp_004", None),
|
||||
]
|
||||
await self._insert_chain(conn, tid, specs)
|
||||
await conn.execute(
|
||||
"UPDATE checkpoints SET type = ?, checkpoint = ? "
|
||||
"WHERE thread_id = ? AND checkpoint_id = ?",
|
||||
("garbage_type", b"not a real blob", tid, "cp_003"),
|
||||
)
|
||||
await conn.commit()
|
||||
await saver._prune_after_put(tid, "")
|
||||
await conn.commit()
|
||||
return await self._surviving_ids(conn, tid)
|
||||
|
||||
survivors = _run(_go())
|
||||
# anchors = cp_004, cp_005. Walk visits cp_003 (added to
|
||||
# extra_preserve before deserialize fails). cp_001, cp_002
|
||||
# pruned. Survivors: cp_003, cp_004, cp_005.
|
||||
assert survivors == ["cp_003", "cp_004", "cp_005"]
|
||||
|
||||
def test_anchor_count_below_keep_is_noop(self):
|
||||
"""When checkpoint count < keep_per_ns, prune returns early without DELETE."""
|
||||
from EvoScientist.sessions import PruningCheckpointer
|
||||
|
||||
tid = "tdsa_006"
|
||||
|
||||
async def _go():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
saver = PruningCheckpointer(conn, keep_per_ns=5)
|
||||
await saver.setup()
|
||||
# Only 3 checkpoints; keep=5. anchor_ids has 3 items,
|
||||
# 3 < 5, prune returns early — all survive untouched.
|
||||
specs = [
|
||||
("cp_001", None, None),
|
||||
("cp_002", "cp_001", None),
|
||||
("cp_003", "cp_002", None),
|
||||
]
|
||||
await self._insert_chain(conn, tid, specs)
|
||||
await saver._prune_after_put(tid, "")
|
||||
await conn.commit()
|
||||
return await self._surviving_ids(conn, tid)
|
||||
|
||||
survivors = _run(_go())
|
||||
assert survivors == ["cp_001", "cp_002", "cp_003"]
|
||||
|
||||
|
||||
class TestMigrationSweep(unittest.TestCase):
|
||||
"""Tests for the legacy-bloat migration sweep."""
|
||||
|
||||
@@ -933,7 +1445,7 @@ class TestMigrationSweep(unittest.TestCase):
|
||||
idx INTEGER NOT NULL,
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT,
|
||||
blob BLOB,
|
||||
value BLOB,
|
||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
||||
)
|
||||
"""
|
||||
@@ -1114,6 +1626,107 @@ class TestMigrationSweep(unittest.TestCase):
|
||||
|
||||
assert _run(_second())
|
||||
|
||||
def test_sweep_preserves_snapshot_ancestor(self):
|
||||
"""Migration sweep must apply the same DeltaChannel walk as steady-state.
|
||||
|
||||
Without this, legacy users upgrading to PR #231 would hit a
|
||||
one-shot silent truncation: the bloat sweep would naive-prune
|
||||
the ``_DeltaSnapshot`` seed out of long threads, then the
|
||||
``user_version`` marker locks the sweep so it never re-runs —
|
||||
leaving permanently empty ``/resume`` history.
|
||||
|
||||
Mirrors ``TestPruningCheckpointerDeltaChannel.test_preserves_
|
||||
snapshot_ancestor`` but drives via ``_run_migration_sweep``.
|
||||
"""
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
from EvoScientist.sessions import _run_migration_sweep
|
||||
|
||||
tid = "tsweep_delta"
|
||||
|
||||
async def _seed():
|
||||
import aiosqlite
|
||||
|
||||
serde = JsonPlusSerializer()
|
||||
snapshot_type, snapshot_blob = serde.dumps_typed(
|
||||
{"channel_values": {"messages": _DeltaSnapshot(value=[])}}
|
||||
)
|
||||
empty_type, empty_blob = serde.dumps_typed({"channel_values": {}})
|
||||
meta = json.dumps({"agent_name": AGENT_NAME})
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS checkpoints (
|
||||
thread_id TEXT NOT NULL,
|
||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||
checkpoint_id TEXT NOT NULL,
|
||||
parent_checkpoint_id TEXT,
|
||||
type TEXT,
|
||||
checkpoint BLOB,
|
||||
metadata TEXT NOT NULL DEFAULT '{}',
|
||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id)
|
||||
)
|
||||
"""
|
||||
)
|
||||
await conn.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS 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)
|
||||
)
|
||||
"""
|
||||
)
|
||||
# cp_001..cp_002 delta-only, cp_003 carries the snapshot,
|
||||
# cp_004..cp_010 delta-only. Anchor window with keep=5 is
|
||||
# cp_006..cp_010; walk from cp_005 backward hits cp_003
|
||||
# (seed) → stop. Survivors: cp_003..cp_010.
|
||||
for i in range(1, 11):
|
||||
cid = f"cp_{i:03d}"
|
||||
parent = f"cp_{i - 1:03d}" if i > 1 else None
|
||||
if i == 3:
|
||||
ct, cb = snapshot_type, snapshot_blob
|
||||
else:
|
||||
ct, cb = empty_type, empty_blob
|
||||
await conn.execute(
|
||||
"INSERT INTO checkpoints (thread_id, checkpoint_ns, "
|
||||
"checkpoint_id, parent_checkpoint_id, type, checkpoint, metadata) "
|
||||
"VALUES (?, '', ?, ?, ?, ?, ?)",
|
||||
(tid, cid, parent, ct, cb, meta),
|
||||
)
|
||||
await conn.commit()
|
||||
|
||||
_run(_seed())
|
||||
pairs = _run(_run_migration_sweep(keep=5))
|
||||
assert pairs == 1
|
||||
|
||||
async def _survivors():
|
||||
import aiosqlite
|
||||
|
||||
async with aiosqlite.connect(self._db_path) as conn:
|
||||
async with conn.execute(
|
||||
"SELECT checkpoint_id FROM checkpoints WHERE thread_id = ? "
|
||||
"ORDER BY checkpoint_id ASC",
|
||||
(tid,),
|
||||
) as cur:
|
||||
return [r[0] for r in await cur.fetchall()]
|
||||
|
||||
survivors = _run(_survivors())
|
||||
# cp_001, cp_002 pruned. cp_003 (snapshot) + walk-through (cp_004,
|
||||
# cp_005) + anchors (cp_006..cp_010) survive.
|
||||
assert survivors == [f"cp_{i:03d}" for i in range(3, 11)]
|
||||
# Explicit absence of the pruned ids — guards against a future
|
||||
# refactor that accidentally returns an empty survivors list.
|
||||
assert "cp_001" not in survivors
|
||||
assert "cp_002" not in survivors
|
||||
|
||||
|
||||
class TestDbStats(unittest.TestCase):
|
||||
"""Tests for the read-only ``db_stats`` diagnostic helper."""
|
||||
@@ -1167,7 +1780,7 @@ class TestDbStats(unittest.TestCase):
|
||||
idx INTEGER NOT NULL,
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT,
|
||||
blob BLOB,
|
||||
value BLOB,
|
||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
||||
)
|
||||
"""
|
||||
@@ -1197,7 +1810,7 @@ class TestDbStats(unittest.TestCase):
|
||||
# (counted by db_stats via the JOIN to checkpoints).
|
||||
for i in range(4):
|
||||
await conn.execute(
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob) "
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) "
|
||||
"VALUES ('evo01', '', 'ce01_0', 't1', ?, 'ch', 'str', X'AA')",
|
||||
(i,),
|
||||
)
|
||||
@@ -1206,7 +1819,7 @@ class TestDbStats(unittest.TestCase):
|
||||
# checkpoints and filters by agent_name).
|
||||
for i in range(2):
|
||||
await conn.execute(
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob) "
|
||||
"INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) "
|
||||
"VALUES ('oth01', '', 'co01_0', 't2', ?, 'ch', 'str', X'BB')",
|
||||
(i,),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user