Team Ai
Apppublic

documentExtractionag051/ExtractDocument

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
invoice_post_processing.py625 linesDownload Raw Back to utils
1# custom > invoice_post_processing.py > extract_invoice_tables2import logging3import re4from typing import List, Dict, Any, Optional5from .geometry_utils import compute_global_bounds, Rect, TextRect6 7# Set up logging - use INFO level by default8logging.basicConfig(level=logging.INFO, format='%(levelname)s:%(name)s:%(message)s')9logger = logging.getLogger(__name__)10 11 12def extract_invoice_tables(13    ocr_elems,14    validation_data,15    field_name='custom_field_name',16    anchor_keywords=None,17    column_headers=None,18    anchor_column=None,19    anchor_pattern=None,20    advanced_options=None,21    page_dimensions=None,22    **kwargs23):24    """25    Specialized post-processing pipeline for Invoice documents.26    27    Fixes applied:28    - Multi-line cell text is now ordered top-to-bottom, then left-to-right.29    - Rows with invalid Material codes are removed based on multiple regex opt.30    31    Args:32        column_headers: Dict mapping field names to header text in document33        anchor_column: Column name used for row detection34        anchor_pattern: Regex pattern string for validating anchor column values35        advanced_options: Dict for special logic options:36            - stop_text: Text marking end of table on last page37            - enable_continuation: Whether to enable continuation column38            - continuation_marker: Text after which continuation starts39    """40    logger.info("Running invoice specialized post-processing...")41    42    # Parse advanced options43    if advanced_options is None:44        advanced_options = {}45    46    stop_text = advanced_options.get("stop_text", "Total Amount:")47    enable_continuation = advanced_options.get("enable_continuation", False)48    continuation_column = advanced_options.get("continuation_column", "Description")49    continuation_marker = advanced_options.get("continuation_marker", "Ref.:")50    51    # Add table end markers for footer detection52    table_end_markers = advanced_options.get("table_end_markers", [])53    54    # Use provided column headers or default55    COLUMN_HEADERS = column_headers if column_headers else {56        "Material Code": "Material Code",57        "Description": "Description",58        "Quantity": "Qty.",59        "unit_price": "Unit Price",60        "total_price": "Total Net Value"61    }62    63    # Use provided anchor column or default64    ANCHOR_COLUMN = anchor_column if anchor_column else "Material Code"65    66    # Use provided pattern or default67    if anchor_pattern:68        MATERIAL_CODE_PATTERNS = [re.compile(anchor_pattern)]69    else:70        MATERIAL_CODE_PATTERNS = [re.compile(r"^\d{8}$")]71    72    # Calculate cumulative heights for multi-page documents73    cumulative_heights = _calculate_cumulative_heights(page_dimensions)74    75    # Step 1: Create column spaces based on header positions76    column_spaces = _create_column_spaces(ocr_elems, page_dimensions, COLUMN_HEADERS)77    78    # Step 2: Find table boundaries (footer detection)79    table_boundaries = _find_table_boundaries(ocr_elems, page_dimensions, table_end_markers)80    81    # Step 3: Assign OCR lines to column spaces82    space_to_lines = _assign_lines_to_spaces(ocr_elems, column_spaces)83    84    # Step 4: Get lines from anchor column and identify rows85    anchor_col_lines = space_to_lines.get(ANCHOR_COLUMN, [])86    rows = _identify_rows(87        anchor_col_lines, 88        page_dimensions, 89        MATERIAL_CODE_PATTERNS[0],90        table_boundaries91    )92    93    logger.info(f"Identified {len(rows)} rows based on anchor column '{ANCHOR_COLUMN}'")94    95    # Step 5: Extract table data from rows96    table_rows = _extract_table_rows(97        rows=rows,98        column_spaces=column_spaces,99        space_to_lines=space_to_lines,100        cumulative_heights=cumulative_heights,101        MATERIAL_CODE_PATTERNS=MATERIAL_CODE_PATTERNS,102        ocr_elems=ocr_elems,103        anchor_column=ANCHOR_COLUMN,104        stop_text=stop_text,105        enable_continuation=enable_continuation,106        continuation_column=continuation_column,107        continuation_marker=continuation_marker,108        table_boundaries=table_boundaries109    )110    111    logger.info(f"Extracted {len(table_rows)} table rows")112    113    # Step 6: Update validation_data with table results114    validation_data["tables"] = {"table": table_rows}115    116    return validation_data117 118 119def _calculate_cumulative_heights(page_dimensions):120    """Calculate cumulative page heights for multi-page documents."""121    cumulative_heights = {}122    cumulative_height = 0123    124    for page_num in sorted(page_dimensions.keys()):125        cumulative_heights[page_num] = cumulative_height126        height = page_dimensions[page_num][1]127        cumulative_height += height128    129    return cumulative_heights130 131 132def _create_column_spaces(ocr_elems, page_dimensions, column_headers):133    """134    Create vertical column spaces based on header text positions.135    136    Uses a multi-pass matching strategy:137    1. Exact match (case-sensitive)138    2. Case-insensitive match139    3. Partial/contains match (header text in element or element in header)140    4. Normalized match (strip whitespace, lowercase)141    142    Args:143        ocr_elems: List of OCR TextRect elements144        page_dimensions: Dict of {page_num: (width, height)}145        column_headers: Dict mapping column keys to header text patterns146    147    Returns:148        Dict mapping column keys to Rect objects representing column spaces149    """150    column_spaces = {key: None for key in column_headers.keys()}151    max_page_height = max(height for width, height in page_dimensions.values()) if page_dimensions else 0152    153    # Get all LINE elements from page 1 for header matching154    page1_lines = [155        elem for elem in ocr_elems156        if elem.page_num == 1 and elem.block_type == "LINE"157    ]158    159    logger.info("=== Column Header Detection ===")160    logger.info(f"Looking for headers: {column_headers}")161    logger.info(f"Found {len(page1_lines)} LINE elements on page 1")162    163    # Log all page 1 LINE elements for debugging164    logger.debug("Page 1 LINE elements:")165    for elem in page1_lines:166        logger.debug(f"  '{elem.text.strip()}' at x={elem.x1}-{elem.x2}, y={elem.y1}-{elem.y2}")167    168    def _normalize_text(text):169        """Normalize text for comparison: lowercase, strip, collapse whitespace."""170        return ' '.join(text.lower().strip().split())171    172    def _match_header(elem_text, header_text):173        """174        Multi-strategy header matching.175        176        Returns: (match_type, confidence) where confidence is 0-100177        """178        elem_clean = elem_text.strip()179        header_clean = header_text.strip()180        181        # Pass 1: Exact match182        if elem_clean == header_clean:183            return ("exact", 100)184        185        # Pass 2: Case-insensitive exact match186        if elem_clean.lower() == header_clean.lower():187            return ("case_insensitive", 95)188        189        # Pass 3: Normalized match (strip + lowercase + collapse whitespace)190        elem_norm = _normalize_text(elem_clean)191        header_norm = _normalize_text(header_clean)192        if elem_norm == header_norm:193            return ("normalized", 90)194        195        # Pass 4: Contains match (element contains header or vice versa)196        if header_norm in elem_norm:197            return ("contains_header", 80)198        if elem_norm in header_norm:199            return ("contains_elem", 75)200        201        # Pass 5: Starts-with match202        if elem_norm.startswith(header_norm):203            return ("starts_with", 70)204        if header_norm.startswith(elem_norm):205            return ("header_starts", 65)206        207        return (None, 0)208    209    # Match each header using multi-pass strategy210    matched_headers = {}211    212    for col_key, header_text in column_headers.items():213        best_match = None214        best_confidence = 0215        best_elem = None216        217        for elem in page1_lines:218            match_type, confidence = _match_header(elem.text, header_text)219            if confidence > best_confidence:220                best_confidence = confidence221                best_match = match_type222                best_elem = elem223        224        if best_match and best_elem:225            column_spaces[col_key] = Rect(226                x1=best_elem.x1,227                y1=0,228                x2=best_elem.x2,229                y2=max_page_height230            )231            matched_headers[col_key] = {232                "header_text": header_text,233                "matched_text": best_elem.text.strip(),234                "match_type": best_match,235                "confidence": best_confidence,236                "x_range": f"{best_elem.x1}-{best_elem.x2}"237            }238            logger.info(239                f"✓ Matched '{col_key}' -> '{best_elem.text.strip()}' "240                f"(type={best_match}, conf={best_confidence}%, x={best_elem.x1}-{best_elem.x2})"241            )242        else:243            logger.warning(f"✗ No match found for column '{col_key}' with header '{header_text}'")244    245    # Log summary246    found_count = sum(1 for space in column_spaces.values() if space is not None)247    missing_cols = [key for key, space in column_spaces.items() if space is None]248    249    logger.info(f"=== Header Detection Summary ===")250    logger.info(f"Found {found_count}/{len(column_headers)} columns")251    252    if missing_cols:253        logger.warning(f"Missing columns: {', '.join(missing_cols)}")254        logger.info("Available LINE texts on page 1 (for debugging):")255        for elem in page1_lines:256            logger.info(f"  - '{elem.text.strip()}'")257    258    return column_spaces259 260 261def _find_table_boundaries(ocr_elems, page_dimensions, table_end_markers):262    """263    Find the table content boundaries for each page.264    Returns dict: {page_num: {'top': y1, 'bottom': y2}}265    266    The bottom boundary is determined by finding footer markers.267    """268    boundaries = {}269    270    for page_num, (page_width, page_height) in page_dimensions.items():271        # Default: full page272        boundaries[page_num] = {273            'top': 0,274            'bottom': page_height275        }276        277        if not table_end_markers:278            continue279        280        # Find the earliest footer marker on this page281        footer_top = page_height282        for elem in ocr_elems:283            if elem.page_num == page_num and elem.block_type == "LINE":284                for marker in table_end_markers:285                    if marker in elem.text:286                        # Use the top of the footer line as the table bottom287                        if elem.y1 < footer_top:288                            footer_top = elem.y1 - 10  # Small buffer above footer289                            logger.info(f"Page {page_num}: Found footer marker '{marker}' "290                                      f"at y={elem.y1}, setting table bottom to {footer_top}")291                        break292        293        boundaries[page_num]['bottom'] = footer_top294    295    return boundaries296 297 298def _assign_lines_to_spaces(ocr_elems, column_spaces):299    """Assign each LINE element to its best-matching column space."""300    space_to_lines = {key: [] for key in column_spaces.keys()}301    302    for line_elem in ocr_elems:303        if line_elem.block_type != "LINE":304            continue305        306        best_space, best_intersection = None, 0307        308        for col_key, space_rect in column_spaces.items():309            if space_rect is None:310                continue311            intersection = space_rect.intersection_pct(line_elem)312            if intersection > best_intersection:313                best_intersection, best_space = intersection, col_key314        315        if best_space and best_intersection > 0:316            space_to_lines[best_space].append(line_elem)317    318    return space_to_lines319 320 321def _identify_rows(material_code_lines, page_dimensions, pattern, table_boundaries=None):322    """323    Identify row boundaries based on material code positions.324    325    Args:326        table_boundaries: Dict with table content boundaries per page327    """328    # Log all lines being tested329    logger.info(f"Testing {len(material_code_lines)} lines against anchor pattern: {pattern.pattern}")330    for line in material_code_lines:331        text = line.text.strip()332        match = pattern.match(text)333        logger.debug(f"  Line '{text}' -> {'MATCH' if match else 'no match'}")334    335    valid_material_lines = [336        line for line in material_code_lines if pattern.match(line.text.strip())337    ]338    339    logger.info(f"Found {len(valid_material_lines)} valid anchor column values")340    if valid_material_lines:341        logger.info(f"Valid material lines: {[line.text for line in valid_material_lines]}")342    else:343        # Show what values ARE in the anchor column for debugging344        logger.warning(f"No lines matched pattern '{pattern.pattern}'")345        logger.warning(f"Available anchor column values:")346        for line in material_code_lines[:20]:  # Show first 20347            logger.info(f"  - '{line.text.strip()}' (page {line.page_num}, y={line.y1})")348    349    page_to_lines, rows = {}, []350    for line in valid_material_lines:351        page_to_lines.setdefault(line.page_num, []).append(line)352    353    for page_num, lines in sorted(page_to_lines.items()):354        lines_sorted = sorted(lines, key=lambda l: l.y1)355        page_width, page_height = page_dimensions[page_num]356        357        # Get table bottom boundary for this page (excludes footer)358        if table_boundaries and page_num in table_boundaries:359            table_bottom = table_boundaries[page_num]['bottom']360        else:361            table_bottom = page_height362        363        for i, line in enumerate(lines_sorted):364            y1 = line.y1365            if i + 1 < len(lines_sorted):366                y2 = lines_sorted[i + 1].y1367            else:368                # Last row on this page - use table bottom (not page height)369                y2 = table_bottom370            371            row_rect = Rect(x1=0, y1=y1, x2=page_width, y2=y2, block_type="ROW", confidence=1.0)372            logger.info(f"Created row rect: page={page_num}, x1={row_rect.x1}, y1={row_rect.y1}")373            rows.append({"page_num": page_num, "row_rect": row_rect})374    375    return rows376 377 378def _extract_table_rows(379    rows,380    column_spaces,381    space_to_lines,382    cumulative_heights,383    MATERIAL_CODE_PATTERNS,384    ocr_elems,385    anchor_column="Material Code",386    stop_text=None,387    enable_continuation=False,388    continuation_column=None,389    continuation_marker=None,390    table_boundaries=None391):392    """393    Extract cell values for each table row by finding line elements at row-column intersections.394    395    Args:396        stop_text: Text that marks end of table on last page397        enable_continuation: Whether to look for continuation on next page398        continuation_column: Which column may continue to next page399        continuation_marker: Text after which continuation starts400        table_boundaries: Dict with table content boundaries per page (to filter out footer)401    """402    # print(f"Rows to process: {rows}")403    # print(f"column_spaces: {column_spaces}")404    # print(f"space_to_lines: {list(space_to_lines)}")405    # print(f"cumulative_heights: {cumulative_heights}")406    # print(f"material_code_patterns: {MATERIAL_CODE_PATTERNS}")407    #print(f"ocr_elems: {ocr_elems}")408    table_rows = []409    410    if not rows:411        return table_rows412    413    highest_page_row_number = max(row["page_num"] for row in rows)414    logger.info(f"Highest page row number: {highest_page_row_number}")415    416    # Sort rows by page_num and y1 for easier next-row lookup417    sorted_rows = sorted(rows, key=lambda r: (r["page_num"], r["row_rect"].y1))418    419    for row_idx, row in enumerate(sorted_rows):420        page_num = row["page_num"]421        row_rect = row["row_rect"]422        height_offset = cumulative_heights.get(page_num, 0)423        row_data = {}424        425        for col_key, col_rect in column_spaces.items():426            print(f"Processing column: {col_key}")427            if col_rect is None:428                continue429            430            # ---------------------------431            # GET TABLE BOUNDARY FOR THIS PAGE432            # ---------------------------433            page_table_bottom = None434            if table_boundaries and page_num in table_boundaries:435                page_table_bottom = table_boundaries[page_num]['bottom']436            437            # ---------------------------438            # CONFIGURABLE LAST-PAGE LOGIC439            # ---------------------------440            effective_col_rect_y2 = col_rect.y2441            if page_num == highest_page_row_number and stop_text:442                # Look for stop text to limit column boundary443                anchor_lines = sorted(444                    space_to_lines.get(anchor_column, []),445                    key=lambda l: l.y1446                )447                for line in anchor_lines:448                    if stop_text in line.text:449                        if line.page_num == page_num:450                            effective_col_rect_y2 = line.y1 - 10451                            break452            453            # ---------------------------454            # FIND BEST MATCHING LINES (filtered by table boundary)455            # ---------------------------456            candidate_lines = []457            for l in space_to_lines.get(col_key, []):458                if l.page_num != page_num:459                    continue460                if col_rect.intersection_pct(l) <= 0:461                    continue462                # Filter out lines below table boundary (footer lines)463                if page_table_bottom and l.y1 >= page_table_bottom:464                    continue465                candidate_lines.append(l)466            467            # Score each line by its intersection with the row468            scored_lines = []469            for line in candidate_lines:470                row_intersection = row_rect.intersection_pct(line)471                if row_intersection > 0:472                    scored_lines.append((line, row_intersection))473            474            # Filter to lines where this row has the best claim475            cell_lines = []476            for line, score in scored_lines:477                best_row_for_line = True478                for other_row in sorted_rows:479                    if other_row is row:480                        continue481                    if other_row["page_num"] != page_num:482                        continue483                    other_intersection = other_row["row_rect"].intersection_pct(line)484                    if other_intersection > score:485                        best_row_for_line = False486                        break487                488                if best_row_for_line:489                    cell_lines.append(line)490            491            # ---------------------------492            # CONFIGURABLE CONTINUATION LOGIC493            # ---------------------------494            continuation_lines = []495            if enable_continuation and col_key == continuation_column:496                # Check if this is the last row on current page497                is_last_row_on_page = True498                for future_row in sorted_rows[row_idx + 1:]:499                    if future_row["page_num"] == page_num:500                        is_last_row_on_page = False501                        brea_y1k502                503                if is_last_row_on_page:504                    next_page_num = page_num + 1505                    506                    # Get table boundaries for next page507                    next_page_table_top = 0508                    next_page_table_bottom = None509                    if table_boundaries and next_page_num in table_boundaries:510                        next_page_table_top = table_boundaries[next_page_num].get('top', 0)511                        next_page_table_bottom = table_boundaries[next_page_num].get('bottom')512                    513                    # Find the first row on the next page (if any)514                    next_row = None515                    for future_row in sorted_rows[row_idx + 1:]:516                        if future_row["page_num"] == next_page_num:517                            next_row_y1 = future_row["row_rect"].y1518                            break519                    520                    # Find the continuation marker line on the next page (optional)521                    marker_line_y2 = 0522                    if continuation_marker:523                        for elem in ocr_elems:524                            if elem.page_num == next_page_num and elem.block_type == "LINE":525                                if continuation_marker in elem.text:526                                    marker_line_y2 = elem.y2527                                    break528                    529                    # If there's a next row on next page, get content above it530                    if next_row_y1 is not None:531                        # Get lines from next page that are above the next row532                        for l in space_to_lines.get(col_key, []):533                            if l.page_num != next_page_num:534                                continue535                            if col_rect.intersection_pct(l) <= 0:536                                continue537                            # Must be above the next row's anchor538                            if l.y1 >= next_row_y1:539                                continue540                            # Must be within table area (not in footer)541                            if next_page_table_bottom and l.y1 >= next_page_table_bottom:542                                continue543                            # Must be after continuation marker (if specified)544                            if continuation_marker and marker_line_y2 and l.y1 <= marker_line_y2:545                                continue546                            547                            continuation_lines.append(l)548                    549                    if continuation_lines:550                        logger.info(551                            f"Found {len(continuation_lines)} continuation lines for {col_key} "552                            f"from page {page_num} to page {next_page_num} (above next row at y={next_row_y1})"553                        )554                else:555                    # No more rows on next page - check if content continues there556                    for l in space_to_lines.get(col_key, []):557                        if l.page_num != next_page_num:558                            continue559                        if col_rect.intersection_pct(l) <= 0:560                            continue561                        # Must be within table area (not in footer)562                        if next_page_table_bottom and l.y1 >= next_page_table_bottom:563                            continue564                        # Must be after continuation marker (if specified)565                        if continuation_marker and marker_line_y2 and l.y1 <= marker_line_y2:566                            continue567                        568                        continuation_lines.append(l)569                    570                    if continuation_lines:571                        logger.info(572                            f"Found {len(continuation_lines)} continuation lines for {col_key} "573                            f"from page {page_num} to page {next_page_num} (no more rows on next page)"574                        )575            576            # ---------------------------577            # SORT + MERGE TEXT578            # ---------------------------579            # Combine cell_lines and continuation_lines for text extraction580            all_lines_for_text = cell_lines + continuation_lines581            all_lines_sorted = sorted(582                all_lines_for_text,583                key=lambda l: (l.page_num, round(l.y1, 1), l.x1)584            )585            586            cell_text = " ".join(l.text.strip() for l in all_lines_sorted)587            588            # Determine bounding box (only from original page cell_lines)589            cell_lines_sorted = sorted(590                cell_lines,591                key=lambda l: (round(l.y1, 1), l.x1)592            )593            594            if cell_lines_sorted:595                x1 = min(l.x1 for l in cell_lines_sorted)596                y1 = min(l.y1 for l in cell_lines_sorted) + height_offset597                x2 = max(l.x2 for l in cell_lines_sorted)598                y2 = max(l.y2 for l in cell_lines_sorted) + height_offset599                bounds = f"{x1}, {y1}, {x2 - x1}, {y2 - y1}"600            else:601                bounds = "0, 0, -1, -1"602            603            row_data[col_key] = {604                "value": cell_text,605                "bounds": bounds606            }607        608        # ---------------------------609        # ANCHOR COLUMN VALIDATION610        # ---------------------------611        anchor_value = row_data.get(anchor_column, {}).get("value", "").strip()612        logger.info(f"Validating anchor column '{anchor_column}' value: '{anchor_value}'")613        614        # if not any(p.match(anchor_value) for p in MATERIAL_CODE_PATTERNS):615        #     logger.info(f"Skipping invalid {anchor_column} row: '{anchor_value}'")616        #     continue617        618        table_rows.append(row_data)619    620    logger.info(f"Final table row count after filtering: {len(table_rows)}")621    # for row in table_rows:622    #     print(row['description'])623    624    return table_rows625