"""FastAPI application. Run: uvicorn docxextract.api:app --host 0.0.0.0 --port 8000 """ from __future__ import annotations import asyncio import time from contextlib import asynccontextmanager from fastapi import FastAPI, File, Form, Request, UploadFile from fastapi.responses import JSONResponse, ORJSONResponse from pydantic import ValidationError as PydanticValidationError from . import parsing from .config import get_settings from .engine import get_engine from .logging import configure_logging, get_logger, new_request_id from .schema import (DocExtractError, ExtractRequest, ExtractResponse, ValidationError) from .service import ExtractionService log = get_logger(__name__) settings = get_settings() @asynccontextmanager async def lifespan(app: FastAPI): configure_logging(settings.log_level, settings.log_json) parsing.configure_cache(settings) log.info("starting up", extra={ "detail": f"model={settings.model_id} threads={settings.torch_threads} " f"batch={settings.batch_size}", }) # Load and warm the model before serving, so the first request does not # absorb a 1-135s cold start. try: await asyncio.to_thread(get_engine(settings).warmup) log.info("model ready") except Exception as exc: # noqa: BLE001 # Startup should not hard-fail: the health endpoint will report the # degraded state and requests will surface a 503. log.error("model warmup failed", exc_info=exc) yield log.info("shutting down") app = FastAPI( title="Document Data Extraction Service", version="1.0.0", description="Extract arbitrary user-defined keys from invoices, receipts " "and scanned documents.", default_response_class=ORJSONResponse, lifespan=lifespan, ) @app.middleware("http") async def request_context(request: Request, call_next): rid = new_request_id() started = time.perf_counter() try: response = await call_next(request) except DocExtractError as exc: return JSONResponse( status_code=exc.status_code, content={"error": type(exc).__name__, "message": exc.message, "detail": exc.detail, "request_id": rid}, ) except Exception as exc: # noqa: BLE001 log.error("unhandled error", exc_info=exc, extra={"request_id": rid}) return JSONResponse( status_code=500, content={"error": "InternalError", "message": "an unexpected error occurred", "request_id": rid}, ) response.headers["X-Request-ID"] = rid response.headers["X-Latency-Ms"] = f"{(time.perf_counter()-started)*1000:.1f}" return response def _service() -> ExtractionService: return ExtractionService(settings) @app.get("/health", tags=["ops"]) async def health() -> dict: """Liveness plus model readiness, for load-balancer checks.""" engine = get_engine(settings) # `warm` is the readiness signal for both backends. Reaching into # engine._pipe would break the ONNX backend, which has no such attribute. ready = bool(getattr(engine, "warm", False)) return { "status": "ok" if ready else "degraded", "model": settings.model_id, "backend": type(engine).__name__, "model_ready": ready, "load_seconds": round(getattr(engine, "load_seconds", 0.0), 2), "parse_cache": parsing.cache_stats(), } @app.get("/v1/config", tags=["ops"]) async def config() -> dict: return { "model": settings.model_id, "device": settings.device, "batch_size": settings.batch_size, "max_keys": settings.max_keys, "max_pages": settings.max_pages, "max_upload_mb": settings.max_upload_mb, "confidence_threshold": settings.confidence_threshold, } @app.post("/v1/extract", response_model=ExtractResponse, tags=["extract"]) async def extract( file: UploadFile = File(..., description="PDF, PNG, JPEG, TIFF or BMP"), keys: str = Form(..., description="Comma-separated extraction keys"), confidence_threshold: float | None = Form(default=None), max_pages: int | None = Form(default=None), ) -> ExtractResponse: """Extract one value per requested key from a document. Keys are comma-separated, e.g. `keys=INVOICE NO,Vendor Id,Invoice Date,Total Amount,GSTIN` """ key_list = [k.strip() for k in keys.split(",") if k.strip()] if not key_list: raise ValidationError("at least one non-empty key is required") if len(key_list) > settings.max_keys: raise ValidationError( f"too many keys: {len(key_list)} > {settings.max_keys}") # Read with a hard cap so a large upload cannot exhaust memory before the # size check ever runs. limit = settings.max_upload_mb * 1024 * 1024 data = await file.read(limit + 1) if len(data) > limit: from .schema import PayloadTooLargeError raise PayloadTooLargeError( f"document exceeds {settings.max_upload_mb} MB") service = _service() try: return await asyncio.wait_for( asyncio.to_thread( service.extract, data, key_list, confidence_threshold, max_pages), timeout=settings.request_timeout_s, ) except asyncio.TimeoutError: from .schema import ExtractionTimeoutError log.error("extraction timed out", extra={"n_keys": len(key_list)}) raise ExtractionTimeoutError( f"extraction exceeded {settings.request_timeout_s}s") from None @app.post("/v1/extract/json", response_model=ExtractResponse, tags=["extract"]) async def extract_json( payload: dict, request: Request, ) -> ExtractResponse: """JSON variant: base64 document plus a keys array. Useful for clients that already hold the file in memory and want to avoid multipart encoding overhead. """ import base64 try: req = ExtractRequest( keys=payload.get("keys", []), confidence_threshold=payload.get("confidence_threshold"), max_pages=payload.get("max_pages"), ) except PydanticValidationError as exc: raise ValidationError("invalid request", detail=exc.errors()) from exc b64 = payload.get("document_base64") if not isinstance(b64, str) or not b64: raise ValidationError("document_base64 is required") if "," in b64[:100] and b64.startswith("data:"): b64 = b64.split(",", 1)[1] try: data = base64.b64decode(b64, validate=True) except Exception as exc: # noqa: BLE001 raise ValidationError("document_base64 is not valid base64") from exc service = _service() return await asyncio.wait_for( asyncio.to_thread(service.extract, data, req.keys, req.confidence_threshold, req.max_pages), timeout=settings.request_timeout_s, )