Spaces:
Runtime error
Runtime error
Download extraction.py from sdudeja/agentic-extractor: direct link, hf CLI and curl.
- Browser
- Download file 70.8 kB
-
https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/extraction.py
- Command line
-
hf download hf://spaces/sdudeja/agentic-extractor/extraction.py
-
curl -L -o extraction.py https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/extraction.py
70.8 kB
| """ | |
| DocuLens - Core Invoice/Document Extraction Engine | |
| Supports multiple extraction backends: | |
| 1. Hugging Face Inference API (Qwen2.5-VL) - Primary, highest accuracy | |
| 2. Tesseract OCR + regex patterns - Fallback, CPU-only, zero cost | |
| """ | |
| import os | |
| import re | |
| import json | |
| import base64 | |
| import logging | |
| from io import BytesIO | |
| from typing import Optional | |
| from dataclasses import dataclass, field, asdict | |
| from datetime import datetime | |
| from concurrent.futures import ThreadPoolExecutor, as_completed | |
| import requests | |
| from PIL import Image | |
| from huggingface_hub import InferenceClient | |
| from hf_client import create_inference_client, robust_chat_completion | |
| from gemini_client import extract_with_gemini, GEMINI_API_KEY | |
| logger = logging.getLogger(__name__) | |
| # --------------------------------------------------------------------------- | |
| # Data Models | |
| # --------------------------------------------------------------------------- | |
| class LineItem: | |
| description: str = "" | |
| quantity: Optional[float] = None | |
| unit_price: Optional[float] = None | |
| amount: Optional[float] = None | |
| class BoundingBox: | |
| field: str = "" | |
| x1: int = 0 | |
| y1: int = 0 | |
| x2: int = 0 | |
| y2: int = 0 | |
| page: int = 0 | |
| class ExtractionResult: | |
| # Legacy named fields (populated for invoice use case, empty for others) | |
| vendor_name: str = "" | |
| vendor_address: str = "" | |
| invoice_number: str = "" | |
| invoice_date: str = "" | |
| due_date: str = "" | |
| currency: str = "" | |
| subtotal: Optional[float] = None | |
| tax_amount: Optional[float] = None | |
| tax_rate: Optional[str] = None | |
| total_amount: Optional[float] = None | |
| line_items: list = field(default_factory=list) | |
| payment_terms: str = "" | |
| customer_name: str = "" | |
| customer_address: str = "" | |
| raw_text: str = "" | |
| extraction_method: str = "" | |
| confidence: float = 0.0 | |
| processing_time_ms: int = 0 | |
| bounding_boxes: list = field(default_factory=list) | |
| page_image: Optional[str] = None # data-URI of rendered page (for PDFs) | |
| # Dynamic fields β primary data carrier for all use cases | |
| document_type: str = "" | |
| extracted_fields: dict = field(default_factory=dict) | |
| def to_dict(self): | |
| d = asdict(self) | |
| d["line_items"] = [asdict(li) if isinstance(li, LineItem) else li for li in self.line_items] | |
| d["bounding_boxes"] = [asdict(bb) if isinstance(bb, BoundingBox) else bb for bb in self.bounding_boxes] | |
| return d | |
| # --------------------------------------------------------------------------- | |
| # Bounding Box Generation - matches extracted values to OCR word positions | |
| # --------------------------------------------------------------------------- | |
| def _get_ocr_word_boxes(image: Image.Image) -> list[dict]: | |
| """Run pytesseract image_to_data to get word-level bounding boxes.""" | |
| try: | |
| import pytesseract | |
| data = pytesseract.image_to_data(image, output_type=pytesseract.Output.DICT, config="--psm 3") | |
| words = [] | |
| for i in range(len(data["text"])): | |
| text = data["text"][i].strip() | |
| conf = int(data["conf"][i]) | |
| if text and conf > 10: | |
| words.append({ | |
| "text": text, | |
| "x": data["left"][i], | |
| "y": data["top"][i], | |
| "w": data["width"][i], | |
| "h": data["height"][i], | |
| }) | |
| return words | |
| except Exception as e: | |
| logger.warning(f"OCR word boxes failed: {e}") | |
| return [] | |
| def _date_search_variants(date_str: str) -> list[str]: | |
| """Generate alternative date representations for OCR matching. | |
| VLM returns YYYY-MM-DD but documents may show DD-Mon-YYYY, DD/MM/YYYY, etc. | |
| """ | |
| variants = [date_str] | |
| # Try parsing YYYY-MM-DD | |
| m = re.match(r"(\d{4})-(\d{2})-(\d{2})", date_str) | |
| if m: | |
| y, mo, d = m.group(1), m.group(2), m.group(3) | |
| d_int = str(int(d)) # strip leading zero | |
| months = ["Jan","Feb","Mar","Apr","May","Jun","Jul","Aug","Sep","Oct","Nov","Dec"] | |
| month_name = months[int(mo) - 1] if 1 <= int(mo) <= 12 else mo | |
| month_full = ["January","February","March","April","May","June", | |
| "July","August","September","October","November","December"] | |
| full_name = month_full[int(mo) - 1] if 1 <= int(mo) <= 12 else mo | |
| variants += [ | |
| f"{d}-{month_name}-{y}", # 15-Aug-2026 | |
| f"{d_int}-{month_name}-{y}", # 5-Aug-2026 | |
| f"{d}/{mo}/{y}", # 15/08/2026 | |
| f"{d_int}/{mo}/{y}", # 5/08/2026 | |
| f"{mo}/{d}/{y}", # 08/15/2026 | |
| f"{d}.{mo}.{y}", # 15.08.2026 | |
| f"{d}-{mo}-{y}", # 15-08-2026 | |
| f"{d_int} {month_name} {y}", # 5 Aug 2026 | |
| f"{d} {month_name} {y}", # 15 Aug 2026 | |
| f"{month_name} {d_int}, {y}", # Aug 5, 2026 | |
| f"{d_int} {full_name} {y}", # 5 August 2026 | |
| f"{d} {full_name} {y}", # 15 August 2026 | |
| f"{full_name} {d_int}, {y}", # August 5, 2026 | |
| ] | |
| return list(dict.fromkeys(variants)) | |
| def _currency_search_variants(currency: str) -> list[str]: | |
| """Generate alternative currency representations for OCR matching.""" | |
| symbol_map = { | |
| "INR": ["βΉ", "Rs", "Rs.", "INR"], | |
| "USD": ["$", "USD", "US$"], | |
| "EUR": ["β¬", "EUR"], | |
| "GBP": ["Β£", "GBP"], | |
| "TRY": ["βΊ", "TL", "TRY"], | |
| "JPY": ["Β₯", "JPY"], | |
| } | |
| return symbol_map.get(currency.upper(), [currency]) | |
| def _word_matches_token(word_text: str, token: str) -> bool: | |
| """Check if an OCR word matches a search token. | |
| Uses full-word matching for alphabetic tokens (avoiding "tax" matching | |
| "taxation") and substring matching for numeric tokens (since OCR often | |
| merges numbers with punctuation like "$1,234.56"). | |
| """ | |
| wt = word_text.lower() | |
| # Numeric tokens: allow substring (OCR may merge currency symbols, commas) | |
| # Also strip common OCR artifacts before checking | |
| if any(c.isdigit() for c in token): | |
| # Direct substring | |
| if token in wt: | |
| return True | |
| # Strip non-digit non-dot non-comma chars and retry | |
| wt_digits = re.sub(r"[^0-9.,]", "", wt) | |
| tk_digits = re.sub(r"[^0-9.,]", "", token) | |
| return tk_digits != "" and tk_digits in wt_digits | |
| # Short tokens (<=2 chars): exact match only to avoid noise | |
| if len(token) <= 2: | |
| return wt == token | |
| # Alphabetic tokens: check word boundaries β strip common punctuation | |
| # from OCR word and compare | |
| wt_clean = re.sub(r"[^a-z0-9]", "", wt) | |
| tk_clean = re.sub(r"[^a-z0-9]", "", token) | |
| if not tk_clean: | |
| return False | |
| return wt_clean == tk_clean or wt_clean.startswith(tk_clean) or wt_clean.endswith(tk_clean) | |
| def _find_all_text_bboxes(words: list[dict], search_text: str) -> list[BoundingBox]: | |
| """Find ALL bounding boxes for a text value by matching against OCR words. | |
| Returns every match location so callers can disambiguate using context | |
| (e.g. proximity to a field label). | |
| """ | |
| if not search_text or not words: | |
| return [] | |
| search_lower = search_text.lower().strip() | |
| search_tokens = search_lower.split() | |
| if not search_tokens: | |
| return [] | |
| results = [] | |
| if len(search_tokens) == 1: | |
| token = search_tokens[0] | |
| for w in words: | |
| if _word_matches_token(w["text"], token): | |
| results.append(BoundingBox( | |
| x1=w["x"], y1=w["y"], | |
| x2=w["x"] + w["w"], y2=w["y"] + w["h"], | |
| )) | |
| # Fallback: try joining the search text without spaces (OCR may merge) | |
| if not results: | |
| merged = search_lower.replace(" ", "") | |
| if len(merged) >= 3: | |
| for w in words: | |
| wt_clean = re.sub(r"[^a-z0-9]", "", w["text"].lower()) | |
| if wt_clean == merged: | |
| results.append(BoundingBox( | |
| x1=w["x"], y1=w["y"], | |
| x2=w["x"] + w["w"], y2=w["y"] + w["h"], | |
| )) | |
| else: | |
| # Multi-word: try anchoring on EACH token position, not just the first. | |
| # OCR may miss or garble the first word but read later words fine. | |
| for anchor_idx in range(len(search_tokens)): | |
| anchor_token = search_tokens[anchor_idx] | |
| for i, w in enumerate(words): | |
| if not _word_matches_token(w["text"], anchor_token): | |
| continue | |
| # Expand backward to find earlier tokens | |
| x1, y1 = w["x"], w["y"] | |
| x2, y2 = w["x"] + w["w"], w["y"] + w["h"] | |
| matched = 1 | |
| # Look backward for tokens before the anchor | |
| if anchor_idx > 0: | |
| back_j = anchor_idx - 1 | |
| for bk in range(1, min(anchor_idx * 2 + 1, i + 1)): | |
| if back_j < 0: | |
| break | |
| prev_w = words[i - bk] | |
| if abs(prev_w["y"] - w["y"]) > 80: | |
| break | |
| if _word_matches_token(prev_w["text"], search_tokens[back_j]): | |
| x1 = min(x1, prev_w["x"]) | |
| y1 = min(y1, prev_w["y"]) | |
| x2 = max(x2, prev_w["x"] + prev_w["w"]) | |
| y2 = max(y2, prev_w["y"] + prev_w["h"]) | |
| matched += 1 | |
| back_j -= 1 | |
| # Look forward for tokens after the anchor | |
| search_j = anchor_idx + 1 | |
| scan_limit = min((len(search_tokens) - anchor_idx) * 2, len(words) - i) | |
| for k in range(1, scan_limit): | |
| if search_j >= len(search_tokens): | |
| break | |
| next_w = words[i + k] | |
| if abs(next_w["y"] - w["y"]) > 80: | |
| break | |
| if _word_matches_token(next_w["text"], search_tokens[search_j]): | |
| x2 = max(x2, next_w["x"] + next_w["w"]) | |
| y2 = max(y2, next_w["y"] + next_w["h"]) | |
| x1 = min(x1, next_w["x"]) | |
| y1 = min(y1, next_w["y"]) | |
| matched += 1 | |
| search_j += 1 | |
| # Require at least half the tokens to match | |
| threshold = max(1, (len(search_tokens) + 1) // 2) | |
| if matched >= threshold: | |
| results.append(BoundingBox(x1=x1, y1=y1, x2=x2, y2=y2)) | |
| # If we found matches anchoring on this token, don't try other anchors | |
| # (avoids duplicate/overlapping results) | |
| if results: | |
| break | |
| return results | |
| def _find_text_bbox(words: list[dict], search_text: str) -> Optional[BoundingBox]: | |
| """Find bounding box for a text value (first match). Legacy convenience wrapper.""" | |
| matches = _find_all_text_bboxes(words, search_text) | |
| return matches[0] if matches else None | |
| def _bbox_center(bb: BoundingBox) -> tuple[float, float]: | |
| """Return (cx, cy) center of a bounding box.""" | |
| return ((bb.x1 + bb.x2) / 2, (bb.y1 + bb.y2) / 2) | |
| def _bbox_distance(a: BoundingBox, b: BoundingBox) -> float: | |
| """Euclidean distance between bounding box centers.""" | |
| ax, ay = _bbox_center(a) | |
| bx, by = _bbox_center(b) | |
| return ((ax - bx) ** 2 + (ay - by) ** 2) ** 0.5 | |
| def _bbox_overlaps(a: BoundingBox, b: BoundingBox, margin: int = 5) -> bool: | |
| """Check if two bounding boxes overlap (with optional margin).""" | |
| return not ( | |
| a.x2 + margin < b.x1 or b.x2 + margin < a.x1 or | |
| a.y2 + margin < b.y1 or b.y2 + margin < a.y1 | |
| ) | |
| def _get_surrounding_text( | |
| candidate: BoundingBox, | |
| words: list[dict], | |
| radius: int = 5, | |
| ) -> list[str]: | |
| """Return the OCR word texts surrounding a candidate bbox. | |
| Finds the OCR word(s) that overlap the candidate, then collects up to | |
| *radius* words before and after in reading order. Returns lowercased | |
| word texts β callers use these to check whether a field's label tokens | |
| appear near the candidate value. | |
| """ | |
| # Find word indices that overlap the candidate | |
| overlap_indices: list[int] = [] | |
| for idx, w in enumerate(words): | |
| w_bb = BoundingBox(x1=w["x"], y1=w["y"], | |
| x2=w["x"] + w["w"], y2=w["y"] + w["h"]) | |
| if _bbox_overlaps(candidate, w_bb, margin=3): | |
| overlap_indices.append(idx) | |
| if not overlap_indices: | |
| # Fallback: find the nearest word by center distance | |
| cx, cy = _bbox_center(candidate) | |
| best_idx, best_dist = 0, float("inf") | |
| for idx, w in enumerate(words): | |
| wx = w["x"] + w["w"] / 2 | |
| wy = w["y"] + w["h"] / 2 | |
| d = ((cx - wx) ** 2 + (cy - wy) ** 2) ** 0.5 | |
| if d < best_dist: | |
| best_dist = d | |
| best_idx = idx | |
| overlap_indices = [best_idx] | |
| lo = max(0, min(overlap_indices) - radius) | |
| hi = min(len(words), max(overlap_indices) + 1 + radius) | |
| # Exclude the candidate's own words | |
| surrounding = [] | |
| for idx in range(lo, hi): | |
| if idx not in overlap_indices: | |
| surrounding.append(words[idx]["text"].lower()) | |
| return surrounding | |
| def _surrounding_text_score( | |
| candidate: BoundingBox, | |
| words: list[dict], | |
| label_tokens: list[str], | |
| radius: int = 5, | |
| ) -> float: | |
| """Score how well the surrounding OCR words match a field's label tokens. | |
| Returns a value between 0.0 (no label tokens found nearby) and 1.0 | |
| (all label tokens present in the surrounding words). This helps | |
| disambiguate when the same value (e.g. "0" or "2790") appears in | |
| multiple places on the page. | |
| """ | |
| if not label_tokens: | |
| return 0.0 | |
| surrounding = _get_surrounding_text(candidate, words, radius) | |
| if not surrounding: | |
| return 0.0 | |
| # Check how many label tokens appear in the surrounding text | |
| joined_surrounding = " ".join(surrounding) | |
| matched = 0 | |
| for lt in label_tokens: | |
| # Allow partial matches for long tokens (e.g. "statement" matching "statements") | |
| if any(lt in sw or sw in lt for sw in surrounding if len(sw) > 2): | |
| matched += 1 | |
| elif lt in joined_surrounding: | |
| matched += 1 | |
| return matched / len(label_tokens) | |
| def _field_name_to_label_tokens(field_name: str) -> list[str]: | |
| """Derive likely label keywords from a snake_case field name. | |
| e.g. 'invoice_date' β ['invoice', 'date'] | |
| 'total_amount' β ['total', 'amount'] | |
| 'pan' β ['pan'] | |
| """ | |
| return [t.lower() for t in field_name.split("_") if len(t) > 1] | |
| def _find_label_bboxes(words: list[dict], field_name: str) -> list[BoundingBox]: | |
| """Find ALL positions where a field's label appears in the document OCR. | |
| Returns all matches so the caller can pick the one closest to the value. | |
| Searches for the label keywords derived from the field name. Looks for | |
| multi-word labels first (e.g. "Invoice Date"), then single keywords. | |
| """ | |
| tokens = _field_name_to_label_tokens(field_name) | |
| if not tokens: | |
| return [] | |
| # Try the full label as a phrase first | |
| full_label = " ".join(tokens) | |
| matches = _find_all_text_bboxes(words, full_label) | |
| if matches: | |
| return matches | |
| # Fall back to the most distinctive single token (longest word) | |
| tokens_by_len = sorted(tokens, key=len, reverse=True) | |
| for token in tokens_by_len: | |
| # Skip very common short words that appear everywhere | |
| if token in ("of", "no", "in", "to", "at", "by", "id", "or"): | |
| continue | |
| matches = _find_all_text_bboxes(words, token) | |
| if matches: | |
| return matches | |
| return [] | |
| def _find_label_bbox(words: list[dict], field_name: str) -> Optional[BoundingBox]: | |
| """Find where a field's label appears β convenience wrapper returning first match.""" | |
| matches = _find_label_bboxes(words, field_name) | |
| return matches[0] if matches else None | |
| def _pick_best_match( | |
| candidates: list[BoundingBox], | |
| label_bbox: Optional[BoundingBox], | |
| claimed: list[BoundingBox], | |
| label_bboxes: Optional[list[BoundingBox]] = None, | |
| words: Optional[list[dict]] = None, | |
| field_name: Optional[str] = None, | |
| ) -> Optional[BoundingBox]: | |
| """Pick the best bbox from candidates, preferring proximity to label and avoiding claimed regions. | |
| Scoring: lower is better. | |
| - Distance to label (if known) is the primary signal. | |
| - **Surrounding-text context** β when ``words`` and ``field_name`` are | |
| provided, candidates whose neighbouring OCR words contain the | |
| field's label tokens receive a significant bonus. This is the | |
| strongest disambiguation signal for repeated short values. | |
| - Values to the right of or below their label get a bonus (documents | |
| typically place values right-of or below labels). | |
| - Values that overlap the label itself are penalised (that's the label, | |
| not the value). | |
| - Overlapping a claimed region adds a large penalty. | |
| - When multiple label positions are given (label_bboxes), the scoring | |
| picks the best (label, value) pair. | |
| """ | |
| if not candidates: | |
| return None | |
| CLAIMED_PENALTY = 100_000 | |
| LABEL_OVERLAP_PENALTY = 50_000 | |
| WRONG_DIRECTION_PENALTY = 200 # values above or far-left of label | |
| CONTEXT_BONUS_MAX = 500 # max bonus for surrounding-text match | |
| all_labels = label_bboxes or ([label_bbox] if label_bbox else []) | |
| # Pre-compute label tokens for surrounding-text scoring | |
| label_tokens = _field_name_to_label_tokens(field_name) if field_name else [] | |
| # Filter out very common/short tokens that add noise | |
| label_tokens = [t for t in label_tokens if t not in ("of", "no", "in", "to", "at", "by", "id", "or", "is")] | |
| # Determine if the value is ambiguous (short/numeric β more likely repeated) | |
| is_ambiguous = len(candidates) > 1 | |
| def score_with_label(bb: BoundingBox, lbl: BoundingBox) -> float: | |
| s = _bbox_distance(bb, lbl) | |
| # Penalise if the candidate overlaps the label itself (likely IS the label) | |
| if _bbox_overlaps(bb, lbl, margin=2): | |
| s += LABEL_OVERLAP_PENALTY | |
| # Bonus for value being to the right of label (same row) or below | |
| bb_cx, bb_cy = _bbox_center(bb) | |
| lbl_cx, lbl_cy = _bbox_center(lbl) | |
| if bb_cx > lbl_cx or bb_cy > lbl_cy + 10: | |
| # Good position β slight bonus (reduce score) | |
| s *= 0.8 | |
| elif bb_cy < lbl_cy - 20 and bb_cx < lbl_cx: | |
| # Value is above-left of label β unlikely | |
| s += WRONG_DIRECTION_PENALTY | |
| return s | |
| def score(bb: BoundingBox) -> float: | |
| s = 0.0 | |
| if all_labels: | |
| # Try each label position and take the best score | |
| s = min(score_with_label(bb, lbl) for lbl in all_labels) | |
| # Surrounding-text context bonus β strongly favours candidates | |
| # whose nearby OCR words contain the field's label tokens. | |
| if words and label_tokens and is_ambiguous: | |
| ctx_score = _surrounding_text_score(bb, words, label_tokens, radius=6) | |
| # Reduce the score (= make it better) proportionally | |
| # A perfect context match (1.0) can override 500px of distance | |
| s -= ctx_score * CONTEXT_BONUS_MAX | |
| # Penalty for overlapping already-claimed regions | |
| for c in claimed: | |
| if _bbox_overlaps(bb, c): | |
| s += CLAIMED_PENALTY | |
| return s | |
| scored = sorted(candidates, key=score) | |
| return scored[0] if scored else None | |
| def generate_bounding_boxes(image: Image.Image, result: "ExtractionResult") -> list[BoundingBox]: | |
| """Generate bounding boxes by matching extracted field values to OCR word positions. | |
| Only returns bounding boxes when OCR can reliably match at least 30% of | |
| extracted fields. For handwritten or complex documents where OCR can't read | |
| the text (but VLM could), returns an empty list β no boxes is better than | |
| wrong boxes. | |
| """ | |
| # Count how many fields actually have values | |
| field_values = _get_field_values(result) | |
| numeric_fields = _get_numeric_fields(result) | |
| total_fields_with_values = sum(1 for v in field_values.values() if v) + len(numeric_fields) | |
| if total_fields_with_values == 0: | |
| return [] | |
| # Try OCR-based matching | |
| words = _get_ocr_word_boxes(image) | |
| if not words: | |
| logger.info("No OCR words detected β skipping bounding boxes") | |
| return [] | |
| bboxes = [] | |
| claimed: list[BoundingBox] = [] # track assigned regions to avoid duplicates | |
| for field_name, value in field_values.items(): | |
| if not value: | |
| continue | |
| # Try alternative representations for dates and currency | |
| if "date" in field_name: | |
| search_variants = _date_search_variants(value) | |
| elif field_name == "currency": | |
| search_variants = _currency_search_variants(value) | |
| else: | |
| search_variants = [value] | |
| # Gather ALL matches across all variants | |
| all_matches: list[BoundingBox] = [] | |
| for variant in search_variants: | |
| all_matches.extend(_find_all_text_bboxes(words, variant)) | |
| # Fallback: for multi-word values with no matches, try progressively | |
| # shorter sub-phrases, then individual distinctive words | |
| if not all_matches and len(value.split()) > 1: | |
| tokens = value.split() | |
| # Try progressively shorter windows, longest first | |
| for window in range(len(tokens) - 1, 0, -1): | |
| for start in range(len(tokens) - window + 1): | |
| sub = " ".join(tokens[start:start + window]) | |
| # For single words, require length >= 5 to be distinctive | |
| min_len = 5 if window == 1 else 4 | |
| if len(sub) >= min_len: | |
| sub_matches = _find_all_text_bboxes(words, sub) | |
| all_matches.extend(sub_matches) | |
| if all_matches: | |
| break # found matches at this window size | |
| if all_matches: | |
| # Find ALL label positions for context-aware disambiguation | |
| label_bboxes = _find_label_bboxes(words, field_name) | |
| label_bbox = label_bboxes[0] if label_bboxes else None | |
| best = _pick_best_match(all_matches, label_bbox, claimed, | |
| label_bboxes=label_bboxes, words=words, | |
| field_name=field_name) | |
| if best: | |
| best.field = field_name | |
| bboxes.append(best) | |
| claimed.append(best) | |
| for field_name, num_value in numeric_fields.items(): | |
| search_variants = _numeric_search_variants(num_value) | |
| all_matches = [] | |
| for variant in search_variants: | |
| all_matches.extend(_find_all_text_bboxes(words, variant)) | |
| if all_matches: | |
| label_bboxes = _find_label_bboxes(words, field_name) | |
| label_bbox = label_bboxes[0] if label_bboxes else None | |
| best = _pick_best_match(all_matches, label_bbox, claimed, | |
| label_bboxes=label_bboxes, words=words, | |
| field_name=field_name) | |
| if best: | |
| best.field = field_name | |
| bboxes.append(best) | |
| claimed.append(best) | |
| # Line items β try to locate each row by its description and/or amount. | |
| # Uses claimed list so successive line items with the same amount | |
| # (e.g. two items at $50) each get their own bbox. | |
| for idx, li in enumerate(result.line_items): | |
| li_field = f"line_item_{idx}" | |
| # Support both LineItem objects and raw dicts | |
| if isinstance(li, dict): | |
| li_desc = str(li.get("description", "") or "") | |
| # For non-standard schemas, pick the longest string value as | |
| # the "description" for bbox matching (e.g. transaction_description) | |
| if not li_desc: | |
| str_vals = [(k, str(v)) for k, v in li.items() | |
| if isinstance(v, str) and len(str(v)) > 3 | |
| and k not in ("amount", "quantity", "unit_price")] | |
| if str_vals: | |
| li_desc = max(str_vals, key=lambda x: len(x[1]))[1] | |
| li_amount = li.get("amount") | |
| else: | |
| li_desc = li.description | |
| li_amount = li.amount | |
| # Find best unclaimed description match | |
| desc_bbox = None | |
| if li_desc: | |
| desc_matches = _find_all_text_bboxes(words, li_desc) | |
| desc_bbox = _pick_best_match(desc_matches, None, claimed) if desc_matches else None | |
| # Find best unclaimed amount match | |
| amt_bbox = None | |
| if li_amount is not None and li_amount != 0: | |
| amt_matches = [] | |
| for variant in _numeric_search_variants(li_amount): | |
| amt_matches.extend(_find_all_text_bboxes(words, variant)) | |
| if amt_matches: | |
| # Prefer an amount on the same line as the description | |
| amt_bbox = _pick_best_match(amt_matches, desc_bbox, claimed) | |
| if desc_bbox and amt_bbox: | |
| # Only merge if both are on roughly the same line (within 60px at 200 DPI) | |
| desc_mid_y = (desc_bbox.y1 + desc_bbox.y2) / 2 | |
| amt_mid_y = (amt_bbox.y1 + amt_bbox.y2) / 2 | |
| if abs(desc_mid_y - amt_mid_y) < 60: | |
| merged = BoundingBox( | |
| field=li_field, | |
| x1=min(desc_bbox.x1, amt_bbox.x1), | |
| y1=min(desc_bbox.y1, amt_bbox.y1), | |
| x2=max(desc_bbox.x2, amt_bbox.x2), | |
| y2=max(desc_bbox.y2, amt_bbox.y2), | |
| ) | |
| bboxes.append(merged) | |
| claimed.append(merged) | |
| else: | |
| desc_bbox.field = li_field | |
| bboxes.append(desc_bbox) | |
| claimed.append(desc_bbox) | |
| elif desc_bbox: | |
| desc_bbox.field = li_field | |
| bboxes.append(desc_bbox) | |
| claimed.append(desc_bbox) | |
| elif amt_bbox: | |
| amt_bbox.field = li_field | |
| bboxes.append(amt_bbox) | |
| claimed.append(amt_bbox) | |
| # Only return OCR matches if we matched a meaningful portion of fields. | |
| # Sparse matches (1-2 out of 10) are usually noise on handwritten docs. | |
| match_ratio = len(bboxes) / total_fields_with_values if total_fields_with_values > 0 else 0 | |
| # Log which fields were NOT matched for diagnostics | |
| matched_fields = {bb.field for bb in bboxes} | |
| all_field_names = set(k for k, v in field_values.items() if v) | set(numeric_fields.keys()) | |
| unmatched = all_field_names - matched_fields | |
| if unmatched: | |
| logger.info(f"Unmatched fields: {sorted(unmatched)}") | |
| if match_ratio < 0.2: | |
| logger.info( | |
| f"OCR matched only {len(bboxes)}/{total_fields_with_values} fields " | |
| f"({match_ratio:.0%}) β too few, skipping bounding boxes" | |
| ) | |
| return [] | |
| logger.info(f"OCR matched {len(bboxes)}/{total_fields_with_values} fields ({match_ratio:.0%})") | |
| return bboxes | |
| def _get_field_values(result: "ExtractionResult") -> dict: | |
| """Get string field nameβvalue mapping for bbox matching. | |
| Prefers extracted_fields (works for any document type). | |
| Falls back to legacy named fields for older responses. | |
| """ | |
| # Primary: extracted_fields dict (dynamic, works for all use cases) | |
| if result.extracted_fields: | |
| return { | |
| k: str(v) for k, v in result.extracted_fields.items() | |
| if v is not None and not isinstance(v, (int, float)) | |
| } | |
| # Fallback: legacy named fields (invoice-specific) | |
| return { | |
| "vendor_name": result.vendor_name, | |
| "vendor_address": result.vendor_address, | |
| "invoice_number": result.invoice_number, | |
| "invoice_date": result.invoice_date, | |
| "due_date": result.due_date, | |
| "currency": result.currency, | |
| "customer_name": result.customer_name, | |
| "customer_address": result.customer_address, | |
| "payment_terms": result.payment_terms, | |
| } | |
| def _get_numeric_fields(result: "ExtractionResult") -> dict: | |
| """Get numeric field nameβvalue mapping for bbox matching. | |
| Prefers extracted_fields (works for any document type). | |
| Falls back to legacy named fields for older responses. | |
| """ | |
| # Primary: extracted_fields dict (dynamic, works for all use cases) | |
| if result.extracted_fields: | |
| return { | |
| k: v for k, v in result.extracted_fields.items() | |
| if isinstance(v, (int, float)) and v is not None | |
| } | |
| # Fallback: legacy named fields | |
| fields = {} | |
| if result.subtotal is not None: | |
| fields["subtotal"] = result.subtotal | |
| if result.tax_amount is not None: | |
| fields["tax_amount"] = result.tax_amount | |
| if result.total_amount is not None: | |
| fields["total_amount"] = result.total_amount | |
| return fields | |
| def _numeric_search_variants(value: float) -> list[str]: | |
| """Generate multiple string representations of a number for OCR matching.""" | |
| variants = [] | |
| # Integer form if it's a whole number: "353" | |
| if value == int(value): | |
| variants.append(str(int(value))) | |
| # Standard float: "353.0" | |
| variants.append(str(value)) | |
| # Two-decimal: "353.00" | |
| variants.append(f"{value:.2f}") | |
| # Comma-separated thousands (Western): "1,250.00" | |
| if value >= 1000: | |
| variants.append(f"{value:,.2f}") | |
| variants.append(f"{value:,.0f}" if value == int(value) else f"{value:,.2f}") | |
| # Indian lakh/crore grouping: "1,23,456.00" | |
| if value >= 1000: | |
| int_part = int(value) | |
| dec_part = f"{value:.2f}".split(".")[1] | |
| s = str(int_part) | |
| if len(s) > 3: | |
| last3 = s[-3:] | |
| rest = s[:-3] | |
| # Group remaining digits in pairs from right | |
| groups = [] | |
| while rest: | |
| groups.append(rest[-2:] if len(rest) >= 2 else rest) | |
| rest = rest[:-2] | |
| groups.reverse() | |
| indian = ",".join(groups) + "," + last3 | |
| else: | |
| indian = s | |
| variants.append(f"{indian}.{dec_part}") | |
| if value == int(value): | |
| variants.append(indian) | |
| return list(dict.fromkeys(variants)) # deduplicate, preserve order | |
| # --------------------------------------------------------------------------- | |
| # Prompt Library β per use-case extraction prompts | |
| # --------------------------------------------------------------------------- | |
| # The library maps (stage, use_case) β prompt text. | |
| # If a use_case-specific prompt doesn't exist, falls back to "_default". | |
| EXTRACTION_PROMPTS: dict[str, str] = {} | |
| EXTRACTION_PROMPTS["invoice"] = """You are a document extraction AI. Extract ALL structured data from this invoice, receipt, or financial document. | |
| Return ONLY valid JSON with this schema (use null for missing fields): | |
| { | |
| "document_type": "invoice", | |
| "confidence": 0.0 to 1.0, | |
| "fields": { | |
| "vendor_name": "string", | |
| "vendor_address": "string", | |
| "invoice_number": "string or null", | |
| "invoice_date": "string (YYYY-MM-DD)", | |
| "due_date": "string (YYYY-MM-DD) or null", | |
| "currency": "string (ISO 4217) or null", | |
| "subtotal": number or null, | |
| "tax_amount": number or null, | |
| "tax_rate": "string (e.g. 18% GST) or null", | |
| "total_amount": number, | |
| "payment_terms": "string or null", | |
| "customer_name": "string or null", | |
| "customer_address": "string or null", | |
| "...any_other_field": "extract ALL other visible fields as additional flat key-value pairs here" | |
| }, | |
| "line_items": [{"description": "string", "quantity": number or null, "unit_price": number or null, "amount": number or null}] | |
| } | |
| IMPORTANT RULES: | |
| - confidence: your self-assessed confidence in the extraction accuracy (0.0 = pure guess, 1.0 = perfectly clear). | |
| - Be precise with numbers. Parse dates into YYYY-MM-DD format. | |
| - currency: ONLY set this if a currency symbol or code is VISIBLE. Do NOT infer. | |
| - line_items: Extract ONLY what is explicitly written. Do NOT invent quantity or unit_price. Never default quantity to 1. | |
| - Subtotal: only set if the document explicitly labels a subtotal. Do not copy total_amount into subtotal. | |
| - CRITICAL: You MUST include ALL other visible fields directly inside "fields" as additional key-value pairs using snake_case names. Examples: "gstin", "hsn_code", "po_number", "email", "phone", "account_number", "credit_card_no", "statement_date", "minimum_due", "reward_points", "disbursed", etc. Do NOT limit yourself to only the fields listed in the schema above. Every piece of text data visible in the document should appear as a field. The schema above is a MINIMUM β add every other field you can find.""" | |
| EXTRACTION_PROMPTS["tax_form"] = """You are a document extraction AI specializing in tax forms, challans, and government tax documents. | |
| Extract ALL structured data from this tax document. | |
| Return ONLY valid JSON: | |
| { | |
| "document_type": "tax_form", | |
| "confidence": 0.0 to 1.0, | |
| "fields": { | |
| "form_type": "string (e.g. Form 26QB, Challan 280, Form 16, TDS Certificate)", | |
| "pan": "string or null", | |
| "tan": "string or null", | |
| "assessment_year": "string (e.g. 2026-27) or null", | |
| "financial_year": "string or null", | |
| "nature_of_payment": "string or null", | |
| "section": "string (e.g. 194IA, 234E) or null", | |
| "deductor_name": "string or null", | |
| "deductee_name": "string or null", | |
| "amount_paid": number or null, | |
| "tds_amount": number or null, | |
| "surcharge": number or null, | |
| "education_cess": number or null, | |
| "total_tax_deposited": number or null, | |
| "date_of_payment": "string (YYYY-MM-DD) or null", | |
| "date_of_deduction": "string (YYYY-MM-DD) or null", | |
| "acknowledgement_number": "string or null", | |
| "challan_number": "string or null", | |
| "bsr_code": "string or null", | |
| "demand_reference_number": "string or null" | |
| }, | |
| "line_items": [] | |
| } | |
| IMPORTANT RULES: | |
| - Include ALL visible fields, even those not listed above β add them to "fields". | |
| - Use snake_case for any extra field names, derived from the label in the document. | |
| - Parse dates into YYYY-MM-DD format. | |
| - Numbers should be numeric types, not strings. | |
| - PAN and TAN should be uppercase strings. | |
| - confidence: your self-assessed confidence in the extraction accuracy (0.0 = pure guess, 1.0 = perfectly clear).""" | |
| EXTRACTION_PROMPTS["receipt"] = """You are a document extraction AI. Extract ALL structured data from this receipt or purchase record. | |
| Return ONLY valid JSON: | |
| { | |
| "document_type": "receipt", | |
| "confidence": 0.0 to 1.0, | |
| "fields": { | |
| "store_name": "string", | |
| "store_address": "string or null", | |
| "receipt_number": "string or null", | |
| "date": "string (YYYY-MM-DD)", | |
| "currency": "string (ISO 4217) or null", | |
| "subtotal": number or null, | |
| "tax_amount": number or null, | |
| "tax_rate": "string or null", | |
| "total_amount": number, | |
| "payment_method": "string or null", | |
| "card_last_four": "string or null" | |
| }, | |
| "line_items": [{"description": "string", "quantity": number or null, "unit_price": number or null, "amount": number or null}] | |
| } | |
| IMPORTANT RULES: | |
| - Include ALL visible fields, even those not listed β add them to "fields". | |
| - currency: ONLY set if a symbol or code is VISIBLE. Do NOT infer. | |
| - line_items: Extract ONLY what is explicitly written. Never invent quantity. | |
| - Parse dates into YYYY-MM-DD format. | |
| - confidence: your self-assessed confidence in the extraction accuracy (0.0 = pure guess, 1.0 = perfectly clear).""" | |
| EXTRACTION_PROMPTS["_default"] = """You are a document extraction AI. Analyze this document and extract ALL visible structured data. | |
| Steps: | |
| 1. Identify the document type (e.g. invoice, receipt, tax_form, contract, bank_statement, id_card, letter, application, certificate, etc.) | |
| 2. Extract every labeled field visible in the document as key-value pairs. | |
| 3. Extract any tabular data as line_items. | |
| Return ONLY valid JSON: | |
| { | |
| "document_type": "string", | |
| "confidence": 0.0 to 1.0, | |
| "fields": { | |
| "field_name": "value", | |
| ... | |
| }, | |
| "line_items": [ | |
| {"column_name": "value", ...}, | |
| ... | |
| ] | |
| } | |
| IMPORTANT RULES: | |
| - confidence: your self-assessed confidence in the extraction accuracy (0.0 = pure guess, 1.0 = perfectly clear). | |
| - Use snake_case for field names, derived from the label text in the document. | |
| - Parse dates into YYYY-MM-DD format. | |
| - Numbers should be numeric types, not strings. | |
| - Include ALL visible fields β headers, form fields, metadata, identifiers, dates, amounts. | |
| - For tables, use column headers as keys in each row object. | |
| - Do NOT infer or guess values. Only extract what is explicitly visible.""" | |
| # --------------------------------------------------------------------------- | |
| # Validation Prompts | |
| # --------------------------------------------------------------------------- | |
| VALIDATION_PROMPTS: dict[str, str] = {} | |
| VALIDATION_PROMPTS["invoice"] = """You are a document validation AI. You will be given: | |
| 1. An image of the original document. | |
| 2. A JSON object of extracted data. | |
| Verify EVERY extracted field against the document image. For each field, check: | |
| - Is the value present in the document? | |
| - Is it accurately transcribed (exact text, numbers, dates)? | |
| - Are numeric fields consistent (subtotal + tax = total, qty Γ price = amount)? | |
| Return ONLY valid JSON: | |
| { | |
| "checks": [ | |
| { | |
| "field": "field_name", | |
| "status": "pass" | "fail" | "warn", | |
| "message": "short explanation", | |
| "expected": "what the document shows (if different)", | |
| "actual": "what was extracted" | |
| } | |
| ], | |
| "summary": "short overall assessment" | |
| } | |
| Focus on: amounts, dates, vendor/customer names, invoice numbers, tax calculations. | |
| Flag any field where the extracted value doesn't match what's visible in the image.""" | |
| VALIDATION_PROMPTS["tax_form"] = """You are a document validation AI specializing in tax forms and government documents. | |
| You will be given: | |
| 1. An image of the original tax document. | |
| 2. A JSON object of extracted data. | |
| Verify EVERY extracted field against the document image. For tax forms, pay special attention to: | |
| - PAN/TAN format validity (10 chars: 5 letters, 4 digits, 1 letter) | |
| - Assessment year / financial year consistency | |
| - TDS amount + surcharge + cess = total tax deposited | |
| - Date fields match what's printed | |
| - Section numbers and challan/acknowledgement numbers are exact | |
| Return ONLY valid JSON: | |
| { | |
| "checks": [ | |
| { | |
| "field": "field_name", | |
| "status": "pass" | "fail" | "warn", | |
| "message": "short explanation", | |
| "expected": "what the document shows (if different)", | |
| "actual": "what was extracted" | |
| } | |
| ], | |
| "summary": "short overall assessment" | |
| }""" | |
| VALIDATION_PROMPTS["receipt"] = """You are a document validation AI. You will be given: | |
| 1. An image of the original receipt. | |
| 2. A JSON object of extracted data. | |
| Verify EVERY extracted field against the receipt image. For receipts, check: | |
| - Store name and details match the header | |
| - Line item amounts are correct | |
| - Subtotal/tax/total arithmetic | |
| - Date and receipt number accuracy | |
| - Payment method matches what's printed | |
| Return ONLY valid JSON: | |
| { | |
| "checks": [ | |
| { | |
| "field": "field_name", | |
| "status": "pass" | "fail" | "warn", | |
| "message": "short explanation", | |
| "expected": "what the document shows (if different)", | |
| "actual": "what was extracted" | |
| } | |
| ], | |
| "summary": "short overall assessment" | |
| }""" | |
| VALIDATION_PROMPTS["_default"] = """You are a document validation AI. You will be given: | |
| 1. An image of the original document. | |
| 2. A JSON object of extracted data. | |
| Verify EVERY extracted field against the document image: | |
| - Is each field value present and accurately transcribed? | |
| - Are numeric fields consistent with each other? | |
| - Are dates in the correct format and matching the document? | |
| - Are identifiers (numbers, codes, references) exact? | |
| Return ONLY valid JSON: | |
| { | |
| "checks": [ | |
| { | |
| "field": "field_name", | |
| "status": "pass" | "fail" | "warn", | |
| "message": "short explanation", | |
| "expected": "what the document shows (if different)", | |
| "actual": "what was extracted" | |
| } | |
| ], | |
| "summary": "short overall assessment" | |
| }""" | |
| # Stage-level prompt library: maps (stage, use_case) β prompt. | |
| PROMPT_LIBRARY: dict[str, dict[str, str]] = { | |
| "extraction": EXTRACTION_PROMPTS, | |
| "validation": VALIDATION_PROMPTS, | |
| } | |
| def get_prompt(stage: str, use_case: str) -> str: | |
| """Look up the prompt for a (stage, use_case) pair. | |
| Priority: | |
| 1. Active prompt from Supabase prompt_versions table (if configured) | |
| 2. Hardcoded PROMPT_LIBRARY dict (in-code defaults) | |
| 3. _default prompt for the stage | |
| """ | |
| # Try DB first β gracefully returns None when Supabase is not configured | |
| try: | |
| from db.supabase import get_active_prompt | |
| db_prompt = get_active_prompt(use_case, stage) | |
| if db_prompt: | |
| return db_prompt | |
| except Exception: | |
| pass # Fall through to hardcoded prompts | |
| stage_prompts = PROMPT_LIBRARY.get(stage, {}) | |
| return stage_prompts.get(use_case, stage_prompts.get("_default", "")) | |
| # Legacy alias kept so callers that reference EXTRACTION_PROMPT still work | |
| EXTRACTION_PROMPT = EXTRACTION_PROMPTS["invoice"] | |
| def image_to_base64(image: Image.Image, max_size: int = 1024) -> str: | |
| """Convert PIL Image to base64 string, resizing if needed. | |
| IMPORTANT: Works on a copy so the original image is never mutated. | |
| This preserves full resolution for downstream OCR bounding-box generation. | |
| """ | |
| img = image.copy() | |
| if max(img.size) > max_size: | |
| img.thumbnail((max_size, max_size), Image.LANCZOS) | |
| if img.mode == "RGBA": | |
| img = img.convert("RGB") | |
| buffer = BytesIO() | |
| img.save(buffer, format="JPEG", quality=85) | |
| return base64.b64encode(buffer.getvalue()).decode("utf-8") | |
| def _parse_vlm_json(raw: str) -> dict: | |
| """Extract JSON from a VLM response, handling markdown fences.""" | |
| json_match = re.search(r"```(?:json)?\s*([\s\S]*?)```", raw) | |
| if json_match: | |
| return json.loads(json_match.group(1).strip()) | |
| json_match = re.search(r"\{[\s\S]*\}", raw) | |
| if json_match: | |
| return json.loads(json_match.group(0)) | |
| return json.loads(raw) | |
| def _compute_confidence(fields: dict, line_items: list, model_id: str) -> float: | |
| """Compute an extraction confidence score based on field completeness. | |
| Heuristic factors: | |
| - Ratio of non-null/non-empty fields β higher is better | |
| - Having line items when fields suggest a tabular document | |
| - Model tier bonus (larger models are more reliable) | |
| """ | |
| if not fields: | |
| return 0.3 | |
| # Count meaningful fields (non-null, non-empty) | |
| total = len(fields) | |
| filled = 0 | |
| for v in fields.values(): | |
| if v is None: | |
| continue | |
| if isinstance(v, str) and v.strip() in ("", "-", "N/A", "n/a", "None", "none"): | |
| continue | |
| filled += 1 | |
| if total == 0: | |
| return 0.3 | |
| field_ratio = filled / total | |
| # Line-item bonus: having structured rows is a strong signal | |
| li_bonus = min(0.05, len(line_items) * 0.01) if line_items else 0.0 | |
| # Model tier adjustment | |
| model_lower = model_id.lower() | |
| if "72b" in model_lower: | |
| tier_bonus = 0.05 | |
| elif "gemini" in model_lower: | |
| tier_bonus = 0.03 # Gemini Flash: between 72B and 7B quality | |
| elif "7b" in model_lower: | |
| tier_bonus = 0.0 | |
| else: | |
| tier_bonus = -0.05 | |
| # Base confidence: scale field ratio into 0.55β0.95 range | |
| confidence = 0.55 + (field_ratio * 0.40) + li_bonus + tier_bonus | |
| return round(max(0.1, min(0.99, confidence)), 2) | |
| def _vlm_data_to_result(data: dict, model_id: str, use_case: str) -> ExtractionResult: | |
| """Convert parsed VLM JSON into an ExtractionResult. | |
| Handles two response shapes: | |
| 1. New format: {"document_type": "...", "fields": {...}, "line_items": [...]} | |
| 2. Legacy flat format: {"vendor_name": "...", "total_amount": 100, ...} | |
| """ | |
| short_model = model_id.split("/")[-1] | |
| # Detect response shape | |
| fields = data.get("fields", {}) | |
| doc_type = data.get("document_type", use_case) | |
| if not fields: | |
| # Legacy flat format β all keys are fields (except line_items) | |
| fields = {k: v for k, v in data.items() | |
| if k not in ("line_items", "document_type", "fields", "confidence")} | |
| line_items = data.get("line_items", []) | |
| # Use model-reported confidence if present, otherwise compute heuristic | |
| reported_confidence = data.get("confidence") | |
| if isinstance(reported_confidence, (int, float)) and 0 < reported_confidence <= 1: | |
| confidence = round(float(reported_confidence), 2) | |
| else: | |
| confidence = _compute_confidence(fields, line_items, model_id) | |
| result = ExtractionResult( | |
| extraction_method=f"vlm:{short_model}", | |
| confidence=confidence, | |
| document_type=doc_type, | |
| extracted_fields=fields, | |
| ) | |
| # Populate legacy named fields from fields dict (backward compat for invoice) | |
| result.vendor_name = str(fields.get("vendor_name", "") or "") | |
| result.vendor_address = str(fields.get("vendor_address", "") or "") | |
| result.invoice_number = str(fields.get("invoice_number", "") or "") | |
| result.invoice_date = str(fields.get("invoice_date", "") or "") | |
| result.due_date = str(fields.get("due_date", "") or "") | |
| result.currency = str(fields.get("currency", "") or "") | |
| result.subtotal = fields.get("subtotal") | |
| result.tax_amount = fields.get("tax_amount") | |
| result.tax_rate = str(fields.get("tax_rate", "") or "") | |
| result.total_amount = fields.get("total_amount") | |
| result.payment_terms = str(fields.get("payment_terms", "") or "") | |
| result.customer_name = str(fields.get("customer_name", "") or "") | |
| result.customer_address = str(fields.get("customer_address", "") or "") | |
| # Parse line items β preserve the raw dict so the frontend can display | |
| # whatever columns the VLM returned (e.g. date, transaction_description | |
| # for credit-card statements). Only coerce into the rigid LineItem | |
| # dataclass when the item carries exactly the invoice-schema keys. | |
| _LINEITEM_KEYS = {"description", "quantity", "unit_price", "amount"} | |
| for item in data.get("line_items", []): | |
| if isinstance(item, dict): | |
| item_keys = set(item.keys()) | |
| if item_keys <= _LINEITEM_KEYS: | |
| # Standard invoice row β use LineItem for backward compat | |
| result.line_items.append(LineItem( | |
| description=item.get("description", ""), | |
| quantity=item.get("quantity"), | |
| unit_price=item.get("unit_price"), | |
| amount=item.get("amount"), | |
| )) | |
| else: | |
| # Non-standard columns (credit-card statement, etc.) | |
| # Keep the raw dict so the frontend's dynamic column | |
| # builder can show all fields. | |
| result.line_items.append(item) | |
| return result | |
| def _is_gemini_model(model_id: str) -> bool: | |
| """Check whether a model id refers to a Gemini model.""" | |
| return model_id.lower().startswith("gemini") | |
| def _extract_with_gemini( | |
| image: Image.Image, | |
| model_id: str, | |
| use_case: str, | |
| ) -> Optional[ExtractionResult]: | |
| """Extract using Google Gemini Vision API.""" | |
| prompt = get_prompt("extraction", use_case) | |
| if not prompt: | |
| logger.warning(f"No extraction prompt for use_case={use_case}, using default") | |
| prompt = get_prompt("extraction", "_default") | |
| img_b64 = image_to_base64(image) | |
| try: | |
| result_text = extract_with_gemini(img_b64, prompt, model=model_id) | |
| data = _parse_vlm_json(result_text) | |
| return _vlm_data_to_result(data, model_id, use_case) | |
| except Exception as e: | |
| error_msg = str(e)[:200] | |
| logger.warning(f"Gemini extraction failed with {model_id}: {e}") | |
| return ExtractionResult( | |
| extraction_method=f"failed:gemini({model_id})", | |
| confidence=0.0, | |
| raw_text=f"{model_id}: {type(e).__name__}({error_msg})", | |
| ) | |
| def _extract_with_hf( | |
| image: Image.Image, | |
| hf_token: str, | |
| model_id: str, | |
| use_case: str, | |
| ) -> Optional[ExtractionResult]: | |
| """Extract using HF Inference API. When model_id is a specific Qwen model, | |
| tries only that model. When empty or unrecognised, cascades through all. | |
| """ | |
| prompt = get_prompt("extraction", use_case) | |
| if not prompt: | |
| logger.warning(f"No extraction prompt for use_case={use_case}, using default") | |
| prompt = get_prompt("extraction", "_default") | |
| # Determine which models to try | |
| all_models = [ | |
| "Qwen/Qwen2.5-VL-72B-Instruct", | |
| "Qwen/Qwen2.5-VL-7B-Instruct", | |
| "Qwen/Qwen2.5-VL-3B-Instruct", | |
| ] | |
| if model_id and model_id in all_models: | |
| # User picked a specific model β try it first, then cascade | |
| models = [model_id] + [m for m in all_models if m != model_id] | |
| else: | |
| models = all_models | |
| img_b64 = image_to_base64(image) | |
| data_url = f"data:image/jpeg;base64,{img_b64}" | |
| errors = [] | |
| for mid in models: | |
| try: | |
| client = create_inference_client(api_key=hf_token) | |
| response = robust_chat_completion( | |
| client, | |
| model=mid, | |
| messages=[ | |
| { | |
| "role": "user", | |
| "content": [ | |
| {"type": "text", "text": prompt}, | |
| { | |
| "type": "image_url", | |
| "image_url": {"url": data_url}, | |
| }, | |
| ], | |
| } | |
| ], | |
| max_tokens=2048, | |
| temperature=0.1, | |
| ) | |
| result_text = response.choices[0].message.content | |
| data = _parse_vlm_json(result_text) | |
| return _vlm_data_to_result(data, mid, use_case) | |
| except Exception as e: | |
| short_model = mid.split("/")[-1] | |
| error_msg = str(e)[:150] | |
| errors.append(f"{short_model}: {type(e).__name__}({error_msg})") | |
| logger.warning(f"VLM extraction failed with {mid}: {e}") | |
| continue | |
| return ExtractionResult( | |
| extraction_method="failed:vlm_all_models", | |
| confidence=0.0, | |
| raw_text="; ".join(errors) if errors else "all models failed silently", | |
| ) | |
| def extract_with_vlm( | |
| image: Image.Image, | |
| hf_token: str, | |
| use_case: str = "invoice", | |
| model_id: str = "", | |
| ) -> Optional[ExtractionResult]: | |
| """Extract document data using a Vision-Language Model. | |
| Routes to the appropriate provider based on model_id: | |
| - gemini-* β Google Gemini API (requires GEMINI_API_KEY) | |
| - Qwen/* β HF Inference API (requires hf_token) | |
| - empty β Gemini if GEMINI_API_KEY is set, else HF cascade | |
| Selects the prompt from the prompt library based on use_case. | |
| Falls back to the generic prompt when no use_case-specific prompt exists. | |
| """ | |
| # Route: Gemini models | |
| if _is_gemini_model(model_id): | |
| if not GEMINI_API_KEY: | |
| logger.warning("Gemini model requested but GEMINI_API_KEY not set, falling back to HF") | |
| return _extract_with_hf(image, hf_token, "", use_case) | |
| return _extract_with_gemini(image, model_id, use_case) | |
| # Route: HF models (explicit Qwen selection or fallback) | |
| if model_id and model_id.startswith("Qwen/"): | |
| return _extract_with_hf(image, hf_token, model_id, use_case) | |
| # Route: No specific model β prefer Gemini if configured, else HF cascade | |
| if GEMINI_API_KEY: | |
| logger.info("No model specified, defaulting to Gemini Flash") | |
| return _extract_with_gemini(image, "gemini-3.5-flash-lite", use_case) | |
| return _extract_with_hf(image, hf_token, "", use_case) | |
| # --------------------------------------------------------------------------- | |
| # OCR + Regex Extraction (Fallback) - CPU only, zero cost | |
| # --------------------------------------------------------------------------- | |
| def extract_text_ocr(image: Image.Image) -> str: | |
| """Extract text from image using Tesseract OCR.""" | |
| try: | |
| import pytesseract | |
| text = pytesseract.image_to_string(image, config="--psm 3") | |
| return text | |
| except ImportError: | |
| logger.warning("pytesseract not available") | |
| return "" | |
| except Exception as e: | |
| logger.warning(f"OCR failed: {e}") | |
| return "" | |
| def parse_amount(text: str) -> Optional[float]: | |
| """Parse a monetary amount from text.""" | |
| text = text.replace(",", "").replace(" ", "") | |
| match = re.search(r"[\d]+\.?\d*", text) | |
| if match: | |
| try: | |
| return float(match.group()) | |
| except ValueError: | |
| pass | |
| return None | |
| def detect_currency(text: str) -> str: | |
| """Detect currency from text symbols and keywords.""" | |
| currency_patterns = { | |
| "INR": [r"βΉ", r"Rs\.?", r"INR", r"rupee"], | |
| "USD": [r"\$", r"USD", r"dollar"], | |
| "EUR": [r"β¬", r"EUR", r"euro"], | |
| "GBP": [r"Β£", r"GBP", r"pound"], | |
| } | |
| for currency, patterns in currency_patterns.items(): | |
| for p in patterns: | |
| if re.search(p, text, re.IGNORECASE): | |
| return currency | |
| return "USD" # default | |
| def extract_with_regex(text: str) -> ExtractionResult: | |
| """Extract structured data from OCR text using regex patterns.""" | |
| result = ExtractionResult( | |
| raw_text=text, | |
| extraction_method="ocr:tesseract+regex", | |
| confidence=0.65, | |
| ) | |
| # Invoice number β require a separator (No, #, :) between label and value | |
| inv_patterns = [ | |
| r"(?:invoice|inv|bill)\s*(?:#|no\.?|number)\s*:?\s*([A-Z0-9\-/]+)", | |
| r"(?:invoice|inv)\s*:\s*([A-Z0-9\-/]+)", | |
| r"(?:#|no\.?)\s*:?\s*([A-Z0-9\-/]+)", | |
| ] | |
| for p in inv_patterns: | |
| m = re.search(p, text, re.IGNORECASE) | |
| if m: | |
| result.invoice_number = m.group(1).strip() | |
| break | |
| # Dates | |
| date_patterns = [ | |
| r"(\d{1,2}[/\-\.]\d{1,2}[/\-\.]\d{2,4})", | |
| r"(\d{4}[/\-\.]\d{1,2}[/\-\.]\d{1,2})", | |
| r"(\w+\s+\d{1,2},?\s+\d{4})", | |
| ] | |
| dates_found = [] | |
| for p in date_patterns: | |
| dates_found.extend(re.findall(p, text)) | |
| if dates_found: | |
| result.invoice_date = dates_found[0] | |
| if len(dates_found) > 1: | |
| result.due_date = dates_found[1] | |
| # Total amount β try grand total / total due first, then plain "total" | |
| # but skip lines that say "subtotal" or specific tax labels | |
| total_patterns = [ | |
| r"(?:grand\s*total|amount\s*due|balance\s*due|total\s*amount)\s*:?\s*(?:INR|USD|EUR|GBP|[βΉ$β¬Β£])?\s*([\d,]+\.?\d*)", | |
| r"(?<!sub)total\s*:?\s*(?:INR|USD|EUR|GBP|[βΉ$β¬Β£])?\s*([\d,]+\.?\d*)", | |
| ] | |
| for p in total_patterns: | |
| matches = list(re.finditer(p, text, re.IGNORECASE | re.MULTILINE)) | |
| if matches: | |
| # Take the last match (usually the grand total at the bottom) | |
| result.total_amount = parse_amount(matches[-1].group(1)) | |
| break | |
| # Subtotal | |
| m = re.search(r"(?:subtotal|sub\s*total)\s*:?\s*[βΉ$β¬Β£]?\s*([\d,]+\.?\d*)", text, re.IGNORECASE) | |
| if m: | |
| result.subtotal = parse_amount(m.group(1)) | |
| # Tax | |
| tax_patterns = [ | |
| r"(?:tax|gst|vat|igst|cgst|sgst)\s*(?:\(?\d+%?\)?)?\s*:?\s*[βΉ$β¬Β£]?\s*([\d,]+\.?\d*)", | |
| ] | |
| for p in tax_patterns: | |
| m = re.search(p, text, re.IGNORECASE) | |
| if m: | |
| result.tax_amount = parse_amount(m.group(1)) | |
| break | |
| # Tax rate β also match "CGST (9%)" and "GST @18%" patterns | |
| tax_rate_patterns = [ | |
| r"(\d+(?:\.\d+)?)\s*%\s*(?:tax|gst|vat|igst)", | |
| r"(?:tax|gst|vat|cgst|sgst|igst)\s*[@(]?\s*(\d+(?:\.\d+)?)\s*%", | |
| ] | |
| for p in tax_rate_patterns: | |
| m = re.search(p, text, re.IGNORECASE) | |
| if m: | |
| result.tax_rate = f"{m.group(1)}%" | |
| break | |
| # Currency | |
| result.currency = detect_currency(text) | |
| # Vendor name (heuristic: first line or line before address) | |
| lines = [l.strip() for l in text.split("\n") if l.strip()] | |
| if lines: | |
| result.vendor_name = lines[0] | |
| return result | |
| # --------------------------------------------------------------------------- | |
| # Main Extraction Pipeline | |
| # --------------------------------------------------------------------------- | |
| def extract_invoice( | |
| image: Image.Image, | |
| hf_token: Optional[str] = None, | |
| force_ocr: bool = False, | |
| use_case: str = "invoice", | |
| model_id: str = "", | |
| ) -> ExtractionResult: | |
| """ | |
| Main extraction pipeline. Tries VLM first, falls back to OCR+regex. | |
| Args: | |
| image: PIL Image of the document | |
| hf_token: Hugging Face API token (for VLM extraction) | |
| force_ocr: Skip VLM and use OCR only | |
| use_case: Document type for prompt selection (invoice, tax_form, receipt, etc.) | |
| model_id: Model to use (e.g. "gemini-3.5-flash-lite", "Qwen/Qwen2.5-VL-72B-Instruct"). | |
| When empty, defaults based on available API keys. | |
| Returns: | |
| ExtractionResult with all extracted fields | |
| """ | |
| start = datetime.now() | |
| result = None | |
| vlm_errors = "" | |
| # Gemini models don't need an HF token β allow VLM when Gemini is selected | |
| is_gemini = _is_gemini_model(model_id) if model_id else False | |
| has_vlm_access = (hf_token or is_gemini or GEMINI_API_KEY) and not force_ocr | |
| # Try VLM extraction first | |
| if has_vlm_access: | |
| vlm_result = extract_with_vlm(image, hf_token or "", use_case=use_case, model_id=model_id) | |
| if vlm_result and not vlm_result.extraction_method.startswith("failed"): | |
| result = vlm_result | |
| else: | |
| # VLM failed β capture diagnostics, fall through to OCR | |
| vlm_errors = vlm_result.raw_text if vlm_result else "unknown" | |
| logger.warning(f"VLM failed ({vlm_errors}), trying OCR fallback") | |
| elif not force_ocr: | |
| vlm_errors = "no_token" | |
| logger.warning("No API keys set (HF_TOKEN / GEMINI_API_KEY) β using OCR fallback only") | |
| # Fallback to OCR + regex | |
| if result is None: | |
| text = extract_text_ocr(image) | |
| if text.strip(): | |
| result = extract_with_regex(text) | |
| if vlm_errors: | |
| # Note that this was a fallback result | |
| result.extraction_method += f" (vlm failed: {vlm_errors})" | |
| else: | |
| method = f"failed:vlm({vlm_errors})+ocr_empty" if vlm_errors else "failed:ocr_empty" | |
| result = ExtractionResult( | |
| extraction_method=method, | |
| confidence=0.0, | |
| ) | |
| # Generate bounding boxes by matching extracted values to OCR word positions. | |
| # Falls back to layout-based estimation for handwritten documents. | |
| try: | |
| result.bounding_boxes = generate_bounding_boxes(image, result) | |
| except Exception as e: | |
| logger.warning(f"Bounding box generation failed: {e}") | |
| elapsed = (datetime.now() - start).total_seconds() * 1000 | |
| result.processing_time_ms = int(elapsed) | |
| return result | |
| def extract_from_pdf( | |
| pdf_bytes: bytes, | |
| hf_token: Optional[str] = None, | |
| force_ocr: bool = False, | |
| use_case: str = "invoice", | |
| max_seconds: int = 240, | |
| model_id: str = "", | |
| ) -> list[ExtractionResult]: | |
| """Extract from PDF - converts each page to image and extracts. | |
| The first page's rendered image is attached as ``page_image`` (a | |
| data-URI) so the frontend can display the PDF and overlay bounding | |
| boxes without needing a separate PDF renderer. | |
| All pages are rendered first (fast), then extraction runs per page | |
| with a time budget. If the budget is exceeded the pages already | |
| extracted are returned so the user still gets partial results and | |
| can see all page images in the viewer. | |
| """ | |
| import time as _time | |
| try: | |
| import fitz # PyMuPDF | |
| except ImportError: | |
| raise ImportError("PyMuPDF (fitz) is required for PDF processing. Install with: pip install PyMuPDF") | |
| deadline = _time.monotonic() + max_seconds | |
| doc = fitz.open(stream=pdf_bytes, filetype="pdf") | |
| num_pages = min(len(doc), 10) # Max 10 pages | |
| # Phase 1 β render all page images up-front (fast, < 1s per page) | |
| page_imgs: list[tuple[Image.Image, str]] = [] | |
| for page_num in range(num_pages): | |
| page = doc[page_num] | |
| pix = page.get_pixmap(dpi=200) | |
| img = Image.frombytes("RGB", [pix.width, pix.height], pix.samples) | |
| buf = BytesIO() | |
| img.save(buf, format="PNG", optimize=True) | |
| b64 = base64.b64encode(buf.getvalue()).decode() | |
| page_imgs.append((img, f"data:image/png;base64,{b64}")) | |
| doc.close() | |
| # Phase 2 β extract pages in parallel (I/O-bound API calls benefit | |
| # from concurrency). Cap at 2 workers to stay within Render free-tier | |
| # memory limits (512 MB) and Gemini rate limits (~15 RPM). | |
| max_workers = min(num_pages, 2) | |
| def _extract_page(page_idx: int) -> tuple[int, ExtractionResult]: | |
| img, data_uri = page_imgs[page_idx] | |
| result = extract_invoice( | |
| img, hf_token=hf_token, force_ocr=force_ocr, | |
| use_case=use_case, model_id=model_id, | |
| ) | |
| result.page_image = data_uri | |
| return page_idx, result | |
| results: list[Optional[ExtractionResult]] = [None] * num_pages | |
| extracted_count = 0 | |
| with ThreadPoolExecutor(max_workers=max_workers) as executor: | |
| futures = { | |
| executor.submit(_extract_page, idx): idx | |
| for idx in range(num_pages) | |
| } | |
| for future in as_completed(futures): | |
| idx = futures[future] | |
| try: | |
| page_idx, result = future.result() | |
| results[page_idx] = result | |
| extracted_count += 1 | |
| except Exception as e: | |
| logger.warning("PDF page %d extraction failed: %s", idx, e) | |
| results[idx] = ExtractionResult( | |
| extraction_method=f"failed:parallel_error", | |
| confidence=0.0, | |
| raw_text=str(e)[:200], | |
| ) | |
| results[idx].page_image = page_imgs[idx][1] # type: ignore[union-attr] | |
| extracted_count += 1 | |
| # Check time budget β cancel remaining futures if exceeded | |
| if _time.monotonic() > deadline and extracted_count < num_pages: | |
| logger.warning( | |
| "PDF extraction time budget exceeded after %d/%d pages", | |
| extracted_count, num_pages, | |
| ) | |
| for f in futures: | |
| f.cancel() | |
| break | |
| # Return only completed pages in order, dropping trailing Nones | |
| return [r for r in results if r is not None] | |
| # --------------------------------------------------------------------------- | |
| # VLM-based Semantic Validation | |
| # --------------------------------------------------------------------------- | |
| def validate_with_vlm( | |
| image: Image.Image, | |
| extraction_data: dict, | |
| hf_token: str, | |
| model_id: str = "Qwen/Qwen2.5-VL-72B-Instruct", | |
| use_case: str = "_default", | |
| ) -> list[dict]: | |
| """Use a VLM to cross-check extracted data against the source image. | |
| Returns a list of check dicts: {field, status, message, expected, actual}. | |
| """ | |
| if not hf_token: | |
| return [{ | |
| "name": "VLM validation", | |
| "status": "skipped", | |
| "message": "HF_TOKEN not set β cannot run VLM validation", | |
| }] | |
| prompt = get_prompt("validation", use_case) | |
| if not prompt: | |
| prompt = get_prompt("validation", "_default") | |
| # Build the extraction summary for the prompt | |
| fields = extraction_data.get("extracted_fields", {}) | |
| if not fields: | |
| # Legacy format | |
| fields = {k: v for k, v in extraction_data.items() | |
| if k not in ("line_items", "bounding_boxes", "raw_text", | |
| "extraction_method", "confidence", "processing_time_ms", | |
| "page_image", "page_images", "run_id", "error", | |
| "document_type", "extracted_fields")} | |
| extraction_json = json.dumps({ | |
| "fields": fields, | |
| "line_items": extraction_data.get("line_items", []), | |
| }, indent=2) | |
| full_prompt = f"{prompt}\n\nExtracted data to validate:\n```json\n{extraction_json}\n```" | |
| try: | |
| b64 = image_to_base64(image) | |
| client = create_inference_client(api_key=hf_token) | |
| response = robust_chat_completion( | |
| client, | |
| model=model_id, | |
| messages=[{ | |
| "role": "user", | |
| "content": [ | |
| {"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{b64}"}}, | |
| {"type": "text", "text": full_prompt}, | |
| ], | |
| }], | |
| max_tokens=2048, | |
| ) | |
| raw = response.choices[0].message.content | |
| data = _parse_vlm_json(raw) | |
| checks = data.get("checks", []) | |
| # Normalize check format | |
| result = [] | |
| for c in checks: | |
| result.append({ | |
| "name": f"VLM: {c.get('field', 'unknown')}", | |
| "status": c.get("status", "warn"), | |
| "message": c.get("message", ""), | |
| "expected": c.get("expected"), | |
| "actual": c.get("actual"), | |
| }) | |
| return result if result else [{ | |
| "name": "VLM validation", | |
| "status": "pass", | |
| "message": data.get("summary", "All fields appear correct"), | |
| }] | |
| except Exception as e: | |
| logger.warning(f"VLM validation failed: {e}") | |
| return [{ | |
| "name": "VLM validation", | |
| "status": "skipped", | |
| "message": f"VLM validation unavailable: {str(e)[:100]}", | |
| }] | |
| def run_dynamic_math_checks(data: dict) -> list[dict]: | |
| """Run math/consistency checks dynamically based on extracted_fields. | |
| Works for any document type by detecting numeric relationships. | |
| """ | |
| checks = [] | |
| fields = data.get("extracted_fields", {}) | |
| if not fields: | |
| # Fall back to top-level keys | |
| fields = data | |
| line_items = data.get("line_items", []) | |
| # ββ Collect numeric fields ββ | |
| numeric = {} | |
| for k, v in fields.items(): | |
| if v is not None: | |
| try: | |
| numeric[k] = float(v) | |
| except (ValueError, TypeError): | |
| pass | |
| # ββ Rule 1: subtotal + tax-like = total-like ββ | |
| subtotal_key = next((k for k in numeric if "subtotal" in k), None) | |
| total_key = next((k for k in numeric if "total" in k and "sub" not in k), None) | |
| tax_keys = [k for k in numeric if any(t in k for t in ("tax", "cess", "surcharge", "duty", "vat", "gst"))] | |
| if subtotal_key and total_key and tax_keys: | |
| tax_sum = sum(numeric[k] for k in tax_keys) | |
| expected = numeric[subtotal_key] + tax_sum | |
| diff = abs(expected - numeric[total_key]) | |
| checks.append({ | |
| "name": "Subtotal + taxes = Total", | |
| "status": "pass" if diff < 0.02 else "warn" if diff < 1.0 else "fail", | |
| "message": f"{subtotal_key} + {' + '.join(tax_keys)} = {total_key}", | |
| "expected": numeric[total_key], | |
| "actual": round(expected, 2), | |
| }) | |
| # ββ Rule 2: Line item amounts sum to subtotal or total ββ | |
| if line_items: | |
| item_sum = sum( | |
| float(item.get("amount", 0) or 0) for item in line_items | |
| ) | |
| compare_key = subtotal_key or total_key | |
| if compare_key and item_sum > 0: | |
| diff = abs(item_sum - numeric[compare_key]) | |
| checks.append({ | |
| "name": "Line items sum", | |
| "status": "pass" if diff < 0.02 else "warn" if diff < 1.0 else "fail", | |
| "message": f"Sum of line item amounts vs {compare_key}", | |
| "expected": numeric[compare_key], | |
| "actual": round(item_sum, 2), | |
| }) | |
| # ββ Rule 3: Each line item qty Γ unit_price = amount ββ | |
| for i, item in enumerate(line_items): | |
| qty = item.get("quantity") | |
| price = item.get("unit_price") | |
| amount = item.get("amount") | |
| if qty is not None and price is not None and amount is not None: | |
| try: | |
| expected = round(float(qty) * float(price), 2) | |
| diff = abs(expected - float(amount)) | |
| checks.append({ | |
| "name": f"Line {i+1} math", | |
| "status": "pass" if diff < 0.02 else "fail", | |
| "message": f"qty ({qty}) Γ unit_price ({price}) = amount", | |
| "expected": float(amount), | |
| "actual": expected, | |
| }) | |
| except (ValueError, TypeError): | |
| pass | |
| # ββ Rule 4: Total positive ββ | |
| if total_key: | |
| checks.append({ | |
| "name": "Total > 0", | |
| "status": "pass" if numeric[total_key] > 0 else "fail", | |
| "message": f"{total_key} should be positive", | |
| "expected": "> 0", | |
| "actual": numeric[total_key], | |
| }) | |
| # ββ Rule 5: PAN/TAN format (tax forms) ββ | |
| for k, v in fields.items(): | |
| if isinstance(v, str): | |
| if k == "pan" and v: | |
| valid = bool(re.match(r'^[A-Z]{5}[0-9]{4}[A-Z]$', v)) | |
| checks.append({ | |
| "name": "PAN format", | |
| "status": "pass" if valid else "fail", | |
| "message": "PAN should be 5 letters + 4 digits + 1 letter", | |
| "expected": "XXXXX9999X", | |
| "actual": v, | |
| }) | |
| elif k == "tan" and v: | |
| valid = bool(re.match(r'^[A-Z]{4}[0-9]{5}[A-Z]$', v)) | |
| checks.append({ | |
| "name": "TAN format", | |
| "status": "pass" if valid else "fail", | |
| "message": "TAN should be 4 letters + 5 digits + 1 letter", | |
| "expected": "XXXX99999X", | |
| "actual": v, | |
| }) | |
| if not checks: | |
| checks.append({ | |
| "name": "Insufficient data", | |
| "status": "skipped", | |
| "message": "Not enough numeric or validatable fields to run checks", | |
| }) | |
| return checks | |