Implement PruningCheckpointer for efficient checkpoint management and… (#194)

* Implement PruningCheckpointer for efficient checkpoint management and add comprehensive tests

- Introduced `PruningCheckpointer` to manage checkpoint pruning after each `aput()`, ensuring only the latest checkpoints are retained based on a configurable limit.
- Added migration sweep functionality to clean up legacy checkpoints and prevent database bloat.
- Enhanced `get_checkpointer()` to utilize the new `PruningCheckpointer` and trigger migration sweeps when necessary.
- Developed a suite of integration tests for `PruningCheckpointer`, covering various scenarios including pruning behavior, concurrent writes, and retention policies.
- Implemented tests for migration sweep functionality, ensuring proper partitioning and user version management.
- Added diagnostic helper `db_stats` to provide insights into the database state, including thread and checkpoint counts.

* feat(sessions): enhance pruning logic to handle legacy DBs without writes table

* fix(tests): prevent atexit hook leakage in TestMigrationSweep

* feat(tests): enhance TestPruningCheckpointer to validate put+prune serialization

* feat(tests): refactor mock path implementation for get_db_path in test cases
This commit is contained in:
Xi Zhang
2026-04-28 22:26:59 +02:00
committed by GitHub
parent 56cc2fef85
commit 50719ef256
5 changed files with 1392 additions and 5 deletions
+7
View File
@@ -51,3 +51,10 @@ app.add_typer(mcp_app, name="mcp")
# Channel subcommand group
channel_app = typer.Typer(help="Channel management commands")
app.add_typer(channel_app, name="channel")
# Sessions subcommand group — diagnostic tools for the LangGraph checkpoint DB
sessions_app = typer.Typer(
help="Inspect and manage the sessions DB (~/.evoscientist/sessions.db)",
invoke_without_command=True,
)
app.add_typer(sessions_app, name="sessions")
+61 -1
View File
@@ -17,7 +17,7 @@ from rich.table import Table
from ..llm.context_window import DEFAULT_CONTEXT_WINDOW_FALLBACK, resolve_context_window
from ..paths import ensure_dirs, set_workspace_root
from ..stream.console import console
from ._app import app, channel_app, config_app, mcp_app
from ._app import app, channel_app, config_app, mcp_app, sessions_app
from ._constants import build_metadata
from .agent import (
_create_session_workspace,
@@ -1281,6 +1281,66 @@ def mcp_install(
_cmd_install_mcp(source or "")
# =============================================================================
# Sessions commands — read-only diagnostics for ~/.evoscientist/sessions.db
# =============================================================================
def _format_bytes(n: int) -> str:
"""Render a byte count as a human-readable string (KB / MB / GB)."""
if n < 1024:
return f"{n} B"
units = ["KB", "MB", "GB", "TB"]
size = float(n) / 1024.0
for unit in units:
if size < 1024.0:
return f"{size:.1f} {unit}"
size /= 1024.0
return f"{size:.1f} PB"
@sessions_app.callback(invoke_without_command=True)
def sessions_callback(ctx: typer.Context):
"""Inspect and manage the sessions DB.
Running ``EvoSci sessions`` with no subcommand defaults to ``stats``
so the bare command is informative rather than silent.
"""
if ctx.invoked_subcommand is None:
sessions_stats()
@sessions_app.command("stats")
def sessions_stats():
"""Show DB size, thread count, total checkpoints, top heaviest threads."""
import asyncio
from ..sessions import db_stats
try:
stats = asyncio.get_event_loop().run_until_complete(db_stats())
except RuntimeError:
stats = asyncio.new_event_loop().run_until_complete(db_stats())
table = Table(title="EvoScientist sessions DB", show_header=True)
table.add_column("Metric", style="cyan")
table.add_column("Value")
table.add_row("Path", stats["db_path"])
table.add_row("Size", _format_bytes(int(stats["size_bytes"])))
table.add_row("Threads", str(stats["thread_count"]))
table.add_row("Checkpoints", str(stats["checkpoint_count"]))
table.add_row("Writes", str(stats["write_count"]))
console.print(table)
if stats["top_threads"]:
top = Table(title="Heaviest threads (checkpoints per thread)")
top.add_column("thread_id", style="yellow")
top.add_column("checkpoints", justify="right")
for row in stats["top_threads"]:
top.add_row(str(row["thread_id"]), str(row["count"]))
console.print(top)
# =============================================================================
# Main callback (default behavior)
# =============================================================================
+4
View File
@@ -201,6 +201,9 @@ class EvoScientistConfig:
# Agent features
enable_ask_user: bool = True # Enable ask_user tool for agent-initiated questions
# Checkpoint pruning (sessions.db retention per (thread_id, checkpoint_ns))
checkpoint_keep_per_thread: int = 10
# DM access control policy
dm_policy: str = "allowlist"
@@ -389,6 +392,7 @@ _ENV_MAPPINGS = {
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
"use_responses_api": "EVOSCIENTIST_USE_RESPONSES_API",
"checkpoint_keep_per_thread": "EVOSCIENTIST_CHECKPOINT_KEEP_PER_THREAD",
}
+535 -4
View File
@@ -4,18 +4,33 @@ Provides thread CRUD operations, prefix-matched resume, and an async
context manager for the shared ``AsyncSqliteSaver`` checkpointer.
Adapted from upstream ``deepagents_cli/sessions.py``.
Per-step pruning:
LangGraph's checkpointer writes a full state snapshot per super-step,
causing unbounded growth (multi-GB sessions.db). EvoScientist never
reads historical checkpoints — resume always reads the latest, HITL
interrupts attach pending writes to the just-written row. So
``get_checkpointer()`` yields a ``PruningCheckpointer`` that prunes
older rows for the same ``(thread_id, checkpoint_ns)`` after every
``aput()``. The first-run migration sweep cleans up legacy bloat.
"""
import asyncio
import atexit
import logging
import uuid
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
import aiosqlite
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
_logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Monkey-patch aiosqlite for langgraph-checkpoint >= 2.1.0 compatibility
# ---------------------------------------------------------------------------
@@ -80,16 +95,219 @@ def generate_thread_id() -> str:
return uuid.uuid4().hex[:8]
# ---------------------------------------------------------------------------
# Checkpoint pruning
# ---------------------------------------------------------------------------
# Default kept when the caller cannot resolve config (tests, unit-init paths).
# Production callers use ``EvoScientistConfig.checkpoint_keep_per_thread``.
_DEFAULT_KEEP_PER_NS = 10
class PruningCheckpointer(AsyncSqliteSaver):
"""``AsyncSqliteSaver`` that prunes stale checkpoints after every ``aput()``.
After a successful ``aput()``, deletes rows in ``checkpoints`` and
``writes`` whose ``(thread_id, checkpoint_ns)`` matches the just-written
row but whose ``checkpoint_id`` is not among the ``keep_per_ns`` most
recent ids. The just-written row is always kept (it is the head of the
descending order and ``keep_per_ns >= 1`` is enforced).
Inherits from ``AsyncSqliteSaver`` (rather than wrapping it) so
LangGraph's ``compile()`` ``isinstance(x, BaseCheckpointSaver)`` check
succeeds. All other behavior — ``aget_tuple``, ``alist``,
``aput_writes``, ``adelete_thread``, ``setup``, the async context
manager protocol, the connection lock — is inherited unchanged.
HITL safety: pregel records pending writes (e.g. ``interrupt``) against
the checkpoint id returned by the most recent ``aput()``, which is
exactly the row we keep. Older rows can never receive new writes after
a newer ``aput()`` lands, so deleting them is provably safe.
Setting ``keep_per_ns <= 0`` disables pruning (escape hatch for debug).
"""
def __init__(
self,
conn: aiosqlite.Connection,
*,
keep_per_ns: int = _DEFAULT_KEEP_PER_NS,
serde: Any = None,
) -> None:
super().__init__(conn, serde=serde)
self._keep_per_ns = max(0, int(keep_per_ns))
# Outer lock guarantees ``super().aput()`` and ``_prune_after_put()``
# are atomic *as a pair*. Without this, a concurrent ``aput()`` on a
# different ``(thread_id, checkpoint_ns)`` could land between the
# two phases and squeeze the earlier caller's just-written row out
# of the top-N retention window (only matters when ``keep_per_ns``
# is small or N parallel writers race; harmless otherwise but the
# invariant "the just-written row is always kept" must hold).
self._aput_lock = asyncio.Lock()
@classmethod
@asynccontextmanager
async def from_conn_string_with_keep(
cls, conn_string: str, keep_per_ns: int = _DEFAULT_KEEP_PER_NS
) -> AsyncIterator["PruningCheckpointer"]:
"""Build a ``PruningCheckpointer`` from a SQLite connection string.
Mirrors ``AsyncSqliteSaver.from_conn_string`` but threads
``keep_per_ns`` into ``__init__``. The native ``from_conn_string``
classmethod cannot accept extra kwargs, so callers that need
retention control should use this method instead.
"""
async with aiosqlite.connect(conn_string) as conn:
yield cls(conn, keep_per_ns=keep_per_ns)
async def aput(
self,
config: Any,
checkpoint: Any,
metadata: Any,
new_versions: Any,
) -> Any:
"""Delegate to ``super().aput``, then prune older rows atomically.
Wraps both the inner write and the prune in ``self._aput_lock`` so
a concurrent ``aput()`` cannot squeeze this caller's just-written
row out of the top-N retention window. The inner ``self.lock``
(held by ``super().aput`` and by ``_prune_after_put``) is a
separate, finer-grained lock that protects the SQLite connection;
the outer lock here is about the put+prune pair invariant.
Pruning is best-effort: any exception is logged at WARNING and
swallowed so a transient SQLite error never fails the agent step.
"""
async with self._aput_lock:
result = await super().aput(config, checkpoint, metadata, new_versions)
if self._keep_per_ns <= 0:
return result
try:
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "") or ""
await self._prune_after_put(str(thread_id), str(checkpoint_ns))
except Exception as exc: # pragma: no cover - defensive
_logger.warning("checkpoint pruning failed: %s", exc, exc_info=True)
return result
async def _prune_after_put(self, thread_id: str, checkpoint_ns: str) -> None:
"""Run the two DELETE queries (writes first, then checkpoints).
Restricted to rows whose ``metadata.agent_name == AGENT_NAME``.
``json_extract(metadata, '$.agent_name') = ?`` evaluates to NULL
(and so fails the predicate) for any row whose metadata does not
carry an ``agent_name`` key — by design, those rows belong to
third-party LangGraph users and must never be pruned by us.
Runs through the saver's connection and lock for atomicity with
concurrent ``aput()`` calls on the same thread.
"""
keep = self._keep_per_ns
agent = AGENT_NAME
# writes first — we look up which checkpoint ids will be deleted, then
# delete the writes pointing at them. Doing checkpoints first would
# leave orphan writes whose `checkpoint_id` we can no longer resolve.
del_writes = (
"DELETE FROM writes "
"WHERE thread_id = ? AND checkpoint_ns = ? "
" AND checkpoint_id IN ("
" SELECT checkpoint_id FROM checkpoints "
" WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" AND checkpoint_id NOT IN ("
" SELECT checkpoint_id FROM checkpoints "
" WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" ORDER BY checkpoint_id DESC LIMIT ?"
" )"
" )"
)
del_checkpoints = (
"DELETE FROM checkpoints "
"WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" AND checkpoint_id NOT IN ("
" SELECT checkpoint_id FROM checkpoints "
" WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" ORDER BY checkpoint_id DESC LIMIT ?"
" )"
)
async with self.lock:
# Skip the writes DELETE on a legacy DB that only has the
# checkpoints table — without this guard, a missing ``writes``
# would raise sqlite3.OperationalError, the outer try/except
# in ``aput`` would log+swallow it, and pruning would silently
# stop working on the very databases this feature is meant to
# clean up. ``AsyncSqliteSaver.setup()`` normally creates both
# tables, but inherited DBs from older builds may lag.
if await _table_exists(self.conn, "writes"):
await self.conn.execute(
del_writes,
(
thread_id,
checkpoint_ns,
thread_id,
checkpoint_ns,
agent,
thread_id,
checkpoint_ns,
agent,
keep,
),
)
await self.conn.execute(
del_checkpoints,
(
thread_id,
checkpoint_ns,
agent,
thread_id,
checkpoint_ns,
agent,
keep,
),
)
await self.conn.commit()
# ---------------------------------------------------------------------------
# Checkpointer context manager
# ---------------------------------------------------------------------------
def _resolve_keep_per_ns() -> int:
"""Resolve the retention count from EvoScientistConfig, with safe fallback."""
try:
from .config import get_effective_config
return max(0, int(get_effective_config().checkpoint_keep_per_thread))
except Exception: # pragma: no cover - defensive (config import errors)
return _DEFAULT_KEEP_PER_NS
@asynccontextmanager
async def get_checkpointer() -> AsyncIterator[AsyncSqliteSaver]:
"""Yield an ``AsyncSqliteSaver`` connected to the global sessions DB."""
async with AsyncSqliteSaver.from_conn_string(str(get_db_path())) as cp:
yield cp
async def get_checkpointer() -> AsyncIterator[PruningCheckpointer]:
"""Yield a pruning-enabled checkpointer connected to the sessions DB.
Wraps ``AsyncSqliteSaver`` with ``PruningCheckpointer`` so every
super-step trims the per-(thread, ns) history to ``keep_per_ns``. The
retention count is read from ``EvoScientistConfig`` at context entry;
setting it to 0 disables pruning entirely.
Also opportunistically kicks the legacy-bloat migration sweep as a
background task on first entry of an oversized DB. The sweep is a
no-op once ``PRAGMA user_version`` has been bumped, so subsequent
invocations cost nothing.
"""
keep = _resolve_keep_per_ns()
async with PruningCheckpointer.from_conn_string_with_keep(
str(get_db_path()), keep_per_ns=keep
) as saver:
# Fire-and-forget; runs concurrent with the agent on the same loop.
maybe_kick_migration_sweep(keep)
yield saver
# ---------------------------------------------------------------------------
@@ -438,3 +656,316 @@ async def get_thread_messages(thread_id: str) -> list:
messages = channel_values.get("messages", [])
event = channel_values.get("_summarization_event")
return _apply_summarization_event(messages, event)
# ---------------------------------------------------------------------------
# Migration sweep & VACUUM (one-time legacy cleanup)
# ---------------------------------------------------------------------------
# PRAGMA user_version is a 32-bit int slot in the SQLite file header. We
# bump this to 1 once the legacy-bloat sweep has run successfully so it
# never runs again. Future structural migrations can use 2, 3, ...
_MIGRATION_VERSION = 1
# Threshold below which the sweep is skipped (DB is already small enough
# that legacy bloat is not the user's problem). 100 MB is chosen so a
# normally-pruned DB after a few months of use never triggers the sweep,
# while the 2.6 GB pathology is comfortably above the line.
_MIGRATION_THRESHOLD_BYTES = 100 * 1024 * 1024
# Inter-pair sleep so the sweep yields to the agent loop and never spikes
# CPU on a large DB. Tunable for tests via monkeypatch.
_SWEEP_YIELD_SECONDS = 0.0
async def _get_user_version(conn: aiosqlite.Connection) -> int:
async with conn.execute("PRAGMA user_version") as cur:
row = await cur.fetchone()
return int(row[0]) if row else 0
async def _set_user_version(conn: aiosqlite.Connection, version: int) -> None:
# PRAGMAs cannot be parameter-bound; the integer is interpolated safely
# because we control the value (constant int).
await conn.execute(f"PRAGMA user_version = {int(version)}")
await conn.commit()
async def _needs_migration() -> bool:
"""Return True if the legacy-bloat sweep should run now.
True iff the DB exists, is larger than ``_MIGRATION_THRESHOLD_BYTES``,
and ``PRAGMA user_version`` is below ``_MIGRATION_VERSION``.
"""
db_path = get_db_path()
if not db_path.exists():
return False
try:
size = db_path.stat().st_size
except OSError:
return False
if size < _MIGRATION_THRESHOLD_BYTES:
return False
try:
async with aiosqlite.connect(str(db_path), timeout=30.0) as conn:
if not await _table_exists(conn, "checkpoints"):
return False
return await _get_user_version(conn) < _MIGRATION_VERSION
except aiosqlite.Error:
return False
async def _run_migration_sweep(keep: int) -> int:
"""Prune all ``(thread_id, checkpoint_ns)`` pairs to ``keep`` rows each.
Iterates pairs in deterministic order, applies the same DELETE pattern
the per-step pruner uses, and yields to the event loop between pairs
so the agent stays responsive. On success bumps ``PRAGMA user_version``
so the sweep never reruns.
Returns the number of pairs pruned.
"""
if keep <= 0:
return 0
db_path = str(get_db_path())
pairs_pruned = 0
async with aiosqlite.connect(db_path, timeout=60.0) as conn:
if not await _table_exists(conn, "checkpoints"):
return 0
if await _get_user_version(conn) >= _MIGRATION_VERSION:
return 0
# ``writes`` is optional on legacy DBs: skip the writes DELETE
# if the table is absent rather than aborting the whole sweep.
has_writes = await _table_exists(conn, "writes")
async with conn.execute(
"SELECT DISTINCT thread_id, checkpoint_ns FROM checkpoints "
"WHERE json_extract(metadata, '$.agent_name') = ?",
(AGENT_NAME,),
) as cur:
pairs = await cur.fetchall()
del_writes = (
"DELETE FROM writes "
"WHERE thread_id = ? AND checkpoint_ns = ? "
" AND checkpoint_id IN ("
" SELECT checkpoint_id FROM checkpoints "
" WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" AND checkpoint_id NOT IN ("
" SELECT checkpoint_id FROM checkpoints "
" WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" ORDER BY checkpoint_id DESC LIMIT ?"
" )"
" )"
)
del_checkpoints = (
"DELETE FROM checkpoints "
"WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" AND checkpoint_id NOT IN ("
" SELECT checkpoint_id FROM checkpoints "
" WHERE thread_id = ? AND checkpoint_ns = ? "
" AND json_extract(metadata, '$.agent_name') = ? "
" ORDER BY checkpoint_id DESC LIMIT ?"
" )"
)
for thread_id, checkpoint_ns in pairs:
ns = checkpoint_ns or ""
if has_writes:
await conn.execute(
del_writes,
(
thread_id,
ns,
thread_id,
ns,
AGENT_NAME,
thread_id,
ns,
AGENT_NAME,
keep,
),
)
await conn.execute(
del_checkpoints,
(
thread_id,
ns,
AGENT_NAME,
thread_id,
ns,
AGENT_NAME,
keep,
),
)
await conn.commit()
pairs_pruned += 1
if _SWEEP_YIELD_SECONDS >= 0:
await asyncio.sleep(_SWEEP_YIELD_SECONDS)
await _set_user_version(conn, _MIGRATION_VERSION)
# Schedule VACUUM at process exit (must run after the long-lived saver
# connection closes to acquire the exclusive lock VACUUM requires).
# Pass ``db_path`` explicitly so test-time monkeypatches of
# ``get_db_path`` don't leak into atexit and hit the real DB.
_schedule_vacuum_atexit(db_path)
return pairs_pruned
_vacuum_scheduled = False
def _schedule_vacuum_atexit(db_path: str) -> None:
"""Register the atexit VACUUM hook exactly once per process.
Captures ``db_path`` at registration time so the hook always operates
on the path that was current when the sweep ran. This matters in
tests, where ``get_db_path`` is monkey-patched to a temp file but the
patch has unwound by the time atexit fires — re-resolving at exit
would point at the user's real ``sessions.db`` and trigger an
unwanted VACUUM on production data.
"""
global _vacuum_scheduled
if _vacuum_scheduled:
return
_vacuum_scheduled = True
atexit.register(_atexit_vacuum, db_path)
def _atexit_vacuum(db_path: str) -> None:
"""Run ``VACUUM`` synchronously at process exit on the captured path.
Uses stdlib ``sqlite3`` (atexit can't await aiosqlite). Best-effort:
swallow any error since this runs during shutdown when stderr may be
closed.
"""
import os
import sqlite3
if not os.path.exists(db_path):
return
try:
with sqlite3.connect(db_path, timeout=60.0) as conn:
# VACUUM cannot run inside a transaction; sqlite3 starts one
# implicitly on the first execute, so isolation_level=None ensures
# we are in autocommit mode for the VACUUM statement.
conn.isolation_level = None
conn.execute("VACUUM")
except sqlite3.Error as exc:
# Best-effort during shutdown: stderr may already be closed, but
# try to log so a persistent VACUUM failure is at least diagnosable
# from the next session. Swallow any logging error in turn.
try:
_logger.warning("VACUUM at exit failed: %s", exc)
except Exception:
pass
def maybe_kick_migration_sweep(keep: int) -> asyncio.Task | None:
"""Schedule the legacy-bloat sweep as a background task if needed.
Returns the spawned ``asyncio.Task`` (caller does NOT await — sweep
runs concurrent with the agent), or ``None`` if migration is not
needed. Safe to call multiple times: the task itself rechecks the
user_version marker so a duplicate fire is a no-op.
"""
if keep <= 0:
return None
try:
loop = asyncio.get_event_loop()
except RuntimeError:
return None
async def _runner() -> None:
try:
if not await _needs_migration():
return
pairs = await _run_migration_sweep(keep)
if pairs:
_logger.info(
"checkpoint migration sweep pruned %d (thread, ns) pairs; "
"VACUUM will run at exit",
pairs,
)
except Exception as exc: # pragma: no cover - defensive
_logger.warning("migration sweep failed: %s", exc, exc_info=True)
return loop.create_task(_runner())
# ---------------------------------------------------------------------------
# DB stats (read-only diagnostic for `EvoSci sessions stats`)
# ---------------------------------------------------------------------------
async def db_stats(top_n: int = 5) -> dict[str, Any]:
"""Return read-only diagnostics about the sessions DB.
Keys: ``db_path`` (str), ``size_bytes`` (int, file size or 0 if absent),
``thread_count`` (int, EvoScientist threads), ``checkpoint_count``
(int, EvoScientist rows), ``write_count`` (int, EvoScientist writes
only — scoped via JOIN to ``checkpoints.metadata.agent_name`` so
co-located non-EvoSci agents are excluded), ``top_threads`` (list of
dicts with ``thread_id`` and ``count``, sorted desc).
"""
db_path = get_db_path()
size = db_path.stat().st_size if db_path.exists() else 0
out: dict[str, Any] = {
"db_path": str(db_path),
"size_bytes": size,
"thread_count": 0,
"checkpoint_count": 0,
"write_count": 0,
"top_threads": [],
}
if not db_path.exists():
return out
try:
async with aiosqlite.connect(str(db_path), timeout=30.0) as conn:
if not await _table_exists(conn, "checkpoints"):
return out
async with conn.execute(
"SELECT COUNT(DISTINCT thread_id), COUNT(*) FROM checkpoints "
"WHERE json_extract(metadata, '$.agent_name') = ?",
(AGENT_NAME,),
) as cur:
row = await cur.fetchone()
if row:
out["thread_count"] = int(row[0] or 0)
out["checkpoint_count"] = int(row[1] or 0)
if await _table_exists(conn, "writes"):
# Scope writes to EvoScientist rows by joining against
# checkpoints — the ``writes`` table itself has no
# ``agent_name`` column, so a bare ``COUNT(*)`` would
# over-report when other LangGraph apps share this DB.
async with conn.execute(
"SELECT COUNT(*) FROM writes w "
"JOIN checkpoints c "
" ON c.thread_id = w.thread_id "
" AND c.checkpoint_ns = w.checkpoint_ns "
" AND c.checkpoint_id = w.checkpoint_id "
"WHERE json_extract(c.metadata, '$.agent_name') = ?",
(AGENT_NAME,),
) as cur:
row = await cur.fetchone()
if row:
out["write_count"] = int(row[0] or 0)
async with conn.execute(
"SELECT thread_id, COUNT(*) AS n FROM checkpoints "
"WHERE json_extract(metadata, '$.agent_name') = ? "
"GROUP BY thread_id ORDER BY n DESC LIMIT ?",
(AGENT_NAME, int(top_n)),
) as cur:
rows = await cur.fetchall()
out["top_threads"] = [
{"thread_id": r[0], "count": int(r[1])} for r in rows
]
except aiosqlite.Error:
# Read-only — corrupt/locked DB → return zeroed stats rather than crash.
return out
return out
+785
View File
@@ -1,5 +1,6 @@
"""Tests for EvoScientist.sessions — thread CRUD, ID generation, helpers."""
import asyncio
import json
import os
import tempfile
@@ -27,6 +28,24 @@ from EvoScientist.sessions import (
from tests.conftest import run_async as _run
def _mock_path(db_path: str):
"""Build a Path-like object for patching ``EvoScientist.sessions.get_db_path``.
Implements the subset of ``pathlib.Path`` that ``sessions.py`` actually
touches: ``__str__``, ``__fspath__``, ``exists``, ``stat``.
"""
return type(
"MockPath",
(),
{
"__str__": lambda s: db_path,
"__fspath__": lambda s: db_path,
"exists": lambda s: os.path.exists(db_path),
"stat": lambda s: os.stat(db_path),
},
)()
class TestGenerateThreadId(unittest.TestCase):
def test_length(self):
tid = generate_thread_id()
@@ -429,5 +448,771 @@ class TestThreadFunctions(unittest.TestCase):
assert "cp_evo_shared" not in remaining
class TestPruningCheckpointer(unittest.TestCase):
"""Integration tests for ``PruningCheckpointer`` against a real
``AsyncSqliteSaver`` backed by a temp SQLite file.
"""
def setUp(self):
self._tmpdir = tempfile.mkdtemp()
self._db_path = os.path.join(self._tmpdir, "prune.db")
def tearDown(self):
try:
os.unlink(self._db_path)
except OSError:
pass
try:
os.rmdir(self._tmpdir)
except OSError:
pass
def _run_with_wrapper(self, keep: int, body):
"""Open ``PruningCheckpointer`` against the temp DB on a single
loop, invoke ``body(saver)`` (an async callable), then close
cleanly.
Required because ``aiosqlite.Connection`` is bound to the event
loop it was opened on; reusing it across separate ``run_async``
calls raises ``ValueError("no active connection")``.
"""
from EvoScientist.sessions import PruningCheckpointer
async def _go():
async with PruningCheckpointer.from_conn_string_with_keep(
self._db_path, keep_per_ns=keep
) as saver:
await saver.setup()
return await body(saver)
return _run(_go())
@staticmethod
def _config(thread_id: str, ns: str = "") -> dict:
return {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": ns,
"checkpoint_id": None,
}
}
@staticmethod
def _checkpoint(cid: str, step: int = 0) -> dict:
# Minimal Checkpoint dict accepted by JsonPlusSerializer.dumps_typed.
return {
"v": 1,
"ts": "2026-01-01T00:00:00+00:00",
"id": cid,
"channel_values": {},
"channel_versions": {},
"versions_seen": {},
"pending_sends": [],
}
@staticmethod
def _metadata() -> dict:
return {"agent_name": AGENT_NAME, "step": 0, "writes": {}, "parents": {}}
def _row_count(self, thread_id: str, ns: str = "") -> int:
async def _count():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
async with conn.execute(
"SELECT COUNT(*) FROM checkpoints WHERE thread_id = ? AND checkpoint_ns = ?",
(thread_id, ns),
) as cur:
row = await cur.fetchone()
return int(row[0]) if row else 0
return _run(_count())
def test_aput_prunes_after_insert(self):
tid = "tprune01"
async def _body(wrapper):
for i in range(7):
await wrapper.aput(
self._config(tid),
self._checkpoint(f"cp_{i:04d}", step=i),
self._metadata(),
{},
)
self._run_with_wrapper(keep=3, body=_body)
assert self._row_count(tid) == 3
def test_aput_keeps_latest_for_resume(self):
"""After pruning, ``aget_tuple`` must return the just-written checkpoint."""
tid = "tresume1"
async def _body(wrapper):
last_cfg = None
for i in range(5):
last_cfg = await wrapper.aput(
self._config(tid),
self._checkpoint(f"cpr_{i:04d}", step=i),
self._metadata(),
{},
)
tuple_ = await wrapper.aget_tuple(
{"configurable": {"thread_id": tid, "checkpoint_ns": ""}}
)
return last_cfg, tuple_
last_cfg, tuple_ = self._run_with_wrapper(keep=2, body=_body)
assert last_cfg["configurable"]["checkpoint_id"] == "cpr_0004"
assert tuple_ is not None
assert tuple_.checkpoint["id"] == "cpr_0004"
def test_aput_writes_against_kept_checkpoint(self):
"""HITL safety: ``aput_writes`` after prune still attaches successfully."""
tid = "twrites1"
async def _body(wrapper):
last = None
for i in range(4):
last = await wrapper.aput(
self._config(tid),
self._checkpoint(f"cpw_{i:04d}", step=i),
self._metadata(),
{},
)
# Attach a write to the just-written checkpoint id (mimics how
# pregel stores ``interrupt`` pending writes).
await wrapper.aput_writes(last, [("__interrupt__", "v")], "task1")
return last
last_cfg = self._run_with_wrapper(keep=2, body=_body)
async def _check():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
async with conn.execute(
"SELECT COUNT(*) FROM writes WHERE thread_id = ? AND checkpoint_id = ?",
(tid, last_cfg["configurable"]["checkpoint_id"]),
) as cur:
row = await cur.fetchone()
return int(row[0]) if row else 0
assert _run(_check()) == 1
def test_aput_partitions_by_ns(self):
"""Two checkpoint namespaces are pruned independently."""
tid = "tns01"
async def _body(wrapper):
for i in range(4):
await wrapper.aput(
self._config(tid, ns=""),
self._checkpoint(f"main_{i:04d}", step=i),
self._metadata(),
{},
)
await wrapper.aput(
self._config(tid, ns="sub:1"),
self._checkpoint(f"sub_{i:04d}", step=i),
self._metadata(),
{},
)
self._run_with_wrapper(keep=2, body=_body)
assert self._row_count(tid, ns="") == 2
assert self._row_count(tid, ns="sub:1") == 2
def test_inherits_base_checkpoint_saver(self):
"""LangGraph's ``compile()`` requires ``isinstance(saver, BaseCheckpointSaver)``.
Inheriting from ``AsyncSqliteSaver`` (which inherits from
``BaseCheckpointSaver``) is what unblocks agent compilation.
"""
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
from EvoScientist.sessions import PruningCheckpointer
async def _body(saver):
assert isinstance(saver, BaseCheckpointSaver)
assert isinstance(saver, AsyncSqliteSaver)
assert isinstance(saver, PruningCheckpointer)
# Critical inherited attributes/methods.
assert saver.serde is not None
assert saver.lock is not None
assert saver.conn is not None
assert callable(saver.aget_tuple)
assert callable(saver.aput_writes)
self._run_with_wrapper(keep=2, body=_body)
def test_prune_failure_does_not_break_aput(self):
"""If pruning raises, ``aput`` still returns successfully."""
tid = "tfail01"
async def _body(wrapper):
async def _boom(*args, **kwargs):
raise RuntimeError("simulated prune failure")
wrapper._prune_after_put = _boom # type: ignore[assignment]
return await wrapper.aput(
self._config(tid),
self._checkpoint("cpf_0001", step=0),
self._metadata(),
{},
)
result = self._run_with_wrapper(keep=2, body=_body)
assert result["configurable"]["checkpoint_id"] == "cpf_0001"
def test_prune_keep_zero_disables(self):
"""``keep_per_ns=0`` is a no-op — all rows survive."""
tid = "tzero01"
async def _body(wrapper):
for i in range(4):
await wrapper.aput(
self._config(tid),
self._checkpoint(f"cz_{i:04d}", step=i),
self._metadata(),
{},
)
self._run_with_wrapper(keep=0, body=_body)
assert self._row_count(tid) == 4
def test_prune_preserves_other_agent(self):
"""A row with a different ``agent_name`` is never deleted."""
tid = "tother1"
async def _body(saver):
# Seed the OtherAgent row through the same connection so it
# shares the loop with the saver.
other_meta = json.dumps({"agent_name": "OtherAgent", "step": 0})
async with saver.lock:
await saver.conn.execute(
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) "
"VALUES (?, '', ?, ?)",
(tid, "cp_other_keep", other_meta),
)
await saver.conn.commit()
for i in range(5):
await saver.aput(
self._config(tid),
self._checkpoint(f"co_{i:04d}", step=i),
self._metadata(),
{},
)
self._run_with_wrapper(keep=2, body=_body)
# OtherAgent's row + 2 EvoScientist rows = 3 total
assert self._row_count(tid) == 3
async def _check_other():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
async with conn.execute(
"SELECT 1 FROM checkpoints WHERE thread_id = ? AND checkpoint_id = ?",
(tid, "cp_other_keep"),
) as cur:
return (await cur.fetchone()) is not None
assert _run(_check_other())
def test_keep_one_boundary(self):
"""``keep_per_ns=1`` keeps only the latest row, deletes the rest."""
tid = "tk1_001"
async def _body(saver):
for i in range(2):
await saver.aput(
self._config(tid),
self._checkpoint(f"k1_{i:04d}", step=i),
self._metadata(),
{},
)
self._run_with_wrapper(keep=1, body=_body)
assert self._row_count(tid) == 1
async def _which():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
async with conn.execute(
"SELECT checkpoint_id FROM checkpoints WHERE thread_id = ?",
(tid,),
) as cur:
row = await cur.fetchone()
return row[0] if row else None
# The newest write (highest checkpoint_id) is the one kept.
assert _run(_which()) == "k1_0001"
def test_concurrent_same_thread_aput_invariant(self):
"""Concurrent ``aput()`` calls cannot squeeze either caller's
just-written row out of the top-N retention window.
Directly validates put+prune serialization: gates ``_prune_after_put``
on the first call so we can launch the second ``aput()`` while
the first is paused mid-prune. The second call must be blocked
by ``_aput_lock`` — without that outer lock, the ``self.lock``
held by ``super().aput`` would not span the prune phase, and
the two callers would interleave with the buggy result.
"""
tid = "tcc_001"
async def _body(saver):
entered_prune = asyncio.Event()
release_prune = asyncio.Event()
orig_prune = saver._prune_after_put
async def _gated_prune(thread_id: str, checkpoint_ns: str):
if not entered_prune.is_set():
entered_prune.set()
await release_prune.wait()
await orig_prune(thread_id, checkpoint_ns)
saver._prune_after_put = _gated_prune # type: ignore[method-assign]
cfg_a = self._config(tid)
cfg_b = self._config(tid)
cp_a = self._checkpoint("cc_a", step=0)
cp_b = self._checkpoint("cc_b", step=1)
t1 = asyncio.create_task(saver.aput(cfg_a, cp_a, self._metadata(), {}))
await entered_prune.wait()
t2 = asyncio.create_task(saver.aput(cfg_b, cp_b, self._metadata(), {}))
await asyncio.sleep(0)
assert not t2.done() # verifies second call is blocked by outer lock
release_prune.set()
results = await asyncio.gather(t1, t2)
return results
results = self._run_with_wrapper(keep=1, body=_body)
# Whichever caller landed last is the one survivor; importantly,
# the row count is exactly 1 (no torn state where both rows
# disappeared or both survived).
assert self._row_count(tid) == 1
async def _winner():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
async with conn.execute(
"SELECT checkpoint_id FROM checkpoints WHERE thread_id = ?",
(tid,),
) as cur:
row = await cur.fetchone()
return row[0] if row else None
survivor = _run(_winner())
# The survivor must be one of the two we wrote, not some torn ID.
assert survivor in {"cc_a", "cc_b"}
# And both aput results must report a valid checkpoint_id (neither
# call raised mid-prune).
for r in results:
assert r["configurable"]["checkpoint_id"] in {"cc_a", "cc_b"}
def test_uuid_ordering_keeps_latest(self):
"""Uses langgraph's actual UUIDv6-shaped checkpoint IDs to confirm
``ORDER BY checkpoint_id DESC`` keeps the chronologically latest.
``checkpoint_id`` is set by pregel from ``uuid6.uuid6()``, which
is monotonic-by-time. Lexicographic sort of the canonical hex form
therefore matches creation order — but the prune SQL relies on
this, so we exercise it explicitly.
"""
# langgraph ships its own ``uuid6`` (60-bit timestamp + counter,
# canonical hex form is monotonic by time). Pregel uses this to
# mint checkpoint ids — the prune SQL relies on
# ``ORDER BY checkpoint_id DESC`` matching chronological order.
from langgraph.checkpoint.base.id import uuid6 as _uuid6
tid = "tuu_001"
async def _body(saver):
ids: list[str] = []
for i in range(5):
cid = str(_uuid6(clock_seq=i))
ids.append(cid)
cp = {
"v": 1,
"ts": "2026-01-01T00:00:00+00:00",
"id": cid,
"channel_values": {},
"channel_versions": {},
"versions_seen": {},
"pending_sends": [],
}
await saver.aput(self._config(tid), cp, self._metadata(), {})
return ids
ids = self._run_with_wrapper(keep=2, body=_body)
assert self._row_count(tid) == 2
async def _check():
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 DESC",
(tid,),
) as cur:
return [r[0] for r in await cur.fetchall()]
survivors = _run(_check())
# The two latest UUIDv6 ids — by chronological generation —
# must be the survivors. Lexicographic DESC ordering must match.
assert survivors == [ids[4], ids[3]]
class TestMigrationSweep(unittest.TestCase):
"""Tests for the legacy-bloat migration sweep."""
def setUp(self):
self._tmpdir = tempfile.mkdtemp()
self._db_path = os.path.join(self._tmpdir, "sweep.db")
# Patch get_db_path so all sessions.py helpers point at our temp DB.
self._patcher = patch(
"EvoScientist.sessions.get_db_path",
return_value=_mock_path(self._db_path),
)
self._patcher.start()
# Mock atexit.register so sweep-spawned hooks don't leak past the fixture.
self._atexit_patcher = patch("EvoScientist.sessions.atexit.register")
self._atexit_patcher.start()
import EvoScientist.sessions as _sessions_mod
self._prev_vacuum_scheduled = _sessions_mod._vacuum_scheduled
_sessions_mod._vacuum_scheduled = False
def tearDown(self):
self._patcher.stop()
self._atexit_patcher.stop()
import EvoScientist.sessions as _sessions_mod
_sessions_mod._vacuum_scheduled = self._prev_vacuum_scheduled
try:
os.unlink(self._db_path)
except OSError:
pass
try:
os.rmdir(self._tmpdir)
except OSError:
pass
def _seed(self, threads_x_ns_x_count: list[tuple[str, str, int]]):
async def _go():
import aiosqlite
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,
blob BLOB,
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
)
"""
)
meta = json.dumps({"agent_name": AGENT_NAME, "step": 0})
for tid, ns, n in threads_x_ns_x_count:
for i in range(n):
await conn.execute(
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) "
"VALUES (?, ?, ?, ?)",
(tid, ns, f"{tid}_{ns}_{i:04d}", meta),
)
await conn.commit()
_run(_go())
def _user_version(self) -> int:
async def _go():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
async with conn.execute("PRAGMA user_version") as cur:
row = await cur.fetchone()
return int(row[0]) if row else 0
return _run(_go())
def _row_count(self, thread_id: str, ns: str) -> int:
async def _go():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
async with conn.execute(
"SELECT COUNT(*) FROM checkpoints WHERE thread_id = ? AND checkpoint_ns = ?",
(thread_id, ns),
) as cur:
row = await cur.fetchone()
return int(row[0]) if row else 0
return _run(_go())
def test_sweep_partitions_threads_and_ns(self):
from EvoScientist.sessions import _run_migration_sweep
self._seed(
[
("t1", "", 8),
("t1", "sub:1", 6),
("t2", "", 4),
]
)
pairs = _run(_run_migration_sweep(keep=3))
assert pairs == 3
assert self._row_count("t1", "") == 3
assert self._row_count("t1", "sub:1") == 3
assert self._row_count("t2", "") == 3
def test_sweep_sets_user_version(self):
from EvoScientist.sessions import _MIGRATION_VERSION, _run_migration_sweep
self._seed([("ta", "", 5)])
assert self._user_version() == 0
_run(_run_migration_sweep(keep=2))
assert self._user_version() == _MIGRATION_VERSION
def test_sweep_skipped_when_marker_set(self):
from EvoScientist.sessions import (
_MIGRATION_VERSION,
_run_migration_sweep,
_set_user_version,
)
self._seed([("tb", "", 5)])
async def _bump():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
await _set_user_version(conn, _MIGRATION_VERSION)
_run(_bump())
# Already at marker → sweep is a no-op even though many rows exist.
pairs = _run(_run_migration_sweep(keep=2))
assert pairs == 0
assert self._row_count("tb", "") == 5
def test_needs_migration_below_threshold(self):
from EvoScientist.sessions import _needs_migration
# Empty DB (file doesn't exist yet) → False
assert not _run(_needs_migration())
# Tiny DB → False
self._seed([("tc", "", 1)])
assert not _run(_needs_migration())
def test_needs_migration_above_threshold(self):
"""Use monkeypatch on the threshold constant so tests stay fast."""
from EvoScientist import sessions as sessions_module
self._seed([("td", "", 3)])
with patch.object(sessions_module, "_MIGRATION_THRESHOLD_BYTES", 1):
# Tiny DB exceeds the 1-byte threshold → marker check kicks in.
assert _run(sessions_module._needs_migration())
def test_keep_zero_short_circuits_sweep(self):
from EvoScientist.sessions import _run_migration_sweep
self._seed([("te", "", 4)])
pairs = _run(_run_migration_sweep(keep=0))
assert pairs == 0
assert self._row_count("te", "") == 4
def test_sweep_handles_missing_writes_table(self):
"""Legacy DB with only ``checkpoints`` (no ``writes``) must still prune.
Regression test: the sweep used to unconditionally
``DELETE FROM writes`` and would abort on the first iteration
with ``no such table: writes``, leaving the bloat in place.
"""
from EvoScientist.sessions import _run_migration_sweep
# Seed creates both tables; drop ``writes`` to simulate legacy.
self._seed([("tw", "", 5)])
async def _drop_writes():
import aiosqlite
async with aiosqlite.connect(self._db_path) as conn:
await conn.execute("DROP TABLE writes")
await conn.commit()
_run(_drop_writes())
pairs = _run(_run_migration_sweep(keep=2))
assert pairs == 1
assert self._row_count("tw", "") == 2
class TestDbStats(unittest.TestCase):
"""Tests for the read-only ``db_stats`` diagnostic helper."""
def setUp(self):
self._tmpdir = tempfile.mkdtemp()
self._db_path = os.path.join(self._tmpdir, "stats.db")
self._patcher = patch(
"EvoScientist.sessions.get_db_path",
return_value=_mock_path(self._db_path),
)
self._patcher.start()
def tearDown(self):
self._patcher.stop()
try:
os.unlink(self._db_path)
except OSError:
pass
try:
os.rmdir(self._tmpdir)
except OSError:
pass
def _seed(self):
async def _go():
import aiosqlite
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,
blob BLOB,
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
)
"""
)
evo = json.dumps({"agent_name": AGENT_NAME, "step": 0})
other = json.dumps({"agent_name": "OtherAgent", "step": 0})
# 2 EvoScientist threads, 5 + 3 = 8 checkpoints
for i in range(5):
await conn.execute(
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) "
"VALUES (?, '', ?, ?)",
("evo01", f"ce01_{i}", evo),
)
for i in range(3):
await conn.execute(
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) "
"VALUES (?, '', ?, ?)",
("evo02", f"ce02_{i}", evo),
)
# 1 OtherAgent thread (excluded from EvoSci counts)
await conn.execute(
"INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) "
"VALUES (?, '', ?, ?)",
("oth01", "co01_0", other),
)
# 4 writes linked to an EvoScientist checkpoint
# (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) "
"VALUES ('evo01', '', 'ce01_0', 't1', ?, 'ch', 'str', X'AA')",
(i,),
)
# 2 writes linked to OtherAgent's checkpoint — must NOT
# be counted in ``write_count`` (db_stats joins to
# 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) "
"VALUES ('oth01', '', 'co01_0', 't2', ?, 'ch', 'str', X'BB')",
(i,),
)
await conn.commit()
_run(_go())
def test_stats_returns_evo_only_counts(self):
"""All counts (incl. ``write_count``) must scope to EvoScientist rows.
Regression for the previous bare ``COUNT(*) FROM writes`` which
over-reported when other LangGraph apps share the DB. The seed
fixture inserts 4 EvoSci writes and 2 OtherAgent writes; only the
4 should count.
"""
from EvoScientist.sessions import db_stats
self._seed()
stats = _run(db_stats())
assert stats["thread_count"] == 2
assert stats["checkpoint_count"] == 8 # OtherAgent's 1 row excluded
assert stats["write_count"] == 4 # 2 OtherAgent writes excluded
assert stats["size_bytes"] > 0
assert stats["db_path"].endswith("stats.db")
def test_stats_top_threads_ordered_desc(self):
from EvoScientist.sessions import db_stats
self._seed()
stats = _run(db_stats(top_n=5))
ids = [row["thread_id"] for row in stats["top_threads"]]
counts = [row["count"] for row in stats["top_threads"]]
# Sorted desc by count: evo01 (5) before evo02 (3); OtherAgent excluded
assert ids == ["evo01", "evo02"]
assert counts == [5, 3]
def test_stats_missing_db(self):
"""No DB on disk → returns zeroed stats, never raises."""
from EvoScientist.sessions import db_stats
# Don't seed — file doesn't exist.
stats = _run(db_stats())
assert stats["thread_count"] == 0
assert stats["checkpoint_count"] == 0
assert stats["write_count"] == 0
assert stats["size_bytes"] == 0
if __name__ == "__main__":
unittest.main()