Spaces:
Runtime error
Runtime error
Download api_batch.py from sdudeja/agentic-extractor: direct link, hf CLI and curl.
- Browser
- Download file 9.66 kB
-
https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/api_batch.py
- Command line
-
hf download hf://spaces/sdudeja/agentic-extractor/api_batch.py
-
curl -L -o api_batch.py https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/api_batch.py
9.66 kB
| """ | |
| DocuLens — Batch processing endpoints. | |
| Accepts multiple files in a single request, processes them sequentially, | |
| and returns per-file results with a batch summary. | |
| """ | |
| import os | |
| import json | |
| import time | |
| import tempfile | |
| import logging | |
| from typing import Optional | |
| from uuid import uuid4 | |
| from fastapi import APIRouter, UploadFile, File, Form, HTTPException, Depends | |
| from PIL import Image | |
| from extraction import ( | |
| extract_invoice, extract_from_pdf, ExtractionResult, | |
| ) | |
| from db.supabase import save_document, create_run, complete_run, save_stage_result | |
| from middleware import check_rate_limit, validate_file_upload, sanitize_filename | |
| from api_webhooks import deliver_webhook | |
| logger = logging.getLogger(__name__) | |
| batch_router = APIRouter() | |
| HF_TOKEN = os.environ.get("HF_TOKEN", "") | |
| # Maximum files per batch (free tier) | |
| MAX_BATCH_SIZE = int(os.environ.get("MAX_BATCH_SIZE", "20")) | |
| # --------------------------------------------------------------------------- | |
| # Batch processing endpoint | |
| # --------------------------------------------------------------------------- | |
| async def batch_extract( | |
| files: list[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 multiple uploaded files. | |
| Returns a batch result with per-file status, extraction data, | |
| and a summary of passed/failed counts. | |
| """ | |
| if not files: | |
| raise HTTPException(status_code=400, detail="No files provided") | |
| # ── Tier-based batch limits ── | |
| effective_max = MAX_BATCH_SIZE | |
| if user_id: | |
| try: | |
| from api_usage import check_feature, check_quota, TIER_FEATURES, get_user_tier | |
| if not check_feature(user_id, "batch_upload"): | |
| raise HTTPException( | |
| status_code=403, | |
| detail="Batch upload is not available on the free plan. Upgrade to Starter or above.", | |
| ) | |
| tier = get_user_tier(user_id) | |
| tier_max = TIER_FEATURES.get(tier, {}).get("max_batch_files", MAX_BATCH_SIZE) | |
| effective_max = min(tier_max, MAX_BATCH_SIZE) if tier_max else MAX_BATCH_SIZE | |
| # Check quota for all files in the batch | |
| quota = check_quota(user_id, pages_requested=len(files)) | |
| if not quota["allowed"]: | |
| raise HTTPException(status_code=429, detail=quota["message"]) | |
| except HTTPException: | |
| raise | |
| except Exception as e: | |
| logger.warning("Batch tier check failed (non-blocking): %s", e) | |
| if len(files) > effective_max: | |
| raise HTTPException( | |
| status_code=400, | |
| detail=f"Batch size exceeds limit. Maximum {effective_max} files per batch.", | |
| ) | |
| batch_id = str(uuid4()) | |
| batch_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") | |
| token = "" if use_ocr_only else HF_TOKEN | |
| results = [] | |
| passed = 0 | |
| failed = 0 | |
| for i, file in enumerate(files): | |
| file_start = time.time() | |
| safe_filename = sanitize_filename(file.filename) | |
| file_result = { | |
| "index": i, | |
| "filename": safe_filename, | |
| "status": "pending", | |
| "run_id": None, | |
| "data": None, | |
| "error": None, | |
| "processing_time_ms": 0, | |
| } | |
| try: | |
| # Validate each file | |
| file_size = file.size or 0 | |
| validate_file_upload(file.filename, file_size, file.content_type) | |
| # Persistence | |
| doc_id = None | |
| run_id = None | |
| if should_persist: | |
| 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, | |
| ) | |
| file_result["run_id"] = run_id | |
| # Read file contents | |
| contents = await file.read() | |
| ext = os.path.splitext(safe_filename)[1].lower() | |
| # Extract | |
| if ext == ".pdf": | |
| page_results = extract_from_pdf( | |
| contents, hf_token=token, force_ocr=use_ocr_only, use_case=use_case, model_id=model, | |
| ) | |
| if not page_results: | |
| raise ValueError("Could not extract data from PDF") | |
| result = page_results[0] | |
| # Merge multi-page line items | |
| base_li = len(result.line_items) | |
| for pi, pr in enumerate(page_results[1:], start=1): | |
| for bb in pr.bounding_boxes: | |
| if bb.field.startswith("line_item_"): | |
| li_idx = int(bb.field.split("_")[-1]) | |
| bb.field = f"line_item_{base_li + li_idx}" | |
| bb.page = pi | |
| result.bounding_boxes.append(bb) | |
| result.line_items.extend(pr.line_items) | |
| base_li += len(pr.line_items) | |
| 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) | |
| result = extract_invoice( | |
| image, hf_token=token, force_ocr=use_ocr_only, use_case=use_case, model_id=model, | |
| ) | |
| finally: | |
| os.unlink(tmp_path) | |
| file_ms = int((time.time() - file_start) * 1000) | |
| result.processing_time_ms = file_ms | |
| is_failure = result.extraction_method.startswith("failed") | |
| if is_failure: | |
| file_result["status"] = "failed" | |
| file_result["error"] = _describe_failure(result.extraction_method) | |
| failed += 1 | |
| else: | |
| file_result["status"] = "passed" | |
| file_result["data"] = result.to_dict() | |
| # Remove large base64 images from batch response to keep payload small | |
| file_result["data"].pop("page_image", None) | |
| passed += 1 | |
| # Increment usage for successful extraction | |
| if user_id: | |
| try: | |
| from api_usage import increment_usage | |
| increment_usage(user_id, pages=1) | |
| except Exception: | |
| pass | |
| file_result["processing_time_ms"] = file_ms | |
| # Persist run result | |
| if run_id: | |
| save_stage_result( | |
| run_id=run_id, | |
| stage_type="extraction", | |
| stage_name="Data Extraction", | |
| status="failed" if is_failure else "passed", | |
| sort_order=1, | |
| output=result.to_dict() if not is_failure else None, | |
| error_message=file_result.get("error"), | |
| duration_ms=file_ms, | |
| ) | |
| complete_run( | |
| run_id=run_id, | |
| status="completed", | |
| overall_result="failed" if is_failure else "passed", | |
| extraction_data=result.to_dict() if not is_failure else {}, | |
| processing_time_ms=file_ms, | |
| error_message=file_result.get("error"), | |
| ) | |
| except HTTPException as he: | |
| file_result["status"] = "failed" | |
| file_result["error"] = he.detail | |
| file_result["processing_time_ms"] = int((time.time() - file_start) * 1000) | |
| failed += 1 | |
| except Exception as e: | |
| logger.warning(f"Batch file {i} ({safe_filename}) failed: {e}") | |
| file_result["status"] = "failed" | |
| file_result["error"] = str(e) | |
| file_result["processing_time_ms"] = int((time.time() - file_start) * 1000) | |
| failed += 1 | |
| results.append(file_result) | |
| batch_ms = int((time.time() - batch_start) * 1000) | |
| batch_response = { | |
| "batch_id": batch_id, | |
| "total": len(files), | |
| "passed": passed, | |
| "failed": failed, | |
| "processing_time_ms": batch_ms, | |
| "results": results, | |
| } | |
| # Fire webhook | |
| try: | |
| deliver_webhook("batch.completed", batch_response) | |
| except Exception: | |
| pass | |
| return batch_response | |
| # --------------------------------------------------------------------------- | |
| # Helpers | |
| # --------------------------------------------------------------------------- | |
| def _describe_failure(method: str) -> str: | |
| if "no_token" in method: | |
| return "HF_TOKEN not set — AI Vision unavailable." | |
| if "vlm" in method: | |
| return "AI Vision models failed. Inference API may be loading or rate-limited." | |
| return "Could not extract readable text. Image may be blurry or handwritten." | |