Allow resuming run-batch manifest replay

This commit is contained in:
Yiorgis Gozadinos 2026-06-22 12:22:40 +03:00
parent 68bbf94577
commit b0706204bd
No known key found for this signature in database
2 changed files with 133 additions and 15 deletions

View file

@ -9,10 +9,11 @@ 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, BatchManifest from haiku.rag.ingester.batch import BatchChange, 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
from haiku.rag.ingester.queue.models import Job, JobStatus
from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo
from haiku.rag.ingester.workers.pool import WorkerPool from haiku.rag.ingester.workers.pool import WorkerPool
from haiku.rag.ingester.workers.retry import RetryPolicy from haiku.rag.ingester.workers.retry import RetryPolicy
@ -22,6 +23,8 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_MANIFEST_EXTRA_KEY = "_manifest"
def _api_access_log_enabled() -> bool: def _api_access_log_enabled() -> bool:
"""Per-request access logging only when the haiku.rag logger is at DEBUG. """Per-request access logging only when the haiku.rag logger is at DEBUG.
@ -39,6 +42,14 @@ class BatchReport(BaseModel):
failed_sweeps: list[str] = [] failed_sweeps: list[str] = []
def _manifest_change_key(change: BatchChange) -> tuple[str, str, str, str | None]:
return (change.source_id, change.uri, change.op.value, change.revision)
def _manifest_job_key(job: Job) -> tuple[str, str, str, str | None]:
return (job.source_id, job.uri, job.op.value, job.revision)
class IngesterApp: class IngesterApp:
"""Top-level lifecycle for the production ingester. """Top-level lifecycle for the production ingester.
@ -307,15 +318,6 @@ class IngesterApp:
"Manifest references unconfigured source(s): " + ", ".join(missing) "Manifest references unconfigured source(s): " + ", ".join(missing)
) )
counts = await self._jobs.counts_by_status()
pending = counts.get("queued", 0) + counts.get("claimed", 0)
if pending:
await self._pollers.close_sources()
raise ValueError(
"Cannot replay manifest while the queue has pending work: "
f"{pending} queued/claimed job(s)"
)
seen: set[tuple[str, str]] = set() seen: set[tuple[str, str]] = set()
duplicates: set[tuple[str, str]] = set() duplicates: set[tuple[str, str]] = set()
for change in manifest.changes: for change in manifest.changes:
@ -330,6 +332,34 @@ class IngesterApp:
) )
raise ValueError(f"Manifest contains duplicate change(s): {rendered}") raise ValueError(f"Manifest contains duplicate change(s): {rendered}")
manifest_key = manifest.generated_at.isoformat()
manifest_change_keys = {
_manifest_change_key(change) for change in manifest.changes
}
live_jobs = [
*await self._jobs.list_jobs(status=JobStatus.QUEUED, limit=10_000),
*await self._jobs.list_jobs(status=JobStatus.CLAIMED, limit=10_000),
]
stale_jobs: list[Job] = []
live_manifest_keys: set[tuple[str, str, str, str | None]] = set()
for job in live_jobs:
extra = job.extra or {}
job_manifest = extra.get(_MANIFEST_EXTRA_KEY) or {}
key = _manifest_job_key(job)
if (
job_manifest.get("generated_at") != manifest_key
or key not in manifest_change_keys
):
stale_jobs.append(job)
continue
live_manifest_keys.add(key)
if stale_jobs:
await self._pollers.close_sources()
raise ValueError(
"Cannot replay manifest while the queue has non-manifest "
f"pending work: {len(stale_jobs)} queued/claimed job(s)"
)
default_max_attempts = self._config.ingester.workers.retry.max_attempts default_max_attempts = self._config.ingester.workers.retry.max_attempts
max_attempts_by_source = { max_attempts_by_source = {
poller.source_id: ( poller.source_id: (
@ -340,6 +370,8 @@ class IngesterApp:
for poller in self._pollers.pollers for poller in self._pollers.pollers
} }
for change in manifest.changes: for change in manifest.changes:
if _manifest_change_key(change) in live_manifest_keys:
continue
job = await self._jobs.enqueue( job = await self._jobs.enqueue(
change.source_id, change.source_id,
change.uri, change.uri,
@ -347,9 +379,9 @@ class IngesterApp:
revision=change.revision, revision=change.revision,
max_attempts=max_attempts_by_source[change.source_id], max_attempts=max_attempts_by_source[change.source_id],
extra={ extra={
"_manifest": { _MANIFEST_EXTRA_KEY: {
"version": manifest.version, "version": manifest.version,
"generated_at": manifest.generated_at.isoformat(), "generated_at": manifest_key,
"discovered_at": change.discovered_at.isoformat(), "discovered_at": change.discovered_at.isoformat(),
} }
}, },

View file

@ -370,7 +370,53 @@ async def test_run_batch_from_manifest_delete_uses_manifest_even_if_file_reappea
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_batch_from_manifest_rejects_pending_work(tmp_path, use_client): async def test_run_batch_from_manifest_resumes_same_manifest_work(tmp_path, use_client):
(tmp_path / "a.md").write_text("hello")
revision = str((tmp_path / "a.md").stat().st_mtime_ns)
config = _config(tmp_path)
client = _mock_client()
use_client(client)
manifest = _manifest(
BatchChange(
op=JobOp.UPSERT,
source_id="local",
uri=(tmp_path / "a.md").as_uri(),
revision=revision,
discovered_at=datetime.now(UTC),
)
)
engine = await open_queue(config.ingester.queue)
try:
jobs = JobRepo(engine)
await jobs.enqueue(
"local",
(tmp_path / "a.md").as_uri(),
op=JobOp.UPSERT,
revision=revision,
extra={
"_manifest": {
"version": manifest.version,
"generated_at": manifest.generated_at.isoformat(),
"discovered_at": manifest.changes[0].discovered_at.isoformat(),
}
},
)
finally:
await engine.dispose()
report = await IngesterApp(
config=config, db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest(manifest)
assert report.succeeded == 1
assert report.dead == 0
client.create_document_from_source.assert_awaited_once()
@pytest.mark.asyncio
async def test_run_batch_from_manifest_rejects_non_manifest_pending_work(
tmp_path, use_client
):
(tmp_path / "a.md").write_text("hello") (tmp_path / "a.md").write_text("hello")
config = _config(tmp_path) config = _config(tmp_path)
client = _mock_client() client = _mock_client()
@ -382,7 +428,7 @@ async def test_run_batch_from_manifest_rejects_pending_work(tmp_path, use_client
finally: finally:
await engine.dispose() await engine.dispose()
with pytest.raises(ValueError, match="pending work"): with pytest.raises(ValueError, match="non-manifest pending work"):
await IngesterApp( await IngesterApp(
config=config, db_path=tmp_path / "db.lancedb" config=config, db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest( ).run_batch_from_manifest(
@ -397,6 +443,46 @@ async def test_run_batch_from_manifest_rejects_pending_work(tmp_path, use_client
) )
@pytest.mark.asyncio
async def test_run_batch_from_manifest_rejects_different_manifest_pending_work(
tmp_path, use_client
):
(tmp_path / "a.md").write_text("hello")
config = _config(tmp_path)
client = _mock_client()
use_client(client)
manifest = _manifest(
BatchChange(
op=JobOp.UPSERT,
source_id="local",
uri=(tmp_path / "a.md").as_uri(),
discovered_at=datetime.now(UTC),
)
)
engine = await open_queue(config.ingester.queue)
try:
jobs = JobRepo(engine)
await jobs.enqueue(
"local",
(tmp_path / "a.md").as_uri(),
op=JobOp.UPSERT,
extra={
"_manifest": {
"version": 1,
"generated_at": "2026-01-01T00:00:00+00:00",
"discovered_at": manifest.changes[0].discovered_at.isoformat(),
}
},
)
finally:
await engine.dispose()
with pytest.raises(ValueError, match="non-manifest pending work"):
await IngesterApp(
config=config, db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest(manifest)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_batch_from_manifest_rejects_unrelated_pending_work( async def test_run_batch_from_manifest_rejects_unrelated_pending_work(
tmp_path, use_client tmp_path, use_client
@ -412,7 +498,7 @@ async def test_run_batch_from_manifest_rejects_unrelated_pending_work(
finally: finally:
await engine.dispose() await engine.dispose()
with pytest.raises(ValueError, match="queue has pending work"): with pytest.raises(ValueError, match="non-manifest pending work"):
await IngesterApp( await IngesterApp(
config=config, db_path=tmp_path / "db.lancedb" config=config, db_path=tmp_path / "db.lancedb"
).run_batch_from_manifest( ).run_batch_from_manifest(