Add run-batch manifest replay

This commit is contained in:
Yiorgis Gozadinos 2026-06-22 11:39:12 +03:00
parent a3b542dba8
commit 7251d104c4
No known key found for this signature in database
6 changed files with 449 additions and 30 deletions

View file

@ -9,7 +9,7 @@ from typing import TYPE_CHECKING
from pydantic import BaseModel from pydantic import BaseModel
from haiku.rag.config import AppConfig from haiku.rag.config import AppConfig
from haiku.rag.ingester.batch import BatchDryRunReport from haiku.rag.ingester.batch import BatchDryRunReport, BatchManifest
from haiku.rag.ingester.metadata import build_providers, load_metadata_providers from haiku.rag.ingester.metadata import build_providers, load_metadata_providers
from haiku.rag.ingester.pollers.manager import PollerManager from haiku.rag.ingester.pollers.manager import PollerManager
from haiku.rag.ingester.queue.migrations import open_queue from haiku.rag.ingester.queue.migrations import open_queue
@ -179,6 +179,28 @@ 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:
assert self._pool is not None and self._jobs is not None
while True:
counts = await self._jobs.counts_by_status()
if not counts.get("queued") and not counts.get("claimed"):
break
if self._pool.live_workers == 0:
outstanding = counts.get("queued", 0) + counts.get("claimed", 0)
logger.error(
"All workers have died with %d outstanding job(s) "
"— aborting batch; stranded jobs will be reaped "
"on next start",
outstanding,
)
break
await asyncio.sleep(0.1)
completed = await self._jobs.counts_by_status_since(started_at)
return BatchReport(
succeeded=completed.get("succeeded", 0),
dead=completed.get("dead", 0),
)
async def serve(self, *, api: bool = True) -> None: async def serve(self, *, api: bool = True) -> None:
"""Run pollers + workers (and the HTTP API when enabled) until a """Run pollers + workers (and the HTTP API when enabled) until a
SIGINT/SIGTERM is received. Drains in-flight work on shutdown.""" SIGINT/SIGTERM is received. Drains in-flight work on shutdown."""
@ -250,26 +272,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()
while True: report = await self._drain_batch(started_at)
counts = await self._jobs.counts_by_status() report.failed_sweeps = failed_sweeps
if not counts.get("queued") and not counts.get("claimed"): return report
break
if self._pool.live_workers == 0:
outstanding = counts.get("queued", 0) + counts.get("claimed", 0)
logger.error(
"All workers have died with %d outstanding job(s) "
"— aborting batch; stranded jobs will be reaped "
"on next start",
outstanding,
)
break
await asyncio.sleep(0.1)
completed = await self._jobs.counts_by_status_since(started_at)
return BatchReport(
succeeded=completed.get("succeeded", 0),
dead=completed.get("dead", 0),
failed_sweeps=failed_sweeps,
)
finally: finally:
await self._stop_pool() await self._stop_pool()
await self._pollers.close_sources() await self._pollers.close_sources()
@ -282,6 +287,91 @@ 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:
"""Enqueue and drain a dry-run manifest without running a fresh
discovery sweep."""
if manifest.version != 1:
raise ValueError(f"Unsupported manifest version: {manifest.version}")
async with self._resources():
assert (
self._pollers is not None
and self._pool is not None
and self._jobs is not None
)
configured = {source.source_id for source in self._pollers.sources}
manifest_sources = {change.source_id for change in manifest.changes}
missing = sorted(manifest_sources - configured)
if missing:
await self._pollers.close_sources()
raise ValueError(
"Manifest references unconfigured source(s): " + ", ".join(missing)
)
pending = [
source_id
for source_id in sorted(manifest_sources)
if await self._jobs.has_pending(source_id)
]
if pending:
await self._pollers.close_sources()
raise ValueError(
"Cannot replay manifest while source(s) have pending work: "
+ ", ".join(pending)
)
seen: set[tuple[str, str]] = set()
duplicates: set[tuple[str, str]] = set()
for change in manifest.changes:
key = (change.source_id, change.uri)
if key in seen:
duplicates.add(key)
seen.add(key)
if duplicates:
await self._pollers.close_sources()
rendered = ", ".join(
f"{source_id}:{uri}" for source_id, uri in duplicates
)
raise ValueError(f"Manifest contains duplicate change(s): {rendered}")
default_max_attempts = self._config.ingester.workers.retry.max_attempts
max_attempts_by_source = {
poller.source_id: (
poller.config.retry.max_attempts
if poller.config.retry is not None
else default_max_attempts
)
for poller in self._pollers.pollers
}
for change in manifest.changes:
job = await self._jobs.enqueue(
change.source_id,
change.uri,
op=change.op,
revision=change.revision,
max_attempts=max_attempts_by_source[change.source_id],
extra={
"_manifest": {
"version": manifest.version,
"generated_at": manifest.generated_at.isoformat(),
"discovered_at": change.discovered_at.isoformat(),
}
},
)
if job is None:
await self._pollers.close_sources()
raise ValueError(
"Cannot replay manifest because a live job already exists "
f"for {change.source_id}:{change.uri}"
)
started_at = datetime.now(UTC)
await self._pool.start()
try:
return await self._drain_batch(started_at)
finally:
await self._stop_pool()
await self._pollers.close_sources()
async def _maybe_start_api(self, api: bool): async def _maybe_start_api(self, api: bool):
"""Spin up the FastAPI control plane on an asyncio task. Returns """Spin up the FastAPI control plane on an asyncio task. Returns
(task, server) or (None, None) when the API is disabled.""" (task, server) or (None, None) when the API is disabled."""

