""" 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."