259 lines
12 KiB
Python
259 lines
12 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 PdfRenderError, process_pdf, render_pdf_pages
|
|
|
|
|
|
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 image(self, job_id: str, page: int) -> tuple[str, bytes] | None:
|
|
row = self.connection.execute("SELECT status,document_sha256,pdf FROM jobs WHERE job_id=?", (job_id,)).fetchone()
|
|
if not row:
|
|
return None
|
|
if row["status"] != "succeeded" or row["pdf"] is None:
|
|
fail(409, "RESULT_NOT_READY", "OCR review image is not ready")
|
|
try:
|
|
return row["document_sha256"], render_pdf_pages(bytes(row["pdf"]), [page])[0].png
|
|
except PdfRenderError as error:
|
|
fail(422, "INVALID_PAGE", str(error))
|
|
|
|
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", "service": "ocr", "version": os.getenv("OCR_VERSION", "0.1.0")}
|
|
|
|
@application.get("/health/ready")
|
|
def ready(response: Response) -> dict[str, int | bool | str]:
|
|
queue.sweep_expired()
|
|
if not engine_ready:
|
|
response.status_code = 503
|
|
return {"ready": engine_ready, "queueDepth": queue.depth(), "queueCapacity": QUEUE_CAPACITY, "concurrency": 1,
|
|
"service": "ocr", "version": os.getenv("OCR_VERSION", "0.1.0")}
|
|
|
|
@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.get("/v1/jobs/{job_id}/pages/{page}/image", dependencies=[Depends(authorize)])
|
|
def get_review_image(job_id: str, page: int) -> Response:
|
|
stored = queue.image(job_id, page)
|
|
if stored is None:
|
|
fail(404, "JOB_NOT_FOUND", "OCR job does not exist")
|
|
document_sha256, png = stored
|
|
return Response(png, media_type="image/png", headers={
|
|
"X-Document-Sha256": document_sha256,
|
|
"X-Content-Sha256": hashlib.sha256(png).hexdigest(),
|
|
"X-Page-Number": str(page),
|
|
})
|
|
|
|
@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,
|
|
)
|