View file

@ -150,6 +150,11 @@ 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")
def _load_manifest(path: Path) -> BatchManifest:
data = yaml.safe_load(path.read_text(encoding="utf-8"))
return BatchManifest.model_validate(data)
@_cli.command("serve") @_cli.command("serve")
def serve( def serve(
db: Path | None = typer.Option( db: Path | None = typer.Option(
@ -212,12 +217,29 @@ def run_batch(
"-o", "-o",
help="Dry-run manifest path (defaults to manifest-<datestamp>.yaml).", help="Dry-run manifest path (defaults to manifest-<datestamp>.yaml).",
), ),
manifest: Path | None = typer.Option(
None,
"--manifest",
help="Replay a dry-run manifest instead of running discovery.",
),
) -> None: ) -> None:
"""Run one discover sweep across every configured source, drain the queue, """Run one discover sweep across every configured source, drain the queue,
then exit. New and changed resources are ingested, resources that vanished then exit. New and changed resources are ingested, resources that vanished
from a source are deleted. Exits non-zero if any job dead-letters or a from a source are deleted. Exits non-zero if any job dead-letters or a
source's sweep does not complete.""" source's sweep does not complete."""
asyncio.run(_run_batch(get_config(), db, dry_run=dry_run, output=output)) if manifest is not None and dry_run:
raise typer.BadParameter("--manifest cannot be combined with --dry-run")
if manifest is not None and output is not None:
raise typer.BadParameter("--output is only valid with --dry-run")
asyncio.run(
_run_batch(
get_config(),
db,
dry_run=dry_run,
output=output,
manifest_path=manifest,
)
)
async def _run_batch( async def _run_batch(
@ -226,6 +248,7 @@ async def _run_batch(
*, *,
dry_run: bool = False, dry_run: bool = False,
output: Path | None = None, output: Path | None = None,
manifest_path: Path | None = None,
) -> None: ) -> None:
db = _resolve_db_path(app_config, db_path) db = _resolve_db_path(app_config, db_path)
app = IngesterApp(config=app_config, db_path=db) app = IngesterApp(config=app_config, db_path=db)
@ -248,6 +271,20 @@ async def _run_batch(
) )
return return
if manifest_path is not None:
try:
manifest = _load_manifest(manifest_path)
report = await app.run_batch_from_manifest(manifest)
except ValueError as exc:
typer.echo(f"Error: {exc}")
raise typer.Exit(1) from exc
typer.echo(
f"Manifest batch complete: {report.succeeded} succeeded, {report.dead} dead"
)
if report.dead:
raise typer.Exit(1)
return
report = await app.run_batch() report = await app.run_batch()
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:

View file

@ -110,13 +110,18 @@ async def run_job(
), ),
): ):
try: try:
manifest_context = extra.get("_manifest")
if job.op is JobOp.DELETE: if job.op is JobOp.DELETE:
# An atomic-rename save can let a spurious DELETE win the # An atomic-rename save can let a spurious DELETE win the
# enqueue race while the file is mid-rewrite. If the resource # enqueue race while the file is mid-rewrite. If the resource
# is already back, skip the delete (it would blackhole a live # is already back, skip the delete (it would blackhole a live
# document) and let the next sweep re-ingest it. # document) and let the next sweep re-ingest it. Manifest
# replay intentionally follows the frozen dry-run changeset.
if manifest_context is None:
try: try:
source = resolve_configured_source(job.uri, job.source_id, sources) source = resolve_configured_source(
job.uri, job.source_id, sources
)
restored = await source.head(job.uri) is not None restored = await source.head(job.uri) is not None
except Exception: except Exception:
restored = False restored = False
@ -127,6 +132,16 @@ async def run_job(
await client.delete_document(doc.id) await client.delete_document(doc.id)
return JobResult(deleted=True) return JobResult(deleted=True)
if manifest_context is not None and job.revision is not None:
source = resolve_configured_source(job.uri, job.source_id, sources)
current_revision = await source.head(job.uri)
if current_revision != job.revision:
raise PermanentError(
"manifest revision is stale for "
f"{job.uri}: expected {job.revision!r}, "
f"current {current_revision!r}"
)
result = await client.create_document_from_source( result = await client.create_document_from_source(
job.uri, job.uri,
sources=sources, sources=sources,

View file

@ -68,6 +68,17 @@ def _fake_dry_run_app(report: BatchDryRunReport, monkeypatch) -> AsyncMock:
return fake return fake
def _fake_manifest_app(report: BatchReport, monkeypatch) -> AsyncMock:
fake = AsyncMock()
fake.run_batch_from_manifest.return_value = report
monkeypatch.setattr("haiku.rag.ingester.cli.IngesterApp", lambda **_: fake)
return fake
def _write_manifest(path) -> None:
path.write_text(yaml.safe_dump(_manifest().model_dump(mode="json")))
def test_run_batch_reports_and_exits_zero(monkeypatch): def test_run_batch_reports_and_exits_zero(monkeypatch):
fake = _fake_app(BatchReport(succeeded=3, dead=0), monkeypatch) fake = _fake_app(BatchReport(succeeded=3, dead=0), monkeypatch)
@ -152,6 +163,80 @@ def test_run_batch_dry_run_exits_nonzero_when_sweep_fails(monkeypatch, tmp_path)
assert not output.exists() assert not output.exists()
def test_run_batch_manifest_replays_manifest(monkeypatch, tmp_path):
manifest_path = tmp_path / "manifest.yaml"
_write_manifest(manifest_path)
fake = _fake_manifest_app(BatchReport(succeeded=2, dead=0), monkeypatch)
result = runner.invoke(
cli, ["run-batch", "--manifest", str(manifest_path), "--db", "x.lancedb"]
)
assert result.exit_code == 0
assert "Manifest batch complete: 2 succeeded, 0 dead" in result.output
fake.run_batch_from_manifest.assert_awaited_once()
loaded = fake.run_batch_from_manifest.await_args.args[0]
assert isinstance(loaded, BatchManifest)
assert loaded.changes[0].uri == "file:///a.md"
fake.run_batch.assert_not_awaited()
fake.run_batch_dry_run.assert_not_awaited()
def test_run_batch_manifest_exits_nonzero_when_dead(monkeypatch, tmp_path):
manifest_path = tmp_path / "manifest.yaml"
_write_manifest(manifest_path)
_fake_manifest_app(BatchReport(succeeded=1, dead=1), monkeypatch)
result = runner.invoke(cli, ["run-batch", "--manifest", str(manifest_path)])
assert result.exit_code == 1
assert "1 dead" in result.output
def test_run_batch_manifest_reports_validation_error(monkeypatch, tmp_path):
manifest_path = tmp_path / "manifest.yaml"
_write_manifest(manifest_path)
fake = AsyncMock()
fake.run_batch_from_manifest.side_effect = ValueError("bad manifest")
monkeypatch.setattr("haiku.rag.ingester.cli.IngesterApp", lambda **_: fake)
result = runner.invoke(cli, ["run-batch", "--manifest", str(manifest_path)])
assert result.exit_code == 1
assert "Error: bad manifest" in result.output
def test_run_batch_manifest_conflicts_with_dry_run(tmp_path):
manifest_path = tmp_path / "manifest.yaml"
_write_manifest(manifest_path)
result = runner.invoke(
cli, ["run-batch", "--manifest", str(manifest_path), "--dry-run"]
)
assert result.exit_code != 0
assert "--manifest cannot be combined with --dry-run" in result.output
def test_run_batch_manifest_conflicts_with_output(tmp_path):
manifest_path = tmp_path / "manifest.yaml"
_write_manifest(manifest_path)
result = runner.invoke(
cli,
[
"run-batch",
"--manifest",
str(manifest_path),
"--output",
str(tmp_path / "out.yaml"),
],
)
assert result.exit_code != 0
assert "--output is only valid with --dry-run" in result.output
# --- serve --- # --- serve ---

View file

@ -16,6 +16,7 @@ def _job(
*, *,
op: JobOp = JobOp.UPSERT, op: JobOp = JobOp.UPSERT,
uri: str = "https://example.com/a.pdf", uri: str = "https://example.com/a.pdf",
revision: str | None = None,
extra: dict | None = None, extra: dict | None = None,
attempts: int = 0, attempts: int = 0,
) -> Job: ) -> Job:
@ -25,6 +26,7 @@ def _job(
source_id="src", source_id="src",
uri=uri, uri=uri,
op=op, op=op,
revision=revision,
status=JobStatus.CLAIMED, status=JobStatus.CLAIMED,
attempts=attempts, attempts=attempts,
max_attempts=5, max_attempts=5,
@ -255,6 +257,22 @@ async def test_delete_skipped_when_resource_restored_on_source():
client.delete_document.assert_not_awaited() client.delete_document.assert_not_awaited()
@pytest.mark.asyncio
async def test_manifest_delete_proceeds_when_resource_restored_on_source():
client = _mock_client()
client.get_document_by_uri.return_value = Document(id="doc-9", content="", uri="u")
sources: list[Source] = [_StubSource("src", "12345")]
result = await run_job(
client,
_job(op=JobOp.DELETE, extra={"_manifest": {"version": 1}}),
sources=sources,
)
assert result.deleted is True
client.delete_document.assert_awaited_once_with("doc-9")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_proceeds_when_resource_absent_on_source(): async def test_delete_proceeds_when_resource_absent_on_source():
client = _mock_client() client = _mock_client()
@ -267,6 +285,42 @@ async def test_delete_proceeds_when_resource_absent_on_source():
client.delete_document.assert_awaited_once_with("doc-9") client.delete_document.assert_awaited_once_with("doc-9")
@pytest.mark.asyncio
async def test_manifest_upsert_rejects_stale_revision_before_fetch():
client = _mock_client()
sources: list[Source] = [_StubSource("src", "r2")]
with pytest.raises(PermanentError, match="manifest revision is stale"):
await run_job(
client,
_job(revision="r1", extra={"_manifest": {"version": 1}}),
sources=sources,
)
client.create_document_from_source.assert_not_awaited()
@pytest.mark.asyncio
async def test_manifest_upsert_ingests_when_revision_matches():
client = _mock_client()
client.create_document_from_source.return_value = Document(
id="doc-42",
content="x",
uri="https://example.com/a.pdf",
metadata={"md5": "abcd", "source_revision": "r1"},
)
sources: list[Source] = [_StubSource("src", "r1")]
result = await run_job(
client,
_job(revision="r1", extra={"_manifest": {"version": 1}}),
sources=sources,
)
assert result.document_id == "doc-42"
client.create_document_from_source.assert_awaited_once()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_proceeds_when_source_unresolvable(): async def test_delete_proceeds_when_source_unresolvable():
"""No configured source for the job: the probe can't run, so the delete """No configured source for the job: the probe can't run, so the delete

View file

@ -5,6 +5,7 @@ pruning, not embedding."""
import asyncio import asyncio
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from datetime import UTC, datetime
from pathlib import Path from pathlib import Path
from unittest.mock import AsyncMock from unittest.mock import AsyncMock
from urllib.parse import unquote, urlparse from urllib.parse import unquote, urlparse
@ -23,6 +24,7 @@ from haiku.rag.config import (
WorkerConfig, WorkerConfig,
) )
from haiku.rag.ingester.app import IngesterApp from haiku.rag.ingester.app import IngesterApp
from haiku.rag.ingester.batch import BatchChange, BatchManifest
from haiku.rag.ingester.pollers.manager import PollerManager from haiku.rag.ingester.pollers.manager import PollerManager
from haiku.rag.ingester.queue.migrations import open_queue from haiku.rag.ingester.queue.migrations import open_queue
from haiku.rag.ingester.queue.models import JobOp from haiku.rag.ingester.queue.models import JobOp
@ -86,6 +88,10 @@ def _mock_client() -> AsyncMock:
return client return client
def _manifest(*changes: BatchChange) -> BatchManifest:
return BatchManifest(generated_at=datetime.now(UTC), changes=list(changes))
@pytest.fixture @pytest.fixture
def use_client(monkeypatch): def use_client(monkeypatch):
"""Make IngesterApp's internally-created HaikuRAG resolve to the given """Make IngesterApp's internally-created HaikuRAG resolve to the given
@ -278,6 +284,138 @@ async def test_run_batch_dry_run_reports_manifest_without_mutating_queue(tmp_pat
await engine.dispose() await engine.dispose()
@pytest.mark.asyncio
async def test_run_batch_from_manifest_drains_changes_without_sweeping(
tmp_path, use_client, monkeypatch
):
(tmp_path / "a.md").write_text("hello")
revision = str((tmp_path / "a.md").stat().st_mtime_ns)
client = _mock_client()
use_client(client)
sweep_all = AsyncMock(side_effect=AssertionError("manifest replay must not sweep"))
monkeypatch.setattr(PollerManager, "sweep_all", sweep_all)
report = await IngesterApp(
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest(
_manifest(
BatchChange(
op=JobOp.UPSERT,
source_id="local",
uri=(tmp_path / "a.md").as_uri(),
revision=revision,
discovered_at=datetime.now(UTC),
)
)
)
assert report.succeeded == 1
assert report.dead == 0
client.create_document_from_source.assert_awaited_once()
sweep_all.assert_not_awaited()
@pytest.mark.asyncio
async def test_run_batch_from_manifest_rejects_stale_upsert_revision(
tmp_path, use_client
):
(tmp_path / "a.md").write_text("hello")
client = _mock_client()
use_client(client)
report = await IngesterApp(
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest(
_manifest(
BatchChange(
op=JobOp.UPSERT,
source_id="local",
uri=(tmp_path / "a.md").as_uri(),
revision="stale",
discovered_at=datetime.now(UTC),
)
)
)
assert report.succeeded == 0
assert report.dead == 1
client.create_document_from_source.assert_not_awaited()
@pytest.mark.asyncio
async def test_run_batch_from_manifest_delete_uses_manifest_even_if_file_reappears(
tmp_path, use_client
):
path = tmp_path / "gone.md"
path.write_text("back")
client = _mock_client()
use_client(client)
report = await IngesterApp(
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest(
_manifest(
BatchChange(
op=JobOp.DELETE,
source_id="local",
uri=path.as_uri(),
discovered_at=datetime.now(UTC),
)
)
)
assert report.succeeded == 1
assert report.dead == 0
client.delete_document.assert_awaited_once()
@pytest.mark.asyncio
async def test_run_batch_from_manifest_rejects_pending_work(tmp_path, use_client):
(tmp_path / "a.md").write_text("hello")
config = _config(tmp_path)
client = _mock_client()
use_client(client)
engine = await open_queue(config.ingester.queue)
try:
jobs = JobRepo(engine)
await jobs.enqueue("local", (tmp_path / "a.md").as_uri(), op=JobOp.UPSERT)
finally:
await engine.dispose()
with pytest.raises(ValueError, match="pending work"):
await IngesterApp(
config=config, db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest(
_manifest(
BatchChange(
op=JobOp.UPSERT,
source_id="local",
uri=(tmp_path / "a.md").as_uri(),
discovered_at=datetime.now(UTC),
)
)
)
@pytest.mark.asyncio
async def test_run_batch_from_manifest_rejects_duplicate_changes(tmp_path, use_client):
path = tmp_path / "a.md"
path.write_text("hello")
client = _mock_client()
use_client(client)
change = BatchChange(
op=JobOp.UPSERT,
source_id="local",
uri=path.as_uri(),
discovered_at=datetime.now(UTC),
)
with pytest.raises(ValueError, match="duplicate change"):
await IngesterApp(
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest(_manifest(change, change))
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_batch_aborts_when_all_workers_die( async def test_run_batch_aborts_when_all_workers_die(
tmp_path, use_client, monkeypatch, caplog tmp_path, use_client, monkeypatch, caplog