Fix/sessions migration sweep race (#195)

* fix(sessions): ensure migration sweep runs before yielding checkpointer to prevent race conditions

* fix(sessions): enhance migration sweep with progress indication and ETA estimation
This commit is contained in:
Xi Zhang
2026-04-29 03:06:25 +02:00
committed by GitHub
parent 50719ef256
commit 7cfec02416
2 changed files with 180 additions and 50 deletions
+139 -50
View File
@@ -18,8 +18,9 @@ Per-step pruning:
import asyncio
import atexit
import logging
import math
import uuid
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Awaitable, Callable
from contextlib import asynccontextmanager
from datetime import UTC, datetime
from pathlib import Path
@@ -296,17 +297,66 @@ async def get_checkpointer() -> AsyncIterator[PruningCheckpointer]:
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.
Also runs the legacy-bloat migration sweep synchronously *before*
yielding the saver when needed. Sequencing sweep ahead of any agent
``aput()`` eliminates the SQLite file-lock contention that produced
"database is locked" when channel inbound raced the sweep mid-DELETE.
The sweep is gated by ``PRAGMA user_version`` so it runs at most
once across all future launches; subsequent invocations cost nothing.
On failure ``user_version`` is NOT bumped, so the next launch retries.
"""
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)
# The whole gate is wrapped in a broad try/except: any unexpected
# failure (incl. preview / status print) must degrade to "yield
# the saver, log a warning" — never to "hang startup".
if keep > 0:
try:
if await _needs_migration():
import time
from rich.console import Console
# stderr-bound so non-interactive callers redirecting
# stdout don't capture migration progress noise.
console = Console(stderr=True)
size_str, pair_count, size_bytes = await _migration_preview()
eta_str = _format_duration(_estimate_sweep_seconds(size_bytes))
console.print(
f"[dim]·[/dim] Compacting sessions DB "
f"([cyan]{size_str}[/cyan], "
f"[yellow]{pair_count} thread-namespace pairs[/yellow]"
f"[dim], est. ~{eta_str}[/dim])"
)
t0 = time.time()
with console.status(
"[dim]Compacting...[/dim]", spinner="dots"
) as status:
async def _on_progress(done: int, total: int) -> None:
elapsed = time.time() - t0
pct = (done * 100 // total) if total else 0
eta = (elapsed / done) * (total - done) if done > 0 else 0
status.update(
f"[dim]Compacting[/dim] "
f"[yellow]{pct}%[/yellow] "
f"[dim]({done}/{total} pairs · "
f"{_format_duration(elapsed)} elapsed · "
f"~{_format_duration(eta)} remaining)[/dim]"
)
await _run_migration_sweep(keep, progress_cb=_on_progress)
console.print(
f"[dim]·[/dim] [green]✓[/green] Compaction done in "
f"[green]{_format_duration(time.time() - t0)}[/green]"
)
except Exception as exc:
_logger.warning("migration sweep failed: %s", exc, exc_info=True)
yield saver
@@ -386,16 +436,6 @@ def _apply_summarization_event(messages: list, event: dict | None) -> list:
return [summary_message, *messages[cutoff_index:]]
async def _count_messages(
conn: aiosqlite.Connection,
thread_id: str,
serde: JsonPlusSerializer,
) -> int:
"""Count messages in the most recent checkpoint for *thread_id*."""
msgs = await _load_checkpoint_messages(conn, thread_id, serde)
return len(msgs)
def _extract_preview(messages: list, max_len: int = 50) -> str:
"""Extract the first human message as a preview string."""
for msg in messages:
@@ -691,6 +731,75 @@ async def _set_user_version(conn: aiosqlite.Connection, version: int) -> None:
await conn.commit()
def _format_duration(seconds: float) -> str:
"""Render a number of seconds as ``45s`` / ``2m 18s`` / ``1h 5m``.
Guards against ``inf`` / ``nan`` so a stray non-finite value (e.g. an
ETA computed before any pair completes) cannot raise ``OverflowError``
via ``int(float('inf'))``.
"""
if not math.isfinite(seconds):
return "?"
seconds = max(0, int(seconds))
if seconds < 60:
return f"{seconds}s"
if seconds < 3600:
m, s = divmod(seconds, 60)
return f"{m}m {s}s" if s else f"{m}m"
h, rem = divmod(seconds, 3600)
m = rem // 60
return f"{h}h {m}m" if m else f"{h}h"
def _estimate_sweep_seconds(size_bytes: int) -> int:
"""Rough ETA for a full migration sweep on the given DB size.
Empirical baseline: 14.53 GB ≈ 138s on a typical SSD, i.e. ~10 s/GB.
Used only for a user-facing pre-sweep hint — real time is reported
after completion so the estimate doesn't need to be precise.
"""
return max(2, round(size_bytes / (1024 * 1024 * 1024) * 9.5))
async def _migration_preview() -> tuple[str, int, int]:
"""Return ``(size_str, pair_count, size_bytes)`` for the pre-sweep status line.
``size_str`` is a human-readable file size (``"2.62 GB"`` / ``"364.3 MB"``).
``pair_count`` is the number of distinct ``(thread_id, checkpoint_ns)``
pairs the sweep will iterate. ``size_bytes`` is the raw byte count
(used by ``_estimate_sweep_seconds`` for the ETA). Best-effort:
returns zeros on any error.
"""
db_path = get_db_path()
try:
size_bytes = db_path.stat().st_size
except OSError:
size_bytes = 0
if size_bytes < 1024 * 1024:
size_str = f"{size_bytes / 1024:.1f} KB"
elif size_bytes < 1024 * 1024 * 1024:
size_str = f"{size_bytes / (1024 * 1024):.1f} MB"
else:
size_str = f"{size_bytes / (1024 * 1024 * 1024):.2f} GB"
pair_count = 0
try:
async with aiosqlite.connect(str(db_path), timeout=30.0) as conn:
if await _table_exists(conn, "checkpoints"):
async with conn.execute(
"SELECT COUNT(*) FROM ("
" SELECT DISTINCT thread_id, checkpoint_ns FROM checkpoints "
" WHERE json_extract(metadata, '$.agent_name') = ?"
")",
(AGENT_NAME,),
) as cur:
row = await cur.fetchone()
pair_count = int(row[0]) if row else 0
except aiosqlite.Error:
pass
return size_str, pair_count, size_bytes
async def _needs_migration() -> bool:
"""Return True if the legacy-bloat sweep should run now.
@@ -715,7 +824,10 @@ async def _needs_migration() -> bool:
return False
async def _run_migration_sweep(keep: int) -> int:
async def _run_migration_sweep(
keep: int,
progress_cb: Callable[[int, int], Awaitable[None]] | None = None,
) -> int:
"""Prune all ``(thread_id, checkpoint_ns)`` pairs to ``keep`` rows each.
Iterates pairs in deterministic order, applies the same DELETE pattern
@@ -723,6 +835,10 @@ async def _run_migration_sweep(keep: int) -> int:
so the agent stays responsive. On success bumps ``PRAGMA user_version``
so the sweep never reruns.
``progress_cb``: optional ``async (done: int, total: int) -> None``
fired after each pair is committed; used by callers to drive a live
progress indicator.
Returns the number of pairs pruned.
"""
if keep <= 0:
@@ -802,6 +918,11 @@ async def _run_migration_sweep(keep: int) -> int:
)
await conn.commit()
pairs_pruned += 1
if progress_cb is not None:
try:
await progress_cb(pairs_pruned, len(pairs))
except Exception: # pragma: no cover - never let UI break sweep
pass
if _SWEEP_YIELD_SECONDS >= 0:
await asyncio.sleep(_SWEEP_YIELD_SECONDS)
@@ -864,38 +985,6 @@ def _atexit_vacuum(db_path: str) -> None:
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`)
# ---------------------------------------------------------------------------
+41
View File
@@ -1073,6 +1073,47 @@ class TestMigrationSweep(unittest.TestCase):
assert pairs == 1
assert self._row_count("tw", "") == 2
def test_get_checkpointer_blocks_on_sweep_then_idempotent(self):
"""End-to-end: ``get_checkpointer()`` must run the sweep BEFORE
yielding the saver so a concurrent ``aput()`` can't race the
DELETEs. After the first call sets ``user_version=1``, subsequent
calls must skip the sweep entirely.
"""
from EvoScientist import sessions as sessions_module
from EvoScientist.sessions import (
_MIGRATION_VERSION,
get_checkpointer,
)
self._seed([("ge", "", 6)])
# Force the sweep to be needed regardless of file size.
with patch.object(sessions_module, "_MIGRATION_THRESHOLD_BYTES", 1):
# First entry: sweep should run, prune to keep=10 (default), and
# set user_version. With only 6 rows in one (thread, ns) pair,
# the prune is a no-op but user_version is still bumped.
async def _first():
async with get_checkpointer() as saver:
return saver is not None
assert _run(_first())
assert self._user_version() == _MIGRATION_VERSION
# Second entry: sweep must be skipped — patch _run_migration_sweep
# to raise so any accidental re-invocation fails the test loudly.
async def _exploding_sweep(*_args, **_kwargs):
raise AssertionError("sweep must not re-run after marker is set")
with patch.object(
sessions_module, "_run_migration_sweep", _exploding_sweep
):
async def _second():
async with get_checkpointer() as saver:
return saver is not None
assert _run(_second())
class TestDbStats(unittest.TestCase):
"""Tests for the read-only ``db_stats`` diagnostic helper."""