agentic-extractor / extraction.py
adudeja's picture
all changes from render
04c4194
Raw History Blame Contribute Delete
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
# ---------------------------------------------------------------------------
@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"(?<!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