Spaces:
Runtime error
Runtime error
Download api.py from sdudeja/agentic-extractor: direct link, hf CLI and curl.
- Browser
- Download file 16.8 kB
-
https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/api.py
- Command line
-
hf download hf://spaces/sdudeja/agentic-extractor/api.py
-
curl -L -o api.py https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/api.py
16.8 kB
| """ | |
| DocuLens - 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 asyncio | |
| import tempfile | |
| import logging | |
| from fastapi import APIRouter, UploadFile, File, Form, HTTPException, Depends | |
| 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 | |
| # Security middleware | |
| from middleware import check_rate_limit, validate_file_upload, sanitize_filename | |
| 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 | |
| "https://ai-extractor.in", # Custom domain (apex) | |
| "https://www.ai-extractor.in", # Custom domain (www) | |
| ] | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| async def list_models(api_key: str = Depends(check_rate_limit)): | |
| """List available extraction and validation models.""" | |
| return { | |
| "extraction": EXTRACTION_MODELS, | |
| "validation": VALIDATION_MODELS, | |
| } | |
| 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(""), | |
| api_key: str = Depends(check_rate_limit), | |
| ): | |
| """ | |
| 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. | |
| """ | |
| # ββ Input validation ββ | |
| file_size = file.size or 0 | |
| validate_file_upload(file.filename, file_size, file.content_type) | |
| safe_filename = sanitize_filename(file.filename) | |
| # ββ Tier quota check (soft β skips if user_id not provided) ββ | |
| if user_id: | |
| try: | |
| from api_usage import check_quota, increment_usage | |
| quota = check_quota(user_id, pages_requested=1) | |
| if not quota["allowed"]: | |
| raise HTTPException(status_code=429, detail=quota["message"]) | |
| except HTTPException: | |
| raise | |
| except Exception as e: | |
| logger.warning("Quota check failed (non-blocking): %s", e) | |
| 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: | |
| if not file_size and hasattr(file, "size") and file.size: | |
| file_size = file.size | |
| ext = os.path.splitext(safe_filename)[1].lower() | |
| file_type = "pdf" if ext == ".pdf" else "image" | |
| doc_id = save_document( | |
| filename=safe_filename, | |
| 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(safe_filename)[1].lower() | |
| # ββ Stage 2: Extraction ββ | |
| extraction_start = time.time() | |
| if ext == ".pdf": | |
| results = await asyncio.to_thread(extract_from_pdf, contents, hf_token=token, force_ocr=use_ocr_only, use_case=use_case, model_id=model) | |
| 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 = await asyncio.to_thread(extract_invoice, image, hf_token=token, force_ocr=use_ocr_only, use_case=use_case, model_id=model) | |
| 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: | |
| fail_data = result.to_dict() | |
| if hasattr(result, '_page_images') and result._page_images: | |
| fail_data["page_images"] = result._page_images | |
| if run_id: | |
| complete_run( | |
| run_id=run_id, | |
| status="completed", | |
| overall_result="failed", | |
| extraction_data=fail_data, | |
| processing_time_ms=total_ms, | |
| error_message=_describe_failure(result.extraction_method), | |
| ) | |
| return { | |
| **fail_data, | |
| "run_id": run_id, | |
| "error": _describe_failure(result.extraction_method), | |
| } | |
| # ββ Success β mark run as passed (validation is a separate call) ββ | |
| resp = result.to_dict() | |
| # Include all page images for multi-page PDF preview | |
| if hasattr(result, '_page_images') and result._page_images: | |
| resp["page_images"] = result._page_images | |
| if run_id: | |
| complete_run( | |
| run_id=run_id, | |
| status="completed", | |
| overall_result="passed", | |
| extraction_data=resp, | |
| processing_time_ms=total_ms, | |
| ) | |
| resp["run_id"] = run_id | |
| # ββ Increment usage counter on successful extraction ββ | |
| if user_id: | |
| try: | |
| from api_usage import increment_usage | |
| increment_usage(user_id, pages=1) | |
| except Exception as e: | |
| logger.warning("Usage increment failed (non-blocking): %s", e) | |
| # Fire webhook for successful extraction | |
| try: | |
| from api_webhooks import deliver_webhook | |
| deliver_webhook("extraction.completed", resp) | |
| except Exception: | |
| pass | |
| 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)) | |
| 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"), | |
| api_key: str = Depends(check_rate_limit), | |
| ): | |
| """ | |
| Validate extraction results using dynamic math checks and VLM semantic | |
| verification. Checks are generated based on the document's use_case. | |
| """ | |
| # ββ Input validation ββ | |
| validate_file_upload(file.filename, file.size or 0, file.content_type) | |
| 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() | |
| safe_name = sanitize_filename(file.filename) | |
| ext = os.path.splitext(safe_name)[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 = await asyncio.to_thread( | |
| 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, | |
| ) | |
| validation_resp = { | |
| "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, | |
| } | |
| # Fire webhook for validation completion | |
| try: | |
| from api_webhooks import deliver_webhook | |
| deliver_webhook("validation.completed", validation_resp) | |
| except Exception: | |
| pass | |
| return validation_resp | |
| async def health(api_key: str = Depends(check_rate_limit)): | |
| 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." | |