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:
+139
-50
@@ -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`)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user