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

235 lines
11 KiB
Python

import hashlib
import hmac
import json
import os
import sqlite3
import threading
import uuid
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Annotated, Any, Callable
from fastapi import BackgroundTasks, Depends, FastAPI, File, Form, Header, HTTPException, Response, UploadFile
from .engine import OcrEngine
from .models import load_runtime_engine
from .render import process_pdf
MAX_UPLOAD_BYTES = 50 * 1024 * 1024
MAX_PAGES = 100
QUEUE_CAPACITY = 3
TRANSIENT_TTL = timedelta(hours=24)
ALLOWED_CONFIG = {
"languages": ["es", "en"],
"dpi": 200,
"engine": "paddleocr",
"engineVersion": "3.4.0",
"runtimeVersion": "3.2.2",
"configVersion": "ocr-v1",
"returnLayout": True,
}
REQUEST_FIELDS = {"documentSha256", "pages", *ALLOWED_CONFIG}
def fail(status: int, code: str, message: str, retryable: bool = False, headers: dict[str, str] | None = None) -> None:
raise HTTPException(status, {"code": code, "message": message, "retryable": retryable}, headers)
def page_hash(pages: list[int]) -> str:
value = json.dumps(pages, separators=(",", ":")).encode()
return hashlib.sha256(value).hexdigest()
class JobQueue:
def __init__(self, path: str | Path, now: Callable[[], datetime]):
self.connection = sqlite3.connect(str(path), check_same_thread=False)
self.connection.row_factory = sqlite3.Row
self.lock = threading.Lock()
self.now = now
self.connection.execute(
"CREATE TABLE IF NOT EXISTS jobs (job_id TEXT PRIMARY KEY, idempotency_key TEXT UNIQUE, "
"payload_hash TEXT, document_sha256 TEXT, pages TEXT, status TEXT, created_at TEXT, "
"pdf BLOB, result TEXT, error TEXT)"
)
columns = {row[1] for row in self.connection.execute("PRAGMA table_info(jobs)")}
for name, kind in (("pdf", "BLOB"), ("result", "TEXT"), ("error", "TEXT")):
if name not in columns:
self.connection.execute(f"ALTER TABLE jobs ADD COLUMN {name} {kind}")
self.connection.commit()
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 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": "queued",
"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:
fail(409, "IDEMPOTENCY_CONFLICT", "The idempotency key is already bound to another request")
return self.ack(row), True
if self.depth() >= QUEUE_CAPACITY:
fail(429, "QUEUE_FULL", "The OCR queue is full", True, {"Retry-After": "2"})
values = (
f"ocr_{uuid.uuid4()}", key, payload_hash, request["documentSha256"],
json.dumps(request["pages"]), "queued", self.now().isoformat(), pdf, None, None,
)
self.connection.execute("INSERT INTO jobs VALUES (?,?,?,?,?,?,?,?,?,?)", values)
self.connection.commit()
return self.ack(self.connection.execute("SELECT * FROM jobs WHERE job_id=?", (values[0],)).fetchone()), True
def execute(self, job_id: str, engine: OcrEngine) -> None:
with self.lock:
row = self.connection.execute("SELECT * FROM jobs WHERE job_id=?", (job_id,)).fetchone()
if row["status"] != "queued":
return
self.connection.execute("UPDATE jobs SET status='running' WHERE job_id=?", (job_id,))
self.connection.commit()
try:
result = process_pdf(job_id, row["document_sha256"], bytes(row["pdf"]), json.loads(row["pages"]), engine)
update = ("succeeded", json.dumps(result, sort_keys=True, separators=(",", ":")), None, job_id)
except Exception as error:
update = ("failed", None, json.dumps({"code": "OCR_PROCESSING_FAILED", "message": str(error)}), job_id)
with self.lock:
self.connection.execute("UPDATE jobs SET status=?, result=?, error=? WHERE job_id=?", update)
self.connection.commit()
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 delete(self, job_id: str) -> None:
with self.lock:
self.connection.execute("DELETE FROM jobs WHERE job_id=?", (job_id,))
self.connection.commit()
def sweep_expired(self) -> None:
with self.lock:
self.connection.execute("DELETE FROM jobs WHERE created_at < ?", ((self.now() - TRANSIENT_TTL).isoformat(),))
self.connection.commit()
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def create_app(
token: str,
db_path: str | Path = ":memory:",
max_upload_bytes: int = MAX_UPLOAD_BYTES,
engine_ready: bool = False,
engine: OcrEngine | None = None,
now: Callable[[], datetime] = utc_now,
) -> FastAPI:
application = FastAPI(title="Private OCR Service", docs_url=None, redoc_url=None)
queue = JobQueue(db_path, now)
def authorize(authorization: Annotated[str | None, Header()] = None) -> None:
scheme, _, supplied = (authorization or "").partition(" ")
if not token or scheme != "Bearer" or not hmac.compare_digest(supplied, token):
fail(401, "UNAUTHORIZED", "Valid bearer authorization is required", headers={"WWW-Authenticate": "Bearer"})
@application.get("/health/live")
def live() -> dict[str, str]:
return {"status": "ok"}
@application.get("/health/ready")
def ready(response: Response) -> dict[str, int | bool]:
queue.sweep_expired()
if not engine_ready:
response.status_code = 503
return {"ready": engine_ready, "queueDepth": queue.depth(), "queueCapacity": QUEUE_CAPACITY, "concurrency": 1}
@application.post("/v1/jobs", status_code=202, dependencies=[Depends(authorize)])
async def create_job(
background_tasks: BackgroundTasks, file: Annotated[UploadFile, File()], request: Annotated[str, Form()],
idempotency_key: Annotated[str | None, Header(alias="Idempotency-Key")] = None,
) -> dict[str, Any]:
content = await file.read(max_upload_bytes + 1)
if len(content) > max_upload_bytes:
fail(413, "UPLOAD_LIMIT_EXCEEDED", "PDF exceeds the upload limit")
try:
payload = json.loads(request)
except (json.JSONDecodeError, TypeError):
fail(400, "INVALID_REQUEST", "Request must be valid JSON")
if not isinstance(payload, dict) or set(payload) != REQUEST_FIELDS:
fail(400, "INVALID_REQUEST", "Request fields do not match the contract")
pages = payload["pages"]
if not isinstance(pages, list) or not pages or any(type(page) is not int or page < 1 for page in pages):
fail(422, "INVALID_PAGES", "Pages must be positive one-based integers")
if len(pages) > MAX_PAGES:
fail(413, "PAGE_LIMIT_EXCEEDED", "OCR jobs accept at most 100 pages")
if pages != sorted(set(pages)):
fail(422, "INVALID_PAGES", "Pages must be unique and ordered")
if any(payload[name] != value for name, value in ALLOWED_CONFIG.items()):
fail(400, "CONFIG_NOT_ALLOWED", "OCR configuration is not allowlisted")
digest = hashlib.sha256(content).hexdigest()
if payload["documentSha256"] != digest:
fail(422, "INTEGRITY_MISMATCH", "PDF bytes do not match documentSha256")
if file.content_type != "application/pdf" or not content.startswith(b"%PDF-") or b"/Encrypt" in content:
fail(422, "UNSUPPORTED_PDF", "PDF is corrupt, encrypted, or unsupported")
expected_key = f'{digest}:ocr-v1:{page_hash(pages)}'
if idempotency_key != expected_key and not queue.contains(idempotency_key or ""):
fail(400, "INVALID_IDEMPOTENCY_KEY", "Idempotency-Key does not match request identity")
ack, created = queue.submit(idempotency_key or "", payload, content)
if created and engine is not None:
background_tasks.add_task(queue.execute, ack["jobId"], engine)
return ack
@application.get("/v1/jobs/{job_id}", dependencies=[Depends(authorize)])
def get_job(job_id: str) -> dict[str, Any]:
status = queue.status(job_id)
if status is None:
fail(404, "JOB_NOT_FOUND", "OCR job does not exist")
return status
@application.get("/v1/jobs/{job_id}/result", dependencies=[Depends(authorize)])
def get_result(job_id: str) -> dict[str, Any]:
stored = queue.result(job_id)
if stored is None:
fail(404, "JOB_NOT_FOUND", "OCR job does not exist")
if stored[0] != "succeeded" or stored[1] is None:
fail(409, "RESULT_NOT_READY", "OCR job result is not ready")
return stored[1]
@application.delete("/v1/jobs/{job_id}", status_code=204, dependencies=[Depends(authorize)])
def delete_job(job_id: str) -> Response:
queue.delete(job_id)
return Response(status_code=204)
return application
runtime_engine = load_runtime_engine()
app = create_app(
os.getenv("OCR_INTERNAL_TOKEN", ""),
os.getenv("OCR_JOBS_DB", ":memory:"),
engine_ready=runtime_engine is not None,
engine=runtime_engine,
)