rag-service/ocr-service/app/jobs.py

435 lines
21 KiB
Python

import hashlib
import json
import logging
import os
import shutil
import sqlite3
import threading
import uuid
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any, Callable
from .engine import OcrEngine
from .render import process_pdf
QUEUE_CAPACITY = 3
LEASE_DURATION = timedelta(minutes=15)
TRANSIENT_TTL = timedelta(hours=24)
SWEEP_INTERVAL_SECONDS = 15 * 60
MAX_TRANSIENT_BYTES = 2 * 1024**3
MIN_FREE_BYTES = 2 * 1024**3
STORAGE_RATIO = 0.10
logger = logging.getLogger("ocr.jobs")
class QueueFullError(RuntimeError):
pass
class StoragePressureError(RuntimeError):
pass
class JobActiveError(RuntimeError):
pass
class JobQueue:
def __init__(self, db_path: str | Path, now: Callable[[], datetime], sweep_interval: float = SWEEP_INTERVAL_SECONDS) -> None:
self.now = now
self.sweep_interval = sweep_interval
self.lock = threading.Lock()
self.wake_worker = threading.Event()
self.wake_sweeper = threading.Event()
self.stop_worker = threading.Event()
self.worker: threading.Thread | None = None
self.sweeper: threading.Thread | None = None
self.worker_error: str | None = None
self.sweeper_error: str | None = None
self.last_sweep_at: str | None = None
self.db_path = Path(db_path)
self.artifacts_dir = self.db_path.parent / "artifacts"
if str(db_path) == ":memory:":
self.artifacts_dir = Path.cwd() / ".ocr-test-artifacts"
self.artifacts_dir.mkdir(parents=True, exist_ok=True)
self.connection = sqlite3.connect(str(db_path), check_same_thread=False)
self.connection.row_factory = sqlite3.Row
if self.connection.execute("PRAGMA auto_vacuum").fetchone()[0] != 2:
self.connection.execute("PRAGMA auto_vacuum=INCREMENTAL")
self.connection.execute("VACUUM")
self.connection.execute("PRAGMA journal_mode=WAL")
self.connection.execute(
"CREATE TABLE IF NOT EXISTS jobs ("
"job_id TEXT PRIMARY KEY, idempotency_key TEXT UNIQUE, payload_hash TEXT NOT NULL, "
"document_sha256 TEXT NOT NULL, pages TEXT NOT NULL, config_version TEXT NOT NULL DEFAULT 'ocr-v2', status TEXT NOT NULL, created_at TEXT NOT NULL, "
"input_path TEXT, result TEXT, error TEXT, attempts INTEGER NOT NULL DEFAULT 0, "
"recovery_attempts INTEGER NOT NULL DEFAULT 0, started_at TEXT, completed_at TEXT, lease_expires_at TEXT)"
)
columns = {row[1] for row in self.connection.execute("PRAGMA table_info(jobs)")}
for name, kind in (
("input_path", "TEXT"), ("result", "TEXT"), ("error", "TEXT"), ("config_version", "TEXT NOT NULL DEFAULT 'ocr-v2'"),
("attempts", "INTEGER NOT NULL DEFAULT 0"), ("recovery_attempts", "INTEGER NOT NULL DEFAULT 0"),
("started_at", "TEXT"), ("completed_at", "TEXT"), ("lease_expires_at", "TEXT"),
):
if name not in columns:
self.connection.execute(f"ALTER TABLE jobs ADD COLUMN {name} {kind}")
self.connection.commit()
self.recover_interrupted()
@staticmethod
def _iso(value: datetime) -> str:
return value.astimezone(timezone.utc).isoformat()
def _artifact_path(self, job_id: str) -> Path:
return self.artifacts_dir / job_id / "input.pdf"
def _write_file(self, target: Path, content: bytes) -> None:
target.parent.mkdir(mode=0o700, parents=True, exist_ok=True)
temporary = target.with_suffix(".tmp")
with temporary.open("wb") as output:
output.write(content)
output.flush()
os.fsync(output.fileno())
os.replace(temporary, target)
directory_fd = os.open(target.parent, os.O_DIRECTORY)
try:
os.fsync(directory_fd)
finally:
os.close(directory_fd)
os.chmod(target, 0o600)
def _write_artifact(self, job_id: str, content: bytes) -> str:
target = self._artifact_path(job_id)
self._write_file(target, content)
return str(target.relative_to(self.db_path.parent))
def _review_image_path(self, job_id: str, page: int) -> Path:
return self.artifacts_dir / job_id / "review-images" / f"page-{page:04d}.png"
def _persist_review_image(self, job_id: str, page: int, png: bytes) -> None:
if not self._has_storage_capacity(len(png)):
raise StoragePressureError("OCR_STORAGE_PRESSURE")
self._write_file(self._review_image_path(job_id, page), png)
def _remove_artifacts(self, job_id: str) -> None:
shutil.rmtree(self.artifacts_dir / job_id, ignore_errors=True)
def _remove_review_images(self, job_id: str) -> None:
shutil.rmtree(self.artifacts_dir / job_id / "review-images", ignore_errors=True)
def _owned_storage_bytes(self) -> int:
total = sum(path.stat().st_size for path in self.artifacts_dir.rglob("*") if path.is_file())
if str(self.db_path) != ":memory:":
for suffix in ("", "-shm", "-wal"):
path = Path(f"{self.db_path}{suffix}")
if path.exists():
total += path.stat().st_size
return total
def _has_storage_capacity(self, required_bytes: int = 0) -> bool:
usage = shutil.disk_usage(self.artifacts_dir)
budget = min(int(usage.total * STORAGE_RATIO), MAX_TRANSIENT_BYTES)
reserve = max(int(usage.total * STORAGE_RATIO), MIN_FREE_BYTES)
return self._owned_storage_bytes() + required_bytes <= budget and usage.free - required_bytes >= reserve
def storage_available(self) -> bool:
try:
if not self._has_storage_capacity():
return False
probe = self.artifacts_dir / ".write-probe"
with probe.open("wb") as output:
output.write(b"ok")
output.flush()
os.fsync(output.fileno())
probe.unlink()
return True
except OSError:
return False
def depth(self) -> int:
row = self.connection.execute("SELECT count(*) AS count FROM jobs WHERE status IN ('queued','running')").fetchone()
return int(row["count"])
def operational_state(self) -> dict[str, Any]:
rows = self.connection.execute("SELECT status,count(*) AS count FROM jobs GROUP BY status").fetchall()
states = {status: 0 for status in ("queued", "running", "succeeded", "failed")}
states.update({row["status"]: int(row["count"]) for row in rows})
worker_alive = self.worker is not None and self.worker.is_alive()
if self.worker_error:
worker_state = "failed"
elif not worker_alive:
worker_state = "stopped"
else:
worker_state = "processing" if states["running"] else "idle"
recovery_attempts = self.connection.execute("SELECT COALESCE(sum(recovery_attempts),0) FROM jobs").fetchone()[0]
return {
"queueDepth": states["queued"] + states["running"], "queueCapacity": QUEUE_CAPACITY,
"queueStates": states, "concurrency": 1, "workerState": worker_state,
"workerOperational": worker_alive and self.worker_error is None,
"sweeperOperational": self.sweeper is not None and self.sweeper.is_alive() and self.sweeper_error is None,
"recoveryAttempts": int(recovery_attempts), "lastSweepAt": self.last_sweep_at,
}
def contains(self, key: str) -> bool:
return self.connection.execute("SELECT 1 FROM jobs WHERE idempotency_key=?", (key,)).fetchone() is not None
@staticmethod
def _identity(row: sqlite3.Row) -> dict[str, Any]:
pages = json.loads(row["pages"])
config_version = row["config_version"]
canonical = json.dumps({
"configVersion": config_version,
"documentSha256": row["document_sha256"],
"idempotencyKey": row["idempotency_key"],
"requestedPages": pages,
}, sort_keys=True, separators=(",", ":")).encode()
return {"idempotencyKey": row["idempotency_key"], "documentSha256": row["document_sha256"],
"requestedPages": pages, "configVersion": config_version,
"requestIdentitySha256": hashlib.sha256(canonical).hexdigest()}
@staticmethod
def ack(row: sqlite3.Row) -> dict[str, Any]:
return {
"jobId": row["job_id"], "status": row["status"], **JobQueue._identity(row), "createdAt": row["created_at"],
}
def submit(self, key: str, request: dict[str, Any], pdf: bytes) -> tuple[dict[str, Any], bool]:
payload_hash = hashlib.sha256(json.dumps(request, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
with self.lock:
row = self.connection.execute("SELECT * FROM jobs WHERE idempotency_key=?", (key,)).fetchone()
if row:
if row["payload_hash"] != payload_hash:
raise ValueError("IDEMPOTENCY_CONFLICT")
return self.ack(row), False
if self.depth() >= QUEUE_CAPACITY:
raise QueueFullError("QUEUE_FULL")
if not self._has_storage_capacity(len(pdf)):
raise StoragePressureError("OCR_STORAGE_PRESSURE")
job_id = f"ocr_{uuid.uuid4()}"
try:
input_path = self._write_artifact(job_id, pdf)
self.connection.execute(
"INSERT INTO jobs (job_id,idempotency_key,payload_hash,document_sha256,pages,config_version,status,created_at,input_path) "
"VALUES (?,?,?,?,?,?,?,?,?)",
(job_id, key, payload_hash, request["documentSha256"], json.dumps(request["pages"]), request["configVersion"], "queued", self._iso(self.now()), input_path),
)
self.connection.commit()
except Exception:
self.connection.rollback()
self._remove_artifacts(job_id)
raise
row = self.connection.execute("SELECT * FROM jobs WHERE job_id=?", (job_id,)).fetchone()
logger.info("ocr_transition job_id=%s from=none to=queued category=accepted", job_id)
self.wake_worker.set()
return self.ack(row), True
def recover_interrupted(self) -> None:
with self.lock:
interrupted = self.connection.execute(
"SELECT job_id,recovery_attempts,started_at FROM jobs WHERE status='running'"
).fetchall()
for row in interrupted:
timed_out = row["started_at"] is not None and datetime.fromisoformat(row["started_at"]) <= self.now() - LEASE_DURATION
if timed_out or row["recovery_attempts"] >= 1:
code = "OCR_PROCESSING_TIMEOUT" if timed_out else "OCR_RECOVERY_EXHAUSTED"
self.connection.execute(
"UPDATE jobs SET status='failed', error=?, completed_at=?, lease_expires_at=NULL WHERE job_id=?",
(json.dumps({"code": code, "message": "OCR processing did not complete safely"}), self._iso(self.now()), row["job_id"]),
)
logger.warning("ocr_transition job_id=%s from=running to=failed category=%s", row["job_id"], code.lower())
else:
self.connection.execute(
"UPDATE jobs SET status='queued', recovery_attempts=recovery_attempts+1, lease_expires_at=NULL WHERE job_id=?",
(row["job_id"],),
)
logger.warning("ocr_transition job_id=%s from=running to=queued category=startup_recovery", row["job_id"])
self.connection.commit()
if interrupted:
self.wake_worker.set()
def start_worker(self, engine: OcrEngine | None) -> None:
self.sweep_expired()
if self.sweeper is None:
self.sweeper = threading.Thread(target=self._run_sweeper, name="ocr-sweeper", daemon=True)
self.sweeper.start()
if engine is not None and self.worker is None:
self.worker = threading.Thread(target=self._run_worker, args=(engine,), name="ocr-worker", daemon=True)
self.worker.start()
def stop(self) -> None:
self.stop_worker.set()
self.wake_worker.set()
self.wake_sweeper.set()
if self.worker is not None:
self.worker.join(timeout=5)
if self.sweeper is not None:
self.sweeper.join(timeout=5)
def _claim(self) -> sqlite3.Row | None:
with self.lock:
row = self.connection.execute("SELECT * FROM jobs WHERE status='queued' ORDER BY created_at LIMIT 1").fetchone()
if row is None:
return None
now = self._iso(self.now())
lease = self._iso(self.now() + LEASE_DURATION)
changed = self.connection.execute(
"UPDATE jobs SET status='running', attempts=attempts+1, started_at=COALESCE(started_at,?), lease_expires_at=? "
"WHERE job_id=? AND status='queued'", (now, lease, row["job_id"]),
).rowcount
self.connection.commit()
claimed = self.connection.execute("SELECT * FROM jobs WHERE job_id=?", (row["job_id"],)).fetchone() if changed else None
if claimed:
logger.info("ocr_transition job_id=%s from=queued to=running category=claimed", row["job_id"])
self.wake_sweeper.set()
return claimed
def _run_worker(self, engine: OcrEngine) -> None:
try:
self._worker_loop(engine)
except Exception:
self.worker_error = "OCR_WORKER_FAILED"
logger.error("ocr_worker category=worker_failed")
def _worker_loop(self, engine: OcrEngine) -> None:
while not self.stop_worker.is_set():
row = self._claim()
if row is None:
self.wake_worker.wait(timeout=1)
self.wake_worker.clear()
continue
try:
input_path = self.db_path.parent / row["input_path"]
result = process_pdf(
row["job_id"], row["document_sha256"], input_path.read_bytes(), json.loads(row["pages"]), engine,
persist_review_image=lambda page, png: self._persist_review_image(row["job_id"], page, png),
)
result.update(self._identity(row))
result["engine"]["configVersion"] = row["config_version"]
update = ("succeeded", json.dumps(result, sort_keys=True, separators=(",", ":")), None)
except StoragePressureError:
self._remove_review_images(row["job_id"])
update = ("failed", None, json.dumps({"code": "OCR_STORAGE_PRESSURE", "message": "OCR storage capacity is unavailable"}))
except Exception:
update = ("failed", None, json.dumps({"code": "OCR_PROCESSING_FAILED", "message": "OCR processing failed"}))
with self.lock:
changed = self.connection.execute(
"UPDATE jobs SET status=?, result=?, error=?, completed_at=?, lease_expires_at=NULL "
"WHERE job_id=? AND status='running'",
(*update, self._iso(self.now()), row["job_id"]),
).rowcount
self.connection.commit()
if changed:
category = "completed" if update[0] == "succeeded" else json.loads(update[2])["code"].lower()
logger.info("ocr_transition job_id=%s from=running to=%s category=%s", row["job_id"], update[0], category)
else:
stored = self.connection.execute("SELECT 1 FROM jobs WHERE job_id=?", (row["job_id"],)).fetchone()
(self._remove_review_images if stored else self._remove_artifacts)(row["job_id"])
def _next_sweep_delay(self) -> float:
with self.lock:
row = self.connection.execute(
"SELECT min(lease_expires_at),min(started_at) FROM jobs WHERE status='running'"
).fetchone()
deadlines = [datetime.fromisoformat(row[0])] if row[0] else []
if row[1]:
deadlines.append(datetime.fromisoformat(row[1]) + LEASE_DURATION)
if not deadlines:
return self.sweep_interval
return max(0.0, min(self.sweep_interval, min((deadline - self.now()).total_seconds() for deadline in deadlines)))
def _run_sweeper(self) -> None:
while not self.stop_worker.is_set():
try:
delay = self._next_sweep_delay()
except Exception:
delay = self.sweep_interval
if self.wake_sweeper.wait(delay):
self.wake_sweeper.clear()
if self.stop_worker.is_set():
return
continue
if self.stop_worker.is_set():
return
try:
self.sweep_expired()
self.sweeper_error = None
except Exception:
self.sweeper_error = "OCR_SWEEPER_FAILED"
logger.error("ocr_maintenance category=sweeper_failed")
def status(self, job_id: str) -> dict[str, Any] | None:
row = self.connection.execute("SELECT * FROM jobs WHERE job_id=?", (job_id,)).fetchone()
if not row:
return None
total_pages = len(json.loads(row["pages"]))
return {"jobId": job_id, "status": row["status"], **self._identity(row),
"completedPages": total_pages if row["status"] == "succeeded" else 0,
"totalPages": total_pages, "error": json.loads(row["error"]) if row["error"] else None}
def result(self, job_id: str) -> tuple[str, dict[str, Any] | None] | None:
row = self.connection.execute("SELECT status,result FROM jobs WHERE job_id=?", (job_id,)).fetchone()
return (row["status"], json.loads(row["result"]) if row["result"] else None) if row else None
def review_image(self, job_id: str, page: int) -> tuple[dict[str, Any], bytes] | None:
row = self.connection.execute("SELECT * FROM jobs WHERE job_id=?", (job_id,)).fetchone()
if not row:
return None
if row["status"] != "succeeded":
raise RuntimeError("RESULT_NOT_READY")
if page not in json.loads(row["pages"]):
raise ValueError("INVALID_PAGE")
try:
return self._identity(row), self._review_image_path(job_id, page).read_bytes()
except FileNotFoundError as error:
raise RuntimeError("RESULT_NOT_READY") from error
def delete(self, job_id: str) -> None:
self.sweep_expired()
with self.lock:
row = self.connection.execute("SELECT status FROM jobs WHERE job_id=?", (job_id,)).fetchone()
if row is not None and row["status"] == "running":
raise JobActiveError("JOB_ACTIVE")
self.connection.execute("DELETE FROM jobs WHERE job_id=?", (job_id,))
self.connection.commit()
self._remove_artifacts(job_id)
if row is not None:
logger.info("ocr_transition job_id=%s from=%s to=deleted category=confirmed_cleanup", job_id, row["status"])
def sweep_expired(self) -> None:
current = self.now()
cutoff = self._iso(current - TRANSIENT_TTL)
timeout_cutoff = self._iso(current - LEASE_DURATION)
with self.lock:
interrupted = self.connection.execute(
"SELECT job_id,recovery_attempts,started_at FROM jobs WHERE status='running' AND "
"(lease_expires_at IS NULL OR lease_expires_at <= ? OR started_at <= ?)",
(self._iso(current), timeout_cutoff),
).fetchall()
for row in interrupted:
timed_out = row["started_at"] is not None and row["started_at"] <= timeout_cutoff
if timed_out or row["recovery_attempts"] >= 1:
code = "OCR_PROCESSING_TIMEOUT" if timed_out else "OCR_RECOVERY_EXHAUSTED"
self.connection.execute(
"UPDATE jobs SET status='failed',error=?,completed_at=?,lease_expires_at=NULL WHERE job_id=?",
(json.dumps({"code": code, "message": "OCR processing did not complete safely"}), self._iso(current), row["job_id"]),
)
logger.warning("ocr_transition job_id=%s from=running to=failed category=%s", row["job_id"], code.lower())
else:
self.connection.execute(
"UPDATE jobs SET status='queued',recovery_attempts=recovery_attempts+1,lease_expires_at=NULL WHERE job_id=?",
(row["job_id"],),
)
logger.warning("ocr_transition job_id=%s from=running to=queued category=lease_recovery", row["job_id"])
rows = self.connection.execute("SELECT job_id FROM jobs WHERE status != 'running' AND created_at < ?", (cutoff,)).fetchall()
self.connection.executemany("DELETE FROM jobs WHERE job_id=?", [(row["job_id"],) for row in rows])
self.connection.commit()
self.connection.execute("PRAGMA wal_checkpoint(PASSIVE)")
self.connection.execute("PRAGMA incremental_vacuum")
self.connection.commit()
self.last_sweep_at = self._iso(current)
for row in rows:
self._remove_artifacts(row["job_id"])
logger.info("ocr_transition job_id=%s from=expired to=deleted category=ttl_cleanup", row["job_id"])