validops-east-3's picture
Deploy 72701a1d9ec33b1842d0a516ffdee753aaa754a9
efff1ff verified
Raw History Blame Contribute Delete
6.99 kB
"""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,
)