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

175 lines
7.6 KiB
Python

import hashlib
import hmac
import json
import os
import sqlite3
import threading
import uuid
from datetime import datetime, timezone
from pathlib import Path
from typing import Annotated, Any
from fastapi import Depends, FastAPI, File, Form, Header, HTTPException, Response, UploadFile
MAX_UPLOAD_BYTES = 50 * 1024 * 1024
MAX_PAGES = 100
QUEUE_CAPACITY = 3
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):
self.connection = sqlite3.connect(str(path), check_same_thread=False)
self.connection.row_factory = sqlite3.Row
self.lock = threading.Lock()
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)"
)
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]) -> dict[str, Any]:
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)
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", datetime.now(timezone.utc).isoformat(),
)
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())
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
return {"jobId": job_id, "status": row["status"], "completedPages": 0,
"totalPages": len(json.loads(row["pages"])), "error": 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 create_app(
token: str,
db_path: str | Path = ":memory:",
max_upload_bytes: int = MAX_UPLOAD_BYTES,
engine_ready: bool = False,
) -> FastAPI:
application = FastAPI(title="Private OCR Service", docs_url=None, redoc_url=None)
queue = JobQueue(db_path)
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]:
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(
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")
return queue.submit(idempotency_key or "", payload)
@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.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
app = create_app(
os.getenv("OCR_INTERNAL_TOKEN", ""),
os.getenv("OCR_JOBS_DB", ":memory:"),
engine_ready=os.getenv("OCR_ENGINE_READY") == "1",
)