File size: 6,992 Bytes
efff1ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
"""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,
    )