adudeja's picture
fixed scroll issue with pdf, line items made editable, bounding boxes accuracy
676a6ec
Raw History Blame Contribute Delete
14.5 kB
"""
DocuMint AI - FastAPI REST endpoints.
Mounted into Gradio's internal FastAPI app via APIRouter so that
ZeroGPU compatibility is preserved (demo.launch handles @spaces.GPU).
"""
import os
import json
import time
import tempfile
import logging
from fastapi import APIRouter, UploadFile, File, Form, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from PIL import Image
from extraction import (
extract_invoice, extract_from_pdf, ExtractionResult,
validate_with_vlm, run_dynamic_math_checks, image_to_base64,
)
# Persistence layer (gracefully no-ops when Supabase is not configured)
from db.supabase import save_document, create_run, complete_run, save_stage_result
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Router (mounted by app.py into Gradio's FastAPI app)
# ---------------------------------------------------------------------------
api_router = APIRouter()
HF_TOKEN = os.environ.get("HF_TOKEN", "")
# ---------------------------------------------------------------------------
# CORS helper β€” called by app.py after include_router
# ---------------------------------------------------------------------------
ALLOWED_ORIGINS = [
"http://localhost:5173", # Vite dev server
"http://localhost:3000",
"https://*.vercel.app", # Vercel preview deploys
]
FRONTEND_URL = os.environ.get("FRONTEND_URL", "")
if FRONTEND_URL:
ALLOWED_ORIGINS.append(FRONTEND_URL)
def add_cors_middleware(app):
"""Add CORS middleware to the given FastAPI/Starlette app."""
app.add_middleware(
CORSMiddleware,
allow_origins=ALLOWED_ORIGINS,
allow_origin_regex=r"https://.*\.vercel\.app",
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ---------------------------------------------------------------------------
# Available models
# ---------------------------------------------------------------------------
EXTRACTION_MODELS = [
{"id": "Qwen/Qwen2.5-VL-72B-Instruct", "name": "Qwen2.5-VL 72B", "tier": "high"},
{"id": "Qwen/Qwen2.5-VL-7B-Instruct", "name": "Qwen2.5-VL 7B", "tier": "mid"},
{"id": "Qwen/Qwen2.5-VL-3B-Instruct", "name": "Qwen2.5-VL 3B", "tier": "fast"},
{"id": "ocr-only", "name": "OCR + Patterns", "tier": "fast"},
]
VALIDATION_MODELS = [
{"id": "Qwen/Qwen2.5-VL-72B-Instruct", "name": "Qwen2.5-VL 72B", "tier": "high"},
{"id": "Qwen/Qwen2.5-VL-7B-Instruct", "name": "Qwen2.5-VL 7B", "tier": "mid"},
]
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@api_router.get("/api/v1/models")
async def list_models():
"""List available extraction and validation models."""
return {
"extraction": EXTRACTION_MODELS,
"validation": VALIDATION_MODELS,
}
@api_router.post("/api/v1/extract")
async def extract(
file: UploadFile = File(...),
model: str = Form("Qwen/Qwen2.5-VL-72B-Instruct"),
force_ocr: str = Form("false"),
persist: str = Form("true"),
use_case: str = Form("invoice"),
org_id: str = Form(""),
user_id: str = Form(""),
):
"""
Extract structured data from an uploaded invoice image or PDF.
Returns all fields, line items, and bounding box coordinates.
When persist=true and Supabase is configured, records the run.
"""
start = time.time()
use_ocr_only = force_ocr.lower() in ("true", "1", "yes") or model == "ocr-only"
should_persist = persist.lower() in ("true", "1", "yes")
# Determine the HF token to use β€” skip VLM when OCR-only
token = "" if use_ocr_only else HF_TOKEN
# ── Persistence: register document & open run ──
doc_id = None
run_id = None
if should_persist:
file_size = 0
if hasattr(file, "size") and file.size:
file_size = file.size
ext = os.path.splitext(file.filename or "")[1].lower()
file_type = "pdf" if ext == ".pdf" else "image"
doc_id = save_document(
filename=file.filename or "unknown",
file_type=file_type,
file_size=file_size,
org_id=org_id or None,
user_id=user_id or None,
)
run_id = create_run(
document_id=doc_id,
use_case=use_case,
org_id=org_id or None,
user_id=user_id or None,
)
# ── Stage 1: Upload (already complete at this point) ──
upload_ms = int((time.time() - start) * 1000)
if run_id:
save_stage_result(
run_id=run_id,
stage_type="upload",
stage_name="Document Upload",
status="passed",
sort_order=0,
duration_ms=upload_ms,
)
try:
contents = await file.read()
ext = os.path.splitext(file.filename or "")[1].lower()
# ── Stage 2: Extraction ──
extraction_start = time.time()
if ext == ".pdf":
results = extract_from_pdf(contents, hf_token=token, force_ocr=use_ocr_only, use_case=use_case)
if not results:
if run_id:
_record_extraction_failure(run_id, extraction_start, "Could not extract data from PDF")
raise HTTPException(status_code=422, detail="Could not extract data from PDF")
result = results[0]
# Tag page-0 bboxes
for bb in result.bounding_boxes:
bb.page = 0
# Collect all page images for multi-page preview
page_images = []
for r in results:
if r.page_image:
page_images.append(r.page_image)
# Merge line items and their bboxes from subsequent pages.
base_li_count = len(result.line_items)
for page_idx, r in enumerate(results[1:], start=1):
for bb in r.bounding_boxes:
if bb.field.startswith("line_item_"):
li_idx = int(bb.field.split("_")[-1])
bb.field = f"line_item_{base_li_count + li_idx}"
bb.page = page_idx
result.bounding_boxes.append(bb)
result.line_items.extend(r.line_items)
base_li_count += len(r.line_items)
result._page_images = page_images
else:
# Save to temp file and open as image
with tempfile.NamedTemporaryFile(suffix=ext or ".jpg", delete=False) as tmp:
tmp.write(contents)
tmp_path = tmp.name
try:
image = Image.open(tmp_path)
result = extract_invoice(image, hf_token=token, force_ocr=use_ocr_only, use_case=use_case)
finally:
os.unlink(tmp_path)
extraction_ms = int((time.time() - extraction_start) * 1000)
total_ms = int((time.time() - start) * 1000)
result.processing_time_ms = total_ms
# Check for failure modes
is_failure = result.extraction_method.startswith("failed")
if run_id:
extraction_status = "failed" if is_failure else "passed"
save_stage_result(
run_id=run_id,
stage_type="extraction",
stage_name="Data Extraction",
status=extraction_status,
sort_order=1,
output=result.to_dict() if not is_failure else None,
error_message=_describe_failure(result.extraction_method) if is_failure else None,
duration_ms=extraction_ms,
)
if is_failure:
if run_id:
complete_run(
run_id=run_id,
status="completed",
overall_result="failed",
extraction_data=result.to_dict(),
processing_time_ms=total_ms,
error_message=_describe_failure(result.extraction_method),
)
return {
**result.to_dict(),
"run_id": run_id,
"error": _describe_failure(result.extraction_method),
}
# ── Success β€” mark run as passed (validation is a separate call) ──
if run_id:
complete_run(
run_id=run_id,
status="completed",
overall_result="passed",
extraction_data=result.to_dict(),
processing_time_ms=total_ms,
)
resp = result.to_dict()
resp["run_id"] = run_id
# Include all page images for multi-page PDF preview
if hasattr(result, '_page_images') and result._page_images:
resp["page_images"] = result._page_images
return resp
except HTTPException:
raise
except Exception as e:
logger.exception("Extraction failed")
if run_id:
total_ms = int((time.time() - start) * 1000)
complete_run(
run_id=run_id,
status="completed",
overall_result="failed",
extraction_data={},
processing_time_ms=total_ms,
error_message=str(e),
)
raise HTTPException(status_code=500, detail=str(e))
@api_router.post("/api/v1/validate")
async def validate(
file: UploadFile = File(...),
extraction: str = Form(...),
model: str = Form("Qwen/Qwen2.5-VL-72B-Instruct"),
run_id: str = Form(""),
use_case: str = Form("_default"),
):
"""
Validate extraction results using dynamic math checks and VLM semantic
verification. Checks are generated based on the document's use_case.
"""
start = time.time()
try:
data = json.loads(extraction)
except json.JSONDecodeError:
raise HTTPException(status_code=400, detail="Invalid extraction JSON")
# ── Dynamic math/consistency checks (use_case-aware) ──
math_checks = run_dynamic_math_checks(data)
# ── VLM semantic validation (cross-check against source image) ──
semantic_checks = []
try:
contents = await file.read()
ext = os.path.splitext(file.filename or "")[1].lower()
if ext == ".pdf":
try:
import fitz
doc = fitz.open(stream=contents, filetype="pdf")
page = doc[0]
pix = page.get_pixmap(dpi=200)
image = Image.frombytes("RGB", [pix.width, pix.height], pix.samples)
doc.close()
except Exception:
image = None
else:
with tempfile.NamedTemporaryFile(suffix=ext or ".jpg", delete=False) as tmp:
tmp.write(contents)
tmp_path = tmp.name
try:
image = Image.open(tmp_path)
finally:
os.unlink(tmp_path)
if image and HF_TOKEN:
semantic_checks = validate_with_vlm(
image=image,
extraction_data=data,
hf_token=HF_TOKEN,
model_id=model,
use_case=use_case,
)
except Exception as e:
logger.warning(f"Semantic validation failed: {e}")
semantic_checks = [{
"name": "VLM validation",
"status": "skipped",
"message": f"Could not run semantic validation: {str(e)[:100]}",
}]
total_ms = int((time.time() - start) * 1000)
# Determine overall status from all check results
all_checks = math_checks + semantic_checks
statuses = [c["status"] for c in all_checks if c["status"] != "skipped"]
if "fail" in statuses:
overall = "invalid"
elif "warn" in statuses:
overall = "needs_review"
elif statuses:
overall = "valid"
else:
overall = "needs_review"
# ── Persist validation stage ──
if run_id:
result_map = {"valid": "passed", "invalid": "failed", "needs_review": "needs_review"}
save_stage_result(
run_id=run_id,
stage_type="validation",
stage_name="Validation",
status=result_map.get(overall, overall),
sort_order=2,
output={"math_checks": math_checks, "semantic_checks": semantic_checks, "overall": overall},
duration_ms=total_ms,
)
complete_run(
run_id=run_id,
status="completed",
overall_result=result_map.get(overall, overall),
extraction_data=data,
processing_time_ms=total_ms,
)
return {
"overall_status": overall,
"math_checks": math_checks,
"semantic_checks": semantic_checks,
"domain_checks": [],
"cross_model_confidence": 0.0,
"validation_model": model,
"processing_time_ms": total_ms,
}
@api_router.get("/api/v1/health")
async def health():
from db.supabase import is_enabled
return {
"status": "ok",
"hf_token_set": bool(HF_TOKEN),
"persistence_enabled": is_enabled(),
}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _record_extraction_failure(run_id: str, extraction_start: float, message: str):
"""Record a failed extraction stage and close the run."""
import time as _time
extraction_ms = int((_time.time() - extraction_start) * 1000)
save_stage_result(
run_id=run_id,
stage_type="extraction",
stage_name="Data Extraction",
status="failed",
sort_order=1,
error_message=message,
duration_ms=extraction_ms,
)
complete_run(
run_id=run_id,
status="completed",
overall_result="failed",
extraction_data={},
processing_time_ms=extraction_ms,
error_message=message,
)
def _describe_failure(method: str) -> str:
if "no_token" in method:
return "HF_TOKEN not set β€” AI Vision unavailable. Set the Space Secret to enable Qwen2.5-VL."
if "vlm" in method:
return "AI Vision models failed. Inference API may be loading or rate-limited. Try again shortly."
return "Could not extract readable text. Image may be blurry or handwritten."