179 lines
7.6 KiB
Python
179 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
|
|
|
|
from .models import load_runtime_engine
|
|
|
|
|
|
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
|
|
|
|
|
|
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,
|
|
)
|