Show progress for run-batch drains
This commit is contained in:
parent
6e85bbbafe
commit
baf8decb27
5 changed files with 171 additions and 11 deletions
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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."""
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue