import asyncio import json import uuid from collections.abc import Mapping from datetime import UTC, datetime, timedelta import sqlalchemy as sa from sqlalchemy.dialects import postgresql from sqlalchemy.dialects import sqlite as sqlite_dialect from sqlalchemy.ext.asyncio import AsyncEngine from haiku.rag.ingester.queue.db import jobs, sync_state from haiku.rag.ingester.queue.models import ( Job, JobOp, JobStatus, SyncRow, SyncStateRow, ) def _utcnow_iso() -> str: return datetime.now(UTC).isoformat() def _parse_dt(value: str | None) -> datetime | None: return datetime.fromisoformat(value) if value else None def _insert(table: sa.Table, dialect: str): """Dialect-specific INSERT exposing on_conflict_* (and `.excluded`).""" if dialect == "postgresql": return postgresql.insert(table) return sqlite_dialect.insert(table) def _attempts_minus_one() -> sa.ColumnElement[int]: """attempts - 1, floored at 0. Renders identically on both dialects.""" return sa.case((jobs.c.attempts - 1 < 0, 0), else_=jobs.c.attempts - 1) def _row_to_job(row: Mapping) -> Job: extra_text = row["extra"] return Job( id=row["id"], source_id=row["source_id"], uri=row["uri"], op=JobOp(row["op"]), content_hash=row["content_hash"], revision=row["revision"], status=JobStatus(row["status"]), attempts=row["attempts"], max_attempts=row["max_attempts"], last_error=row["last_error"], extra=json.loads(extra_text) if extra_text else None, enqueued_at=datetime.fromisoformat(row["enqueued_at"]), scheduled_at=datetime.fromisoformat(row["scheduled_at"]), claimed_at=_parse_dt(row["claimed_at"]), claimed_by=row["claimed_by"], completed_at=_parse_dt(row["completed_at"]), ) def _row_to_sync_state(row: Mapping) -> SyncStateRow: return SyncStateRow( source_id=row["source_id"], uri=row["uri"], revision=row["revision"], content_hash=row["content_hash"], last_seen_at=datetime.fromisoformat(row["last_seen_at"]), last_ingested_at=_parse_dt(row["last_ingested_at"]), ) class JobRepo: def __init__(self, engine: AsyncEngine): self._engine = engine self._dialect = engine.dialect.name # Notified after a successful enqueue so workers in this process wake # immediately instead of polling on a fixed sleep interval. Workers in # other processes (a shared Postgres queue) fall back to polling. self.job_available = asyncio.Condition() async def enqueue( self, source_id: str, uri: str, op: JobOp = JobOp.UPSERT, *, revision: str | None = None, content_hash: str | None = None, max_attempts: int = 5, extra: dict | None = None, ) -> Job | None: """Enqueue an upsert/delete job. Returns the inserted Job, or None if a live (queued/claimed) job already exists for the same (source_id, uri). The partial unique index enforces atomicity.""" job_id = str(uuid.uuid4()) now = _utcnow_iso() extra_json = json.dumps(extra) if extra is not None else None stmt = ( _insert(jobs, self._dialect) .values( id=job_id, source_id=source_id, uri=uri, op=op.value, content_hash=content_hash, revision=revision, status="queued", attempts=0, max_attempts=max_attempts, last_error=None, extra=extra_json, enqueued_at=now, scheduled_at=now, ) .on_conflict_do_nothing() .returning(*jobs.c) ) async with self._engine.begin() as conn: row = (await conn.execute(stmt)).mappings().one_or_none() if row is not None: async with self.job_available: self.job_available.notify_all() return _row_to_job(row) if row else None async def claim_next(self, worker_id: str) -> Job | None: """Atomically claim the oldest queued job whose scheduled_at <= now. A single `UPDATE ... WHERE id = (SELECT ... LIMIT 1) RETURNING` keeps the claim atomic across connections: on Postgres the subquery adds `FOR UPDATE SKIP LOCKED`; on SQLite the whole statement runs under one write lock, so a racing connection re-evaluates the subquery against the committed state and finds the row already claimed.""" now = _utcnow_iso() candidate = ( sa.select(jobs.c.id) .where(jobs.c.status == "queued", jobs.c.scheduled_at <= now) .order_by(jobs.c.scheduled_at, jobs.c.id) .limit(1) .with_for_update(skip_locked=True) .scalar_subquery() ) claim = ( sa.update(jobs) .where(jobs.c.id == candidate) .values( status="claimed", claimed_at=now, claimed_by=worker_id, attempts=jobs.c.attempts + 1, ) .returning(*jobs.c) ) async with self._engine.begin() as conn: row = (await conn.execute(claim)).mappings().one_or_none() return _row_to_job(row) if row else None async def get_job(self, job_id: str) -> Job | None: async with self._engine.connect() as conn: row = ( (await conn.execute(sa.select(jobs).where(jobs.c.id == job_id))) .mappings() .one_or_none() ) return _row_to_job(row) if row else None async def mark_succeeded(self, job_id: str, claimed_by: str) -> bool: """Transition a still-claimed job to `succeeded`. Guarded on `status='claimed' AND claimed_by=?` so a reaper-resurrected job picked up by a different worker isn't clobbered by the original slow worker. Returns True when the row was updated.""" stmt = ( sa.update(jobs) .where( jobs.c.id == job_id, jobs.c.status == "claimed", jobs.c.claimed_by == claimed_by, ) .values(status="succeeded", completed_at=_utcnow_iso()) .returning(jobs.c.id) ) async with self._engine.begin() as conn: row = (await conn.execute(stmt)).first() return row is not None async def mark_dead(self, job_id: str, error: str, claimed_by: str) -> bool: """Transition a still-claimed job to `dead`. See `mark_succeeded` for the guard semantics.""" stmt = ( sa.update(jobs) .where( jobs.c.id == job_id, jobs.c.status == "claimed", jobs.c.claimed_by == claimed_by, ) .values(status="dead", completed_at=_utcnow_iso(), last_error=error) .returning(jobs.c.id) ) async with self._engine.begin() as conn: row = (await conn.execute(stmt)).first() return row is not None async def reschedule( self, job_id: str, delay_seconds: float, error: str, claimed_by: str ) -> bool: """Reset a still-claimed job back to `queued` with a future scheduled_at. Guarded on `status='claimed' AND claimed_by=?` so a slow worker can't clobber a re-claim that happened after the reaper reset its claim. Returns True when the row was updated.""" scheduled = (datetime.now(UTC) + timedelta(seconds=delay_seconds)).isoformat() stmt = ( sa.update(jobs) .where( jobs.c.id == job_id, jobs.c.status == "claimed", jobs.c.claimed_by == claimed_by, ) .values( status="queued", scheduled_at=scheduled, claimed_at=None, claimed_by=None, last_error=error, ) .returning(jobs.c.id) ) async with self._engine.begin() as conn: row = (await conn.execute(stmt)).first() return row is not None async def retry(self, job_id: str) -> Job: """Reset a `dead` or `queued` job: status='queued', attempts=0, error cleared, scheduled for immediate re-claim. Refuses `claimed` rows (would race with the worker still processing) and `succeeded` rows (re-ingest via UPSERT instead). Raises KeyError when the row is missing or in a non-retryable state.""" now = _utcnow_iso() stmt = ( sa.update(jobs) .where(jobs.c.id == job_id, jobs.c.status.in_(["dead", "queued"])) .values( status="queued", attempts=0, last_error=None, claimed_at=None, claimed_by=None, completed_at=None, scheduled_at=now, ) .returning(*jobs.c) ) async with self._engine.begin() as conn: row = (await conn.execute(stmt)).mappings().one_or_none() if not row: raise KeyError(f"Job {job_id!r} not found or not retryable") return _row_to_job(row) async def cancel(self, job_id: str) -> bool: """True iff a queued/claimed row was removed; terminal jobs aren't cancellable (succeeded/dead rows are kept for history).""" stmt = ( sa.delete(jobs) .where(jobs.c.id == job_id, jobs.c.status.in_(["queued", "claimed"])) .returning(jobs.c.id) ) async with self._engine.begin() as conn: row = (await conn.execute(stmt)).first() return row is not None async def list_jobs( self, *, status: JobStatus | None = None, source_id: str | None = None, uri: str | None = None, limit: int = 50, offset: int = 0, ) -> list[Job]: query = sa.select(jobs) if status is not None: query = query.where(jobs.c.status == status.value) if source_id is not None: query = query.where(jobs.c.source_id == source_id) if uri is not None: query = query.where(jobs.c.uri == uri) query = query.order_by(jobs.c.enqueued_at.desc()).limit(limit).offset(offset) async with self._engine.connect() as conn: rows = (await conn.execute(query)).mappings().all() return [_row_to_job(r) for r in rows] async def has_pending(self, source_id: str) -> bool: """True iff at least one queued/claimed job exists for the source. Cheap probe used by pollers to skip sweeps when there's already outstanding work — the queue's unique index would dedupe new enqueues anyway, so a sweep into a saturated queue is pure wasted listing work. """ query = ( sa.select(jobs.c.id) .where( jobs.c.source_id == source_id, jobs.c.status.in_(["queued", "claimed"]), ) .limit(1) ) async with self._engine.connect() as conn: row = (await conn.execute(query)).first() return row is not None async def counts_by_status(self) -> dict[str, int]: query = sa.select(jobs.c.status, sa.func.count().label("n")).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 counts_by_status_since(self, since: datetime) -> dict[str, int]: """status -> count of jobs that reached a terminal state at or after `since` (by completed_at). Only succeeded/dead set completed_at, so those are the only keys returned. Lets a one-shot batch report the work it finished, independent of terminal rows from earlier runs.""" query = ( sa.select(jobs.c.status, sa.func.count().label("n")) .where(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: """How many jobs reached `succeeded` in the last `seconds` seconds. Drives the dashboard's rolling-throughput chips.""" threshold = (datetime.now(UTC) - timedelta(seconds=seconds)).isoformat() query = sa.select(sa.func.count()).where( jobs.c.status == "succeeded", jobs.c.completed_at >= threshold ) async with self._engine.connect() as conn: count = (await conn.execute(query)).scalar() return int(count or 0) async def oldest_queued_age_seconds(self) -> float | None: """Age (in seconds) of the oldest job sitting in `queued` whose scheduled_at is in the past. Returns None when nothing is waiting. Tells operators whether work is backing up.""" now = datetime.now(UTC) query = sa.select(sa.func.min(jobs.c.scheduled_at)).where( jobs.c.status == "queued", jobs.c.scheduled_at <= now.isoformat() ) async with self._engine.connect() as conn: oldest = (await conn.execute(query)).scalar() if oldest is None: return None return (now - datetime.fromisoformat(oldest)).total_seconds() async def counts_by_source(self, *statuses: str) -> dict[str, int]: """source_id → count of jobs in any of the given statuses. Drives the dashboard's per-source DLQ and backlog summaries.""" if not statuses: return {} query = ( sa.select(jobs.c.source_id, sa.func.count().label("n")) .where(jobs.c.status.in_(statuses)) .group_by(jobs.c.source_id) ) async with self._engine.connect() as conn: rows = (await conn.execute(query)).all() return {source_id: n for source_id, n in rows} async def release_if_claimed(self, job_id: str, claimed_by: str) -> bool: """Reset a still-claimed job back to queued, immediately reclaimable. Guarded on `status='claimed' AND claimed_by=?` so the cancel-cleanup of a slow worker doesn't strip the claim of a different worker that re-claimed after a reaper reset. Decrements attempts to undo the increment from `claim_next`, since a cancellation isn't a failed attempt. Returns True if the row was released.""" stmt = ( sa.update(jobs) .where( jobs.c.id == job_id, jobs.c.status == "claimed", jobs.c.claimed_by == claimed_by, ) .values( status="queued", claimed_at=None, claimed_by=None, scheduled_at=_utcnow_iso(), attempts=_attempts_minus_one(), ) .returning(jobs.c.id) ) async with self._engine.begin() as conn: row = (await conn.execute(stmt)).first() return row is not None async def prune_dead(self, source_id: str, uri: str) -> int: """Delete dead jobs for the given (source_id, uri). Called after a successful DELETE to clear stale UPSERT failures for the same URI — the document is gone, so a "couldn't ingest this" entry is no longer actionable. Returns the number of rows removed.""" stmt = sa.delete(jobs).where( jobs.c.source_id == source_id, jobs.c.uri == uri, jobs.c.status == "dead", ) async with self._engine.begin() as conn: result = await conn.execute(stmt) return result.rowcount or 0 async def prune_terminal(self, max_age_seconds: int) -> int: """Delete terminal jobs (succeeded/dead) whose completed_at is older than max_age_seconds. General housekeeping so the table doesn't grow without bound. Returns the number of rows removed.""" threshold = (datetime.now(UTC) - timedelta(seconds=max_age_seconds)).isoformat() stmt = sa.delete(jobs).where( jobs.c.status.in_(["succeeded", "dead"]), jobs.c.completed_at < threshold, ) async with self._engine.begin() as conn: result = await conn.execute(stmt) return result.rowcount or 0 async def reap_stale(self, claim_timeout_seconds: int) -> int: """Reset claimed jobs whose claimed_at is older than the timeout back to `queued`. Decrements `attempts` to undo the increment from `claim_next` — a crashed worker isn't a consumed attempt.""" threshold = ( datetime.now(UTC) - timedelta(seconds=claim_timeout_seconds) ).isoformat() stmt = ( sa.update(jobs) .where(jobs.c.status == "claimed", jobs.c.claimed_at < threshold) .values( status="queued", claimed_at=None, claimed_by=None, attempts=_attempts_minus_one(), ) ) async with self._engine.begin() as conn: result = await conn.execute(stmt) return result.rowcount or 0 class SyncStateRepo: def __init__(self, engine: AsyncEngine): self._engine = engine self._dialect = engine.dialect.name async def get_revision_snapshot(self, source_id: str) -> dict[str, str]: """uri -> revision map for URIs that have a stored revision. Sources compare current revision against this map to decide UPSERT vs UNCHANGED. Rows without a revision (HTTP without ETag, or a worker that didn't complete) are excluded — they have no revision to compare against; the closing-loop DELETE diff uses list_known_uris instead.""" query = sa.select(sync_state.c.uri, sync_state.c.revision).where( sync_state.c.source_id == source_id, sync_state.c.revision.is_not(None), ) async with self._engine.connect() as conn: rows = (await conn.execute(query)).all() return {uri: revision for uri, revision in rows} async def list_known_uris(self, source_id: str) -> set[str]: """Every URI the source has ever produced. Used by the closing-loop diff in discover() so a URI previously seen but no longer visible (FS file deleted, HTTP URL removed from config) emits DELETE.""" query = sa.select(sync_state.c.uri).where(sync_state.c.source_id == source_id) async with self._engine.connect() as conn: rows = (await conn.execute(query)).all() return {uri for (uri,) in rows} async def get_row(self, source_id: str, uri: str) -> SyncStateRow | None: query = sa.select(sync_state).where( sync_state.c.source_id == source_id, sync_state.c.uri == uri ) async with self._engine.connect() as conn: row = (await conn.execute(query)).mappings().one_or_none() return _row_to_sync_state(row) if row else None def _upsert_stmt( self, source_id: str, uri: str, revision: str | None, content_hash: str | None, last_seen_at: str, last_ingested_at: str | None, ): ins = _insert(sync_state, self._dialect).values( source_id=source_id, uri=uri, revision=revision, content_hash=content_hash, last_seen_at=last_seen_at, last_ingested_at=last_ingested_at, ) excluded = ins.excluded return ins.on_conflict_do_update( index_elements=[sync_state.c.source_id, sync_state.c.uri], set_={ "revision": sa.func.coalesce(excluded.revision, sync_state.c.revision), "content_hash": sa.func.coalesce( excluded.content_hash, sync_state.c.content_hash ), "last_seen_at": excluded.last_seen_at, "last_ingested_at": sa.func.coalesce( excluded.last_ingested_at, sync_state.c.last_ingested_at ), }, ) async def upsert( self, source_id: str, uri: str, *, revision: str | None = None, content_hash: str | None = None, ingested: bool = False, ) -> None: """Insert-or-update the sync_state row. `ingested=True` stamps last_ingested_at; otherwise only last_seen_at is bumped. `revision=None` and `content_hash=None` leave any existing values untouched.""" now = _utcnow_iso() stmt = self._upsert_stmt( source_id, uri, revision, content_hash, now, now if ingested else None ) async with self._engine.begin() as conn: await conn.execute(stmt) async def batch_upsert(self, rows: list[SyncRow]) -> None: """Batch insert-or-update sync_state rows in a single transaction.""" if not rows: return now = _utcnow_iso() async with self._engine.begin() as conn: for source_id, uri, revision, content_hash, ingested in rows: stmt = self._upsert_stmt( source_id, uri, revision, content_hash, now, now if ingested else None, ) await conn.execute(stmt) async def delete(self, source_id: str, uri: str) -> None: stmt = sa.delete(sync_state).where( sync_state.c.source_id == source_id, sync_state.c.uri == uri ) async with self._engine.begin() as conn: await conn.execute(stmt)