haiku.rag/haiku_rag_slim/haiku/rag/ingester/workers/pool.py
2026-06-16 16:03:51 +03:00

343 lines
14 KiB
Python

import asyncio
import logging
import time
from typing import TYPE_CHECKING
from haiku.rag.config import CircuitBreakerConfig
from haiku.rag.ingester.exceptions import PermanentError, TransientError
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
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 collections.abc import Mapping
from haiku.rag.client import HaikuRAG
from haiku.rag.ingester.metadata import MetadataProvider
from haiku.rag.ingester.sources.base import Source
logger = logging.getLogger(__name__)
_WORKER_BREAKER_THRESHOLD = 5
_WORKER_BREAKER_COOLDOWN_S = 60.0
class WorkerPool:
"""`worker_count` async tasks each pull jobs from the queue and run them
through `run_job`. A 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,
retry_policy: RetryPolicy | None = None,
poll_idle_interval_s: float = 1.0,
claim_timeout_s: int = 1800,
reaper_interval_s: int = 60,
retention_s: int | None = None,
sources: "list[Source] | None" = None,
metadata_providers: "Mapping[str, MetadataProvider] | None" = None,
):
self._client = client
self._jobs = job_repo
self._sync = sync_repo
self._worker_count = worker_count
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._retention_s = retention_s
self._sources: list[Source] = list(sources) if sources else []
self._metadata_providers: dict[str, MetadataProvider] = (
dict(metadata_providers) if metadata_providers else {}
)
self._stop = asyncio.Event()
self._workers: list[asyncio.Task] = []
self._reaper: asyncio.Task | None = None
self._pending_releases: set[asyncio.Task] = set()
self._breakers: dict[str, CircuitBreaker] = {}
@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())
@property
def breaker_open(self) -> bool:
return any(b.is_open for b in self._breakers.values())
@property
def breaker_consecutive_failures(self) -> int:
return max((b.consecutive_failures for b in self._breakers.values()), default=0)
def _breaker_for(self, source_id: str) -> CircuitBreaker:
breaker = self._breakers.get(source_id)
if breaker is None:
breaker = CircuitBreaker(
CircuitBreakerConfig(
failure_threshold=_WORKER_BREAKER_THRESHOLD,
cooldown_s=_WORKER_BREAKER_COOLDOWN_S,
)
)
self._breakers[source_id] = breaker
return breaker
def _paused_source_ids(self) -> set[str]:
return {sid for sid, b in self._breakers.items() if b.is_open}
async def start(self) -> None:
if self._workers:
raise RuntimeError("WorkerPool already started")
self._stop.clear()
# Any rows in `claimed` at start time are owned by workers from a
# previous process that didn't get to release them (SIGKILL, OOM,
# host reboot). Reset them so fresh workers can claim immediately
# instead of waiting on the reaper's claim_timeout_s.
reset = await self._jobs.reap_stale(claim_timeout_seconds=0)
if reset:
logger.info("Boot-reaped %d stale claim(s) from previous process", reset)
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()
# Wake workers parked on job_available.wait() so they notice _stop
# immediately instead of sleeping out the full poll_idle interval.
async with self._jobs.job_available:
self._jobs.job_available.notify_all()
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_pending_releases(self, timeout: float = 2.0) -> int:
"""Wait for any in-flight cancel-cleanup release Tasks to finish.
Returns how many completed. Called by the lifecycle owner after
stop() (success or timeout) so orphans land their SQL update before
the queue connection closes. Tasks left running after `timeout`
will be reclaimed by the reaper on the next start instead."""
pending = list(self._pending_releases)
if not pending:
return 0
done, _ = await asyncio.wait(pending, timeout=timeout)
return len(done)
async def drain_once(self, worker_id: str = "drain") -> int:
"""Drain every currently-claimable job to completion on the calling
coroutine. Used 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():
job = await self._jobs.claim_next(
worker_id, exclude_source_ids=self._paused_source_ids()
)
if job is None:
try:
async with self._jobs.job_available:
await asyncio.wait_for(
self._jobs.job_available.wait(),
timeout=self._poll_idle_s,
)
except TimeoutError:
pass
continue
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
reset = await self._jobs.reap_stale(self._claim_timeout_s)
if reset:
logger.info("Reaper reset %d stale claim(s)", reset)
if self._retention_s is not None:
pruned = await self._jobs.prune_terminal(self._retention_s)
if pruned:
logger.info("Reaper pruned %d terminal job(s)", pruned)
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:
assert job.claimed_by is not None, "_process only runs on claimed jobs"
worker_id = job.claimed_by
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,
sources=self._sources,
metadata_providers=self._metadata_providers,
)
except asyncio.CancelledError:
# Graceful shutdown cancelled us mid-flight. Spawn the release
# as an independent Task tracked in _pending_releases — that
# way a second cancel (e.g. shutdown_grace_s elapses and
# wait_for cancels stop() again) can interrupt our await
# without interrupting the SQL update, and the lifecycle owner
# can drain the orphans before closing the queue connection.
release_task = asyncio.create_task(
self._jobs.release_if_claimed(job.id, worker_id)
)
self._pending_releases.add(release_task)
release_task.add_done_callback(self._pending_releases.discard)
try:
await asyncio.shield(release_task)
except asyncio.CancelledError:
logger.info(
"Job %s cancel-cleanup interrupted; orphan release Task "
"will be drained by the lifecycle owner",
job.id,
)
else:
logger.info("Job %s released back to queue on cancel", job.id)
raise
except PermanentError as e:
if not await self._jobs.mark_dead(job.id, str(e), worker_id):
logger.warning(
"Job %s lost claim before mark_dead (likely reaper race); "
"letting the re-claiming worker drive",
job.id,
)
return
logger.info("Job %s dead (permanent): %s", job.id, e)
# Record the failed revision so discovery treats the unchanged file as
# accounted-for and stops re-enqueuing it every sweep. sync_state.revision
# means "last accounted-for revision" — ingested OR permanently failed.
# Revision-less sources (no ETag) can't be suppressed this way.
if job.revision is not None:
try:
await self._sync.upsert(
job.source_id,
job.uri,
revision=job.revision,
content_hash=job.content_hash,
ingested=False,
)
except Exception:
# The job is already dead. A failed marker write only means
# the next sweep may re-enqueue this URI — not worth crashing
# the worker and shrinking the pool over.
logger.exception(
"Job %s dead but failure marker write failed for %s; "
"next sweep may re-enqueue",
job.id,
job.uri,
)
return
except TransientError as e:
breaker = self._breaker_for(job.source_id)
was_closed = not breaker.is_open
breaker.record_failure()
if was_closed and breaker.is_open:
logger.warning(
"Worker breaker opened for source %s after %d consecutive "
"transient failures; pausing its claims for %.0fs",
job.source_id,
_WORKER_BREAKER_THRESHOLD,
_WORKER_BREAKER_COOLDOWN_S,
)
if job.attempts >= job.max_attempts:
await self._jobs.mark_dead(job.id, str(e), worker_id)
logger.info(
"Job %s dead (max attempts %d): %s", job.id, job.max_attempts, e
)
return
delay = compute_backoff(job.attempts, self._retry)
if not await self._jobs.reschedule( # pragma: no cover - reaper race
job.id, delay, str(e), worker_id
):
logger.warning(
"Job %s lost claim before reschedule (likely reaper race); "
"letting the re-claiming worker drive retry instead",
job.id,
)
return
logger.info(
"Job %s rescheduled in %.1fs (attempt %d/%d): %s",
job.id,
delay,
job.attempts,
job.max_attempts,
e,
)
return
# Guard against the reaper race: if our claim was reset and another
# worker re-claimed the job, mark_succeeded is a no-op. Don't write
# sync_state in that case — the new worker will write it when it
# finishes.
if not await self._jobs.mark_succeeded(job.id, worker_id):
logger.warning(
"Job %s lost claim before mark_succeeded (likely reaper race); "
"skipping sync_state write",
job.id,
)
return
breaker = self._breaker_for(job.source_id)
was_open = breaker.is_open
breaker.record_success()
if was_open:
logger.info(
"Worker breaker closed for source %s after successful probe",
job.source_id,
)
try:
if job.op is JobOp.DELETE:
await self._sync.delete(job.source_id, job.uri)
# A successful DELETE resolves any earlier UPSERT failures for
# the same (source_id, uri): the document is gone, the original
# error is no longer actionable, the DLQ entry is visual noise.
pruned = await self._jobs.prune_dead(job.source_id, job.uri)
if pruned:
logger.info(
"Pruned %d dead job(s) for %s after successful DELETE",
pruned,
job.uri,
)
else:
await self._sync.upsert(
job.source_id,
job.uri,
revision=result.revision,
content_hash=result.content_hash,
ingested=True,
)
except Exception:
# The job is already marked succeeded — the document was ingested
# correctly. A sync_state write failure means the next sweep may
# redundantly re-ingest this URI, but that's better than crashing
# the worker and blocking the rest of the queue.
logger.exception(
"Job %s succeeded but sync_state write failed for %s; "
"next sweep may re-ingest",
job.id,
job.uri,
)
logger.info(
"Job %s succeeded in %.2fs: %s", job.id, time.monotonic() - started, job.uri
)