""" 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 # --------------------------------------------------------------------------- @dataclass class LineItem: description: str = "" quantity: Optional[float] = None unit_price: Optional[float] = None amount: Optional[float] = None @dataclass class BoundingBox: field: str = "" x1: int = 0 y1: int = 0 x2: int = 0 y2: int = 0 page: int = 0 @dataclass 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"(? 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