Show progress for run-batch drains

This commit is contained in:
Yiorgis Gozadinos 2026-06-22 12:51:38 +03:00
parent 6e85bbbafe
commit baf8decb27
No known key found for this signature in database
5 changed files with 171 additions and 11 deletions

View file

@ -4,6 +4,7 @@
### Added ### Added
- `haiku-ingester run-batch --dry-run` writes a YAML manifest of planned upserts/deletes without mutating queue jobs or `sync_state`; `run-batch --manifest <path>` replays that frozen changeset without another discovery sweep. - `haiku-ingester run-batch --dry-run` writes a YAML manifest of planned upserts/deletes without mutating queue jobs or `sync_state`; `run-batch --manifest <path>` replays that frozen changeset without another discovery sweep.
- `haiku-ingester run-batch` and `run-batch --manifest` show an interactive progress bar with ETA while draining queued jobs.
### Fixed ### Fixed

View file

@ -1,6 +1,7 @@
import asyncio import asyncio
import logging import logging
import signal import signal
from collections.abc import Callable
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from datetime import UTC, datetime from datetime import UTC, datetime
from pathlib import Path from pathlib import Path
@ -42,6 +43,20 @@ class BatchReport(BaseModel):
failed_sweeps: list[str] = [] failed_sweeps: list[str] = []
class BatchProgress(BaseModel):
"""Snapshot emitted while a one-shot batch drains queued work."""
total: int = 0
completed: int = 0
succeeded: int = 0
dead: int = 0
queued: int = 0
claimed: int = 0
BatchProgressCallback = Callable[[BatchProgress], None]
def _manifest_change_key(change: BatchChange) -> tuple[str, str, str, str | None]: def _manifest_change_key(change: BatchChange) -> tuple[str, str, str, str | None]:
return (change.source_id, change.uri, change.op.value, change.revision) return (change.source_id, change.uri, change.op.value, change.revision)
@ -190,14 +205,37 @@ class IngesterApp:
if landed: if landed:
logger.info("Drained %d cancel-cleanup release(s) before close", landed) logger.info("Drained %d cancel-cleanup release(s) before close", landed)
async def _drain_batch(self, started_at: datetime) -> BatchReport: async def _drain_batch(
self,
started_at: datetime,
*,
progress_callback: BatchProgressCallback | None = None,
) -> BatchReport:
assert self._pool is not None and self._jobs is not None assert self._pool is not None and self._jobs is not None
total = 0
while True: while True:
counts = await self._jobs.counts_by_status() counts = await self._jobs.batch_progress_counts_since(started_at)
if not counts.get("queued") and not counts.get("claimed"): queued = counts.get("queued", 0)
claimed = counts.get("claimed", 0)
outstanding = queued + claimed
succeeded = counts.get("succeeded", 0)
dead = counts.get("dead", 0)
completed_count = succeeded + dead
total = max(total, outstanding + completed_count)
if progress_callback is not None:
progress_callback(
BatchProgress(
total=total,
completed=min(completed_count, total),
succeeded=succeeded,
dead=dead,
queued=queued,
claimed=claimed,
)
)
if not outstanding:
break break
if self._pool.live_workers == 0: if self._pool.live_workers == 0:
outstanding = counts.get("queued", 0) + counts.get("claimed", 0)
logger.error( logger.error(
"All workers have died with %d outstanding job(s) " "All workers have died with %d outstanding job(s) "
"— aborting batch; stranded jobs will be reaped " "— aborting batch; stranded jobs will be reaped "
@ -265,7 +303,9 @@ class IngesterApp:
await self._stop_pool() await self._stop_pool()
await self._pollers.close_sources() await self._pollers.close_sources()
async def run_batch(self) -> BatchReport: async def run_batch(
self, *, progress_callback: BatchProgressCallback | None = None
) -> BatchReport:
"""Run one discover() sweep across every configured source, drain the """Run one discover() sweep across every configured source, drain the
queue to completion, then stop. Unlike `serve`, the periodic poller queue to completion, then stop. Unlike `serve`, the periodic poller
loops never start discovery is driven explicitly, so the run is loops never start discovery is driven explicitly, so the run is
@ -283,7 +323,9 @@ class IngesterApp:
await self._pool.start() await self._pool.start()
try: try:
failed_sweeps = await self._pollers.sweep_all() failed_sweeps = await self._pollers.sweep_all()
report = await self._drain_batch(started_at) report = await self._drain_batch(
started_at, progress_callback=progress_callback
)
report.failed_sweeps = failed_sweeps report.failed_sweeps = failed_sweeps
return report return report
finally: finally:
@ -298,7 +340,12 @@ class IngesterApp:
manifest, failed_sweeps = await self._pollers.dry_run_manifest() manifest, failed_sweeps = await self._pollers.dry_run_manifest()
return BatchDryRunReport(manifest=manifest, failed_sweeps=failed_sweeps) return BatchDryRunReport(manifest=manifest, failed_sweeps=failed_sweeps)
async def run_batch_from_manifest(self, manifest: BatchManifest) -> BatchReport: async def run_batch_from_manifest(
self,
manifest: BatchManifest,
*,
progress_callback: BatchProgressCallback | None = None,
) -> BatchReport:
"""Enqueue and drain a dry-run manifest without running a fresh """Enqueue and drain a dry-run manifest without running a fresh
discovery sweep.""" discovery sweep."""
if manifest.version != 1: if manifest.version != 1:
@ -396,7 +443,9 @@ class IngesterApp:
started_at = datetime.now(UTC) started_at = datetime.now(UTC)
await self._pool.start() await self._pool.start()
try: try:
return await self._drain_batch(started_at) return await self._drain_batch(
started_at, progress_callback=progress_callback
)
finally: finally:
await self._stop_pool() await self._stop_pool()
await self._pollers.close_sources() await self._pollers.close_sources()

View file

@ -1,11 +1,21 @@
import asyncio import asyncio
import sys import sys
from collections.abc import Iterator
from contextlib import contextmanager
from datetime import UTC, datetime from datetime import UTC, datetime
from pathlib import Path from pathlib import Path
import typer import typer
import yaml import yaml
from dotenv import find_dotenv, load_dotenv from dotenv import find_dotenv, load_dotenv
from rich.console import Console
from rich.progress import (
BarColumn,
Progress,
TextColumn,
TimeElapsedColumn,
TimeRemainingColumn,
)
from sqlalchemy import make_url from sqlalchemy import make_url
load_dotenv(find_dotenv(usecwd=True)) load_dotenv(find_dotenv(usecwd=True))
@ -18,7 +28,11 @@ from haiku.rag.config import ( # noqa: E402
load_yaml_config, load_yaml_config,
set_config, set_config,
) )
from haiku.rag.ingester.app import IngesterApp # noqa: E402 from haiku.rag.ingester.app import ( # noqa: E402
BatchProgress,
BatchProgressCallback,
IngesterApp,
)
from haiku.rag.ingester.batch import BatchManifest # noqa: E402 from haiku.rag.ingester.batch import BatchManifest # noqa: E402
from haiku.rag.ingester.queue.migrations import open_queue # noqa: E402 from haiku.rag.ingester.queue.migrations import open_queue # noqa: E402
from haiku.rag.logging import configure_cli_logging # noqa: E402 from haiku.rag.logging import configure_cli_logging # noqa: E402
@ -150,6 +164,47 @@ def _write_manifest(manifest: BatchManifest, path: Path) -> None:
path.write_text(yaml.safe_dump(data, sort_keys=False), encoding="utf-8") path.write_text(yaml.safe_dump(data, sort_keys=False), encoding="utf-8")
@contextmanager
def _batch_progress(description: str) -> Iterator[BatchProgressCallback | None]:
console = Console(file=sys.stdout)
if not console.is_terminal:
yield None
return
progress = Progress(
TextColumn("[progress.description]{task.description}"),
BarColumn(),
TextColumn("{task.completed}/{task.total}"),
TimeRemainingColumn(),
TimeElapsedColumn(),
console=console,
transient=True,
)
task_id = None
def _update(snapshot: BatchProgress) -> None:
nonlocal task_id
task_description = (
f"{description} ({snapshot.succeeded} ok, {snapshot.dead} dead)"
)
if task_id is None:
task_id = progress.add_task(
task_description,
total=snapshot.total,
completed=snapshot.completed,
)
return
progress.update(
task_id,
description=task_description,
total=snapshot.total,
completed=snapshot.completed,
)
with progress:
yield _update
def _load_manifest(path: Path) -> BatchManifest: def _load_manifest(path: Path) -> BatchManifest:
data = yaml.safe_load(path.read_text(encoding="utf-8")) data = yaml.safe_load(path.read_text(encoding="utf-8"))
return BatchManifest.model_validate(data) return BatchManifest.model_validate(data)
@ -276,7 +331,13 @@ async def _run_batch(
if manifest_path is not None: if manifest_path is not None:
try: try:
manifest = _load_manifest(manifest_path) manifest = _load_manifest(manifest_path)
report = await app.run_batch_from_manifest(manifest) with _batch_progress("Replaying manifest") as progress_callback:
if progress_callback is None:
report = await app.run_batch_from_manifest(manifest)
else:
report = await app.run_batch_from_manifest(
manifest, progress_callback=progress_callback
)
except ValueError as exc: except ValueError as exc:
typer.echo(f"Error: {exc}") typer.echo(f"Error: {exc}")
raise typer.Exit(1) from exc raise typer.Exit(1) from exc
@ -287,7 +348,11 @@ async def _run_batch(
raise typer.Exit(1) raise typer.Exit(1)
return return
report = await app.run_batch() with _batch_progress("Running batch") as progress_callback:
if progress_callback is None:
report = await app.run_batch()
else:
report = await app.run_batch(progress_callback=progress_callback)
typer.echo(f"Batch complete: {report.succeeded} succeeded, {report.dead} dead") typer.echo(f"Batch complete: {report.succeeded} succeeded, {report.dead} dead")
if report.failed_sweeps: if report.failed_sweeps:
typer.echo( typer.echo(

View file

@ -371,6 +371,27 @@ class JobRepo:
rows = (await conn.execute(query)).all() rows = (await conn.execute(query)).all()
return {status: n for status, n in rows} return {status: n for status, n in rows}
async def batch_progress_counts_since(self, since: datetime) -> dict[str, int]:
"""Counts for a one-shot batch progress snapshot.
Live rows are counted regardless of enqueue time because run-batch
drains the whole pending queue. Terminal rows are counted only when
this run completed them, matching BatchReport semantics.
"""
query = (
sa.select(jobs.c.status, sa.func.count().label("n"))
.where(
sa.or_(
jobs.c.status.in_(["queued", "claimed"]),
jobs.c.completed_at >= since.isoformat(),
)
)
.group_by(jobs.c.status)
)
async with self._engine.connect() as conn:
rows = (await conn.execute(query)).all()
return {status: n for status, n in rows}
async def count_succeeded_since(self, seconds: int) -> int: async def count_succeeded_since(self, seconds: int) -> int:
"""How many jobs reached `succeeded` in the last `seconds` seconds. """How many jobs reached `succeeded` in the last `seconds` seconds.
Drives the dashboard's rolling-throughput chips.""" Drives the dashboard's rolling-throughput chips."""

View file

@ -132,6 +132,30 @@ async def test_run_batch_drains_upserts(tmp_path, use_client):
} }
@pytest.mark.asyncio
async def test_run_batch_reports_progress(tmp_path, use_client):
(tmp_path / "a.md").write_text("hello")
(tmp_path / "b.md").write_text("world")
client = _mock_client()
use_client(client)
progress = []
report = await IngesterApp(
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
).run_batch(progress_callback=progress.append)
assert report.succeeded == 2
assert report.dead == 0
assert progress
assert progress[-1].total == 2
assert progress[-1].completed == 2
assert progress[-1].succeeded == 2
assert progress[-1].dead == 0
assert progress[-1].queued == 0
assert progress[-1].claimed == 0
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_batch_prunes_orphans(tmp_path, use_client): async def test_run_batch_prunes_orphans(tmp_path, use_client):
(tmp_path / "a.md").write_text("hello") (tmp_path / "a.md").write_text("hello")