documentExtractionag051/ExtractDocument
0
1# form_field_extractor.py2"""3Form Field Extractor Module4 5Extracts key-value form fields from OCR data produced by pdf_utils.py.6Supports multiple matching strategies:7- Regex patterns8- Line-based matching 9- Key-value proximity detection10 11Usage:12 extractor = FormFieldExtractor()13 results = extractor.extract_form_fields(14 form_fields=["Invoice Number", "Invoice Date", "Total Amount"],15 ocr_data=ocr_blocks16 )17"""18 19import re20import logging21from typing import List, Dict, Any, Optional, Tuple22from dataclasses import dataclass, field23 24# Set up logging25logging.basicConfig(level=logging.INFO, format='%(levelname)s:%(name)s:%(message)s')26logger = logging.getLogger(__name__)27 28 29# =============================================================================30# Data Classes31# =============================================================================32 33@dataclass34class ExtractedField:35 """Represents an extracted form field with its value and metadata."""36 field_name: str37 value: str38 confidence: float = 1.039 geometry: Optional[Dict[str, int]] = None40 page_num: int = 141 match_type: str = "unknown" # regex, inline, proximity, pattern42 43 def to_dict(self) -> Dict[str, Any]:44 """Convert to dictionary format."""45 result = {46 "field_name": self.field_name,47 "value": self.value,48 "confidence": self.confidence,49 "page_num": self.page_num,50 "match_type": self.match_type51 }52 if self.geometry:53 result["geometry"] = self.geometry54 result["bounds"] = f"{self.geometry['x1']}, {self.geometry['y1']}, " \55 f"{self.geometry['x2'] - self.geometry['x1']}, " \56 f"{self.geometry['y2'] - self.geometry['y1']}"57 return result58 59 60# =============================================================================61# Field Variation Patterns62# =============================================================================63 64# Common variations for field names (key -> list of variations)65FIELD_VARIATIONS = {66 "invoice number": ["invoice no", "invoice #", "inv no", "inv #", "invoice num", "inv number"],67 "invoice date": ["inv date", "date of invoice", "invoice dt"],68 "po number": ["po no", "po #", "purchase order", "purchase order no", "p.o. number"],69 "order number": ["order no", "order #", "order id"],70 "customer name": ["customer", "cust name", "bill to", "sold to"],71 "vendor name": ["vendor", "supplier", "supplier name"],72 "total amount": ["total", "grand total", "net total", "amount due", "total due"],73 "subtotal": ["sub total", "sub-total", "net amount"],74 "tax": ["tax amount", "gst", "vat", "sales tax", "tax total"],75 "gstin": ["gst no", "gst number", "gstin no", "gst in"],76 "due date": ["payment due", "due by", "pay by"],77 "ship to": ["shipping address", "deliver to", "ship address"],78 "bill to": ["billing address", "invoice to", "billed to"],79}80 81# Separators between key and value82KEY_VALUE_SEPARATORS = [':', '-', 'โ', '|', '=', '.']83 84 85# =============================================================================86# FormFieldExtractor Class87# =============================================================================88 89class FormFieldExtractor:90 """91 Extracts form fields (key-value pairs) from OCR data.92 93 Supports multiple extraction strategies:94 1. Inline extraction: "Invoice Number: 12345" in same line95 2. Proximity extraction: Key and value in separate but adjacent elements96 3. Pattern-based extraction: Using regex patterns for specific value types97 """98 99 def __init__(100 self,101 max_horizontal_gap: float = 0.15, # 15% of page width for proximity102 max_vertical_gap: float = 0.03, # 3% of page height for proximity103 case_sensitive: bool = False104 ):105 """106 Initialize the form field extractor.107 108 Args:109 max_horizontal_gap: Maximum horizontal gap (% of page width) for value detection110 max_vertical_gap: Maximum vertical gap (% of page height) for value detection111 case_sensitive: Whether field matching should be case-sensitive112 """113 self.max_horizontal_gap = max_horizontal_gap114 self.max_vertical_gap = max_vertical_gap115 self.case_sensitive = case_sensitive116 117 # Build variation lookup118 self._variation_map = self._build_variation_map()119 120 def _build_variation_map(self) -> Dict[str, str]:121 """Build a map from variation -> canonical field name."""122 variation_map = {}123 for canonical, variations in FIELD_VARIATIONS.items():124 variation_map[canonical] = canonical125 for var in variations:126 variation_map[var] = canonical127 return variation_map128 129 def extract_form_fields(130 self,131 form_fields: List[str],132 ocr_data: List[Dict[str, Any]],133 page_dimensions: Optional[Dict[int, Tuple[int, int]]] = None134 ) -> Dict[str, Any]:135 """136 Extract and return key-value pairs for requested form fields.137 138 Args:139 form_fields: List of field names to extract (e.g., ["Invoice Number", "Invoice Date"])140 ocr_data: List of OCR blocks from pdf_utils.py output141 page_dimensions: Optional dict of {page_num: (width, height)}142 143 Returns:144 Dict with extracted form fields:145 {146 "Invoice Number": {"value": "INV-2024-001", "confidence": 0.95, ...},147 "Invoice Date": {"value": "2024-01-15", "confidence": 0.90, ...},148 ...149 }150 """151 logger.info(f"Extracting {len(form_fields)} form fields from {len(ocr_data)} OCR blocks")152 153 # Filter to LINE elements only (more reliable than WORD)154 line_blocks = [155 block for block in ocr_data 156 if block.get('blockType') == 'LINE'157 ]158 logger.info(f"Found {len(line_blocks)} LINE blocks")159 160 # Infer page dimensions if not provided161 if not page_dimensions:162 page_dimensions = self._infer_page_dimensions(ocr_data)163 164 # Extract each requested field165 results = {}166 for field_name in form_fields:167 extracted = self._extract_single_field(field_name, line_blocks, page_dimensions)168 if extracted:169 results[field_name] = extracted.to_dict()170 else:171 # Return empty value if not found172 results[field_name] = {173 "field_name": field_name,174 "value": "",175 "confidence": 0.0,176 "match_type": "not_found"177 }178 179 logger.info(f"Extracted {sum(1 for v in results.values() if v['value'])} of {len(form_fields)} fields")180 return results181 182 def _extract_single_field(183 self,184 field_name: str,185 line_blocks: List[Dict[str, Any]],186 page_dimensions: Dict[int, Tuple[int, int]]187 ) -> Optional[ExtractedField]:188 """189 Extract a single form field using multiple strategies.190 191 Strategies (in order of priority):192 1. Inline extraction with separator193 2. Pattern-based extraction (for known patterns)194 3. Proximity-based extraction195 """196 # Normalize field name for matching197 field_norm = self._normalize_text(field_name)198 199 # Get all variations to search for200 variations = self._get_field_variations(field_name)201 logger.debug(f"Searching for '{field_name}' with variations: {variations}")202 203 # Strategy 1: Inline extraction (key: value in same line)204 result = self._extract_inline(field_name, variations, line_blocks)205 if result:206 logger.debug(f"Found '{field_name}' via inline extraction: {result.value}")207 return result208 209 # Strategy 2: Pattern-based extraction210 result = self._extract_by_pattern(field_name, variations, line_blocks)211 if result:212 logger.debug(f"Found '{field_name}' via pattern extraction: {result.value}")213 return result214 215 # Strategy 3: Proximity-based extraction216 result = self._extract_by_proximity(field_name, variations, line_blocks, page_dimensions)217 if result:218 logger.debug(f"Found '{field_name}' via proximity extraction: {result.value}")219 return result220 221 logger.debug(f"Could not find value for '{field_name}'")222 return None223 224 def _extract_inline(225 self,226 field_name: str,227 variations: List[str],228 line_blocks: List[Dict[str, Any]]229 ) -> Optional[ExtractedField]:230 """231 Extract field where key and value are in the same line.232 E.g., "Invoice Number: INV-2024-001" or "Invoice Number - INV-2024-001"233 """234 for block in line_blocks:235 text = block.get('text', '')236 text_norm = self._normalize_text(text)237 238 for variation in variations:239 var_norm = self._normalize_text(variation)240 241 # Check if line contains the field name242 if var_norm not in text_norm:243 continue244 245 # Try to extract value after separator246 for sep in KEY_VALUE_SEPARATORS:247 # Build pattern: field_name + separator + value248 pattern = re.compile(249 rf'{re.escape(var_norm)}\s*{re.escape(sep)}\s*(.+)',250 re.IGNORECASE251 )252 match = pattern.search(text)253 if match:254 value = match.group(1).strip()255 # Clean up value (remove trailing separators)256 value = self._clean_value(value)257 if value:258 return ExtractedField(259 field_name=field_name,260 value=value,261 confidence=0.95,262 geometry=block.get('geometry'),263 page_num=block.get('pageNum', 1),264 match_type="inline"265 )266 267 # Try without separator (field name followed by value)268 # E.g., "Invoice Number INV-2024-001"269 pattern = re.compile(270 rf'{re.escape(var_norm)}\s+(\S+.*?)$',271 re.IGNORECASE272 )273 match = pattern.search(text)274 if match:275 value = match.group(1).strip()276 value = self._clean_value(value)277 # Only accept if value doesn't look like another field name278 if value and not self._looks_like_field_name(value):279 return ExtractedField(280 field_name=field_name,281 value=value,282 confidence=0.85,283 geometry=block.get('geometry'),284 page_num=block.get('pageNum', 1),285 match_type="inline_no_sep"286 )287 288 return None289 290 def _extract_by_pattern(291 self,292 field_name: str,293 variations: List[str],294 line_blocks: List[Dict[str, Any]]295 ) -> Optional[ExtractedField]:296 """297 Extract field using known value patterns.298 E.g., dates, amounts, reference numbers.299 """300 field_lower = field_name.lower()301 302 # Define patterns for specific field types303 patterns = {}304 305 # Date patterns306 if any(kw in field_lower for kw in ['date', 'dt']):307 patterns['date'] = [308 r'\d{1,2}[/-]\d{1,2}[/-]\d{2,4}',309 r'\d{4}[/-]\d{1,2}[/-]\d{1,2}',310 r'(?:Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec)[a-z]*[\s.,]+\d{1,2}[\s.,]+\d{2,4}',311 r'\d{1,2}[\s]+(?:Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec)[a-z]*[\s]+\d{2,4}',312 ]313 314 # Amount patterns315 if any(kw in field_lower for kw in ['amount', 'total', 'price', 'value', 'tax', 'subtotal']):316 patterns['amount'] = [317 r'[\$โฌยฃโน][\d,]+\.?\d*',318 r'[\d,]+\.?\d*\s*(?:USD|EUR|GBP|INR)',319 r'(?:Rs\.?|INR)\s*[\d,]+\.?\d*',320 ]321 322 # Reference number patterns323 if any(kw in field_lower for kw in ['number', 'no', '#', 'id', 'ref']):324 patterns['reference'] = [325 r'[A-Z]{2,4}[-/]?\d{4,}[-/]?\d*', # INV-2024-001, PO/12345326 r'\d{6,}', # Pure numeric references327 ]328 329 # GSTIN pattern330 if 'gstin' in field_lower or 'gst' in field_lower:331 patterns['gstin'] = [332 r'\d{2}[A-Z]{5}\d{4}[A-Z]{1}[A-Z\d]{1}[Z]{1}[A-Z\d]{1}', # Indian GSTIN333 ]334 335 if not patterns:336 return None337 338 # Search for field name first, then look for pattern nearby339 for block in line_blocks:340 text = block.get('text', '')341 text_norm = self._normalize_text(text)342 343 for variation in variations:344 var_norm = self._normalize_text(variation)345 346 if var_norm in text_norm:347 # Found field name, look for pattern in same line348 for pattern_type, pattern_list in patterns.items():349 for pattern in pattern_list:350 # Look for pattern after the field name351 idx = text_norm.find(var_norm)352 search_text = text[idx + len(variation):]353 354 match = re.search(pattern, search_text, re.IGNORECASE)355 if match:356 value = match.group(0).strip()357 return ExtractedField(358 field_name=field_name,359 value=value,360 confidence=0.90,361 geometry=block.get('geometry'),362 page_num=block.get('pageNum', 1),363 match_type=f"pattern_{pattern_type}"364 )365 366 return None367 368 def _extract_by_proximity(369 self,370 field_name: str,371 variations: List[str],372 line_blocks: List[Dict[str, Any]],373 page_dimensions: Dict[int, Tuple[int, int]]374 ) -> Optional[ExtractedField]:375 """376 Extract field where key and value are in separate but adjacent elements.377 """378 # Find blocks containing the field name379 label_blocks = []380 for block in line_blocks:381 text = block.get('text', '')382 text_norm = self._normalize_text(text)383 384 for variation in variations:385 var_norm = self._normalize_text(variation)386 if var_norm in text_norm or text_norm in var_norm:387 # Check if this looks like a label (ends with separator or is the field name)388 if self._looks_like_label(text, variation):389 label_blocks.append((block, variation))390 break391 392 if not label_blocks:393 return None394 395 # For each label, find the best value candidate396 for label_block, matched_variation in label_blocks:397 label_geom = label_block.get('geometry', {})398 label_page = label_block.get('pageNum', 1)399 400 page_dims = page_dimensions.get(label_page, (2479, 3508))401 page_width, page_height = page_dims402 403 max_h_gap = page_width * self.max_horizontal_gap404 max_v_gap = page_height * self.max_vertical_gap405 406 best_value = None407 best_score = 0408 409 for block in line_blocks:410 if block.get('pageNum', 1) != label_page:411 continue412 if block.get('text', '').strip() == label_block.get('text', '').strip():413 continue414 415 value_geom = block.get('geometry', {})416 value_text = block.get('text', '').strip()417 418 # Skip if looks like another label419 if self._looks_like_field_name(value_text):420 continue421 422 # Calculate position score423 score = self._calculate_proximity_score(424 label_geom, value_geom, max_h_gap, max_v_gap425 )426 427 if score > best_score:428 best_score = score429 best_value = block430 431 if best_value and best_score > 0.3:432 value_text = best_value.get('text', '').strip()433 value_text = self._clean_value(value_text)434 435 return ExtractedField(436 field_name=field_name,437 value=value_text,438 confidence=round(best_score, 2),439 geometry=best_value.get('geometry'),440 page_num=best_value.get('pageNum', 1),441 match_type="proximity"442 )443 444 return None445 446 def _calculate_proximity_score(447 self,448 label_geom: Dict[str, int],449 value_geom: Dict[str, int],450 max_h_gap: float,451 max_v_gap: float452 ) -> float:453 """Calculate how likely a value is associated with a label based on position."""454 if not label_geom or not value_geom:455 return 0.0456 457 label_x1 = label_geom.get('x1', 0)458 label_x2 = label_geom.get('x2', 0)459 label_y1 = label_geom.get('y1', 0)460 label_y2 = label_geom.get('y2', 0)461 462 value_x1 = value_geom.get('x1', 0)463 value_x2 = value_geom.get('x2', 0)464 value_y1 = value_geom.get('y1', 0)465 value_y2 = value_geom.get('y2', 0)466 467 score = 0.0468 469 # Check if value is to the right of label (same row)470 label_center_y = (label_y1 + label_y2) / 2471 value_center_y = (value_y1 + value_y2) / 2472 473 vertical_overlap = (min(label_y2, value_y2) - max(label_y1, value_y1))474 min_height = min(label_y2 - label_y1, value_y2 - value_y1)475 476 if min_height > 0 and vertical_overlap / min_height > 0.5:477 # Same row - value should be to the right478 if value_x1 >= label_x2:479 h_gap = value_x1 - label_x2480 if h_gap < max_h_gap:481 proximity = 1.0 - (h_gap / max_h_gap)482 score = max(score, 0.5 + 0.4 * proximity)483 484 # Check if value is below label485 if value_y1 >= label_y2:486 v_gap = value_y1 - label_y2487 if v_gap < max_v_gap:488 # Check horizontal alignment (value should be near label x position)489 horizontal_overlap = (min(label_x2, value_x2) - max(label_x1, value_x1))490 if horizontal_overlap > 0 or abs(value_x1 - label_x1) < max_h_gap:491 proximity = 1.0 - (v_gap / max_v_gap)492 score = max(score, 0.4 + 0.3 * proximity)493 494 return score495 496 def _get_field_variations(self, field_name: str) -> List[str]:497 """Get all variations of a field name to search for."""498 field_norm = self._normalize_text(field_name)499 variations = [field_name]500 501 # Check if this maps to a canonical name502 canonical = self._variation_map.get(field_norm)503 if canonical:504 # Add all known variations505 variations.extend(FIELD_VARIATIONS.get(canonical, []))506 507 # Add common transformations508 variations.append(field_name.replace(' ', '')) # NoSpaces509 variations.append(field_name.replace(' ', '_')) # With_underscores510 511 # Deduplicate while preserving order512 seen = set()513 unique = []514 for v in variations:515 v_norm = self._normalize_text(v)516 if v_norm not in seen:517 seen.add(v_norm)518 unique.append(v)519 520 return unique521 522 def _normalize_text(self, text: str) -> str:523 """Normalize text for comparison."""524 if not text:525 return ""526 result = ' '.join(text.lower().strip().split())527 return result528 529 def _clean_value(self, value: str) -> str:530 """Clean extracted value."""531 if not value:532 return ""533 534 # Remove leading/trailing separators535 value = value.strip()536 for sep in KEY_VALUE_SEPARATORS + [',', ';']:537 value = value.strip(sep).strip()538 539 return value540 541 def _looks_like_label(self, text: str, field_name: str) -> bool:542 """Check if text looks like a field label."""543 text = text.strip()544 545 # Ends with separator546 for sep in KEY_VALUE_SEPARATORS:547 if text.endswith(sep):548 return True549 550 # Is exactly the field name (possibly with separator)551 text_clean = text.rstrip(':- ')552 if self._normalize_text(text_clean) == self._normalize_text(field_name):553 return True554 555 return True # Default to treating as label if field name matches556 557 def _looks_like_field_name(self, text: str) -> bool:558 """Check if text looks like a field label (not a value)."""559 text = text.strip()560 561 # Very short text might be a value562 if len(text) < 3:563 return False564 565 # Contains separator at end566 for sep in KEY_VALUE_SEPARATORS:567 if text.endswith(sep):568 return True569 570 # Is a known field variation571 text_norm = self._normalize_text(text.rstrip(':- '))572 if text_norm in self._variation_map:573 return True574 575 return False576 577 def _infer_page_dimensions(self, ocr_data: List[Dict[str, Any]]) -> Dict[int, Tuple[int, int]]:578 """Infer page dimensions from OCR data."""579 page_dims = {}580 581 for block in ocr_data:582 page_num = block.get('pageNum', 1)583 geom = block.get('geometry', {})584 585 if page_num not in page_dims:586 page_dims[page_num] = (0, 0)587 588 current_w, current_h = page_dims[page_num]589 page_dims[page_num] = (590 max(current_w, geom.get('x2', 0)),591 max(current_h, geom.get('y2', 0))592 )593 594 # Add buffer for page dimensions595 for page_num in page_dims:596 w, h = page_dims[page_num]597 page_dims[page_num] = (int(w * 1.1), int(h * 1.1))598 599 return page_dims if page_dims else {1: (2479, 3508)}600 601 602# =============================================================================603# Standalone Functions604# =============================================================================605 606def extract_form_fields(607 form_fields: List[str],608 ocr_data: List[Dict[str, Any]],609 page_dimensions: Optional[Dict[int, Tuple[int, int]]] = None610) -> Dict[str, Any]:611 """612 Extract and return key-value pairs for requested form fields.613 614 This is a convenience function that creates a FormFieldExtractor instance615 and calls its extract_form_fields method.616 617 Args:618 form_fields: List of field names to extract (e.g., ["Invoice Number", "Invoice Date"])619 ocr_data: List of OCR blocks from pdf_utils.py output (ocrBlocks array)620 page_dimensions: Optional dict of {page_num: (width, height)}621 622 Returns:623 Dict with extracted form fields624 """625 extractor = FormFieldExtractor()626 return extractor.extract_form_fields(form_fields, ocr_data, page_dimensions)627 628 629def merge_extraction_results(630 pdf_output: Dict[str, Any],631 form_fields_result: Dict[str, Any],632 table_fields_result: Optional[Dict[str, Any]] = None633) -> Dict[str, Any]:634 """635 Merge form field and table extraction results with PDF output.636 637 Args:638 pdf_output: Original output from pdf_utils.py639 form_fields_result: Result from extract_form_fields()640 table_fields_result: Optional result from table extraction641 642 Returns:643 Merged output with structure:644 {645 "version": "1.0",646 "metadata": {...},647 "pages": [...],648 "ocrBlocks": {649 "formFields": {...},650 "tableFields": {...}651 }652 }653 """654 result = {655 "version": pdf_output.get("version", "1.0"),656 "metadata": pdf_output.get("metadata", {}),657 "pages": pdf_output.get("pages", []),658 "ocrBlocks": {659 "formFields": form_fields_result,660 "tableFields": table_fields_result or {}661 }662 }663 664 return result665 666 667# =============================================================================668# CLI Entry Point669# =============================================================================670 671if __name__ == '__main__':672 import sys673 import json674 675 # Example usage676 print("Form Field Extractor")677 print("=" * 50)678 679 if len(sys.argv) < 2:680 print("Usage: python form_field_extractor.py <json_file> [field1,field2,...]")681 print("\nExample:")682 print(" python form_field_extractor.py output.json \"Invoice Number,Invoice Date,Total Amount\"")683 sys.exit(1)684 685 json_file = sys.argv[1]686 687 # Default fields to extract688 default_fields = ["Invoice Number", "Invoice Date", "Total Amount", "GSTIN", "PO Number"]689 690 if len(sys.argv) > 2:691 fields_str = sys.argv[2]692 form_fields = [f.strip() for f in fields_str.split(',')]693 else:694 form_fields = default_fields695 696 print(f"Input file: {json_file}")697 print(f"Fields to extract: {form_fields}")698 print("-" * 50)699 700 # Load JSON701 with open(json_file, 'r') as f:702 data = json.load(f)703 704 # Extract form fields705 ocr_blocks = data.get('ocrBlocks', [])706 707 # Get page dimensions708 page_dimensions = {}709 for page in data.get('pages', []):710 page_num = page.get('pageNum', len(page_dimensions) + 1)711 page_dimensions[page_num] = (page.get('width', 2479), page.get('height', 3508))712 713 # Run extraction714 extractor = FormFieldExtractor()715 results = extractor.extract_form_fields(form_fields, ocr_blocks, page_dimensions)716 717 # Print results718 print("\nExtracted Fields:")719 print("-" * 50)720 for field_name, field_data in results.items():721 value = field_data.get('value', '')722 match_type = field_data.get('match_type', 'unknown')723 confidence = field_data.get('confidence', 0)724 725 status = "โ" if value else "โ"726 print(f"{status} {field_name}: {value or '(not found)'}")727 if value:728 print(f" Match: {match_type}, Confidence: {confidence}")729 730 print("\n" + "-" * 50)731 print("JSON Output:")732 print(json.dumps(results, indent=2))733 