419 lines
20 KiB
Python
419 lines
20 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, 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"),
|
|
("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 ack(row: sqlite3.Row) -> dict[str, Any]:
|
|
return {
|
|
"jobId": row["job_id"], "status": row["status"], "documentSha256": row["document_sha256"],
|
|
"requestedPages": json.loads(row["pages"]), "configVersion": "ocr-v1", "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,status,created_at,input_path) "
|
|
"VALUES (?,?,?,?,?,?,?,?)",
|
|
(job_id, key, payload_hash, request["documentSha256"], json.dumps(request["pages"]), "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),
|
|
)
|
|
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"], "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[str, bytes] | None:
|
|
row = self.connection.execute("SELECT document_sha256,pages,status 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 row["document_sha256"], 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"])
|