haiku.rag/haiku_rag_slim/haiku/rag/ingester/workers/pool.py

175 lines
6.6 KiB
Python

import asyncio
import logging
import time
from typing import TYPE_CHECKING
from haiku.rag.ingester.exceptions import PermanentError, TransientError
from haiku.rag.ingester.queue.models import Job, JobOp
from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo
from haiku.rag.ingester.workers.pipeline import run_job
from haiku.rag.ingester.workers.retry import RetryPolicy, compute_backoff
if TYPE_CHECKING:
from haiku.rag.client import HaikuRAG
logger = logging.getLogger(__name__)
class WorkerPool:
"""Asyncio-based pool. N worker tasks share a bounded Semaphore, each
pulling jobs from the queue and running them through `run_job`. Reaper
task resets claims older than `claim_timeout_s` so a crashed worker
doesn't strand its job.
Lifecycle: build it, await start(), let it run, await stop().
"""
def __init__(
self,
*,
client: "HaikuRAG",
job_repo: JobRepo,
sync_repo: SyncStateRepo,
worker_count: int = 4,
max_concurrent: int = 4,
retry_policy: RetryPolicy | None = None,
poll_idle_interval_s: float = 1.0,
claim_timeout_s: int = 1800,
reaper_interval_s: int = 60,
):
self._client = client
self._jobs = job_repo
self._sync = sync_repo
self._worker_count = worker_count
self._semaphore = asyncio.Semaphore(max_concurrent)
self._retry = retry_policy or RetryPolicy()
self._poll_idle_s = poll_idle_interval_s
self._claim_timeout_s = claim_timeout_s
self._reaper_interval_s = reaper_interval_s
self._stop = asyncio.Event()
self._workers: list[asyncio.Task] = []
self._reaper: asyncio.Task | None = None
@property
def live_workers(self) -> int:
"""Worker tasks that are still running. Equal to worker_count under
normal operation; less when a worker has crashed."""
return sum(1 for t in self._workers if not t.done())
async def start(self) -> None:
if self._workers:
raise RuntimeError("WorkerPool already started")
self._stop.clear()
for i in range(self._worker_count):
self._workers.append(asyncio.create_task(self._worker_loop(f"worker-{i}")))
self._reaper = asyncio.create_task(self._reaper_loop())
async def stop(self) -> None:
self._stop.set()
tasks = list(self._workers)
if self._reaper is not None:
tasks.append(self._reaper)
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
self._workers.clear()
self._reaper = None
async def drain_once(self, worker_id: str = "drain") -> int:
"""Drain every currently-claimable job to completion on the calling
coroutine. Used by run-once and by tests; not by `start()`."""
processed = 0
while True:
job = await self._jobs.claim_next(worker_id)
if job is None:
return processed
await self._process(job)
processed += 1
async def _worker_loop(self, worker_id: str) -> None:
while not self._stop.is_set():
try:
job = await self._jobs.claim_next(worker_id)
except Exception: # pragma: no cover - defensive against DB hiccups
logger.exception("claim_next failed in %s", worker_id)
await self._sleep_or_stop(self._poll_idle_s)
continue
if job is None:
await self._sleep_or_stop(self._poll_idle_s)
continue
async with self._semaphore:
await self._process(job)
async def _reaper_loop(self) -> None:
while not self._stop.is_set():
await self._sleep_or_stop(self._reaper_interval_s)
if self._stop.is_set():
return
try:
reset = await self._jobs.reap_stale(self._claim_timeout_s)
if reset:
logger.info("Reaper reset %d stale claim(s)", reset)
except Exception: # pragma: no cover - defensive against DB hiccups
logger.exception("reaper failed")
async def _sleep_or_stop(self, seconds: float) -> None:
try:
await asyncio.wait_for(self._stop.wait(), timeout=seconds)
except TimeoutError:
pass
async def _process(self, job: Job) -> None:
started = time.monotonic()
logger.info("Processing %s %s (job %s)", job.op.value, job.uri, job.id)
try:
result = await run_job(self._client, job)
except asyncio.CancelledError:
# Graceful shutdown cancelled us mid-flight. Release the claim so
# the next process can pick the job up immediately instead of
# waiting on the reaper's claim_timeout_s.
await self._jobs.release_if_claimed(job.id)
logger.info("Job %s released back to queue on cancel", job.id)
raise
except PermanentError as e:
await self._jobs.mark_dead(job.id, str(e))
logger.info("Job %s dead (permanent): %s", job.id, e)
return
except TransientError as e:
if job.attempts >= job.max_attempts:
await self._jobs.mark_dead(job.id, str(e))
logger.info(
"Job %s dead (max attempts %d): %s", job.id, job.max_attempts, e
)
return
delay = compute_backoff(job.attempts, self._retry)
await self._jobs.reschedule(job.id, delay, str(e))
logger.info(
"Job %s rescheduled in %.1fs (attempt %d/%d): %s",
job.id,
delay,
job.attempts,
job.max_attempts,
e,
)
return
except Exception as e: # pragma: no cover - pipeline classifier net
# Defensive: pipeline classifier should have caught everything.
await self._jobs.mark_dead(job.id, f"unclassified: {e!r}")
logger.exception("Unclassified error in job %s", job.id)
return
await self._jobs.mark_succeeded(job.id)
if job.op is JobOp.DELETE:
await self._sync.delete(job.source_id, job.uri)
else:
await self._sync.upsert(
job.source_id,
job.uri,
revision=result.revision,
content_hash=result.content_hash,
ingested=True,
)
logger.info(
"Job %s succeeded in %.2fs: %s", job.id, time.monotonic() - started, job.uri
)