From c84decca8d7f213cdec6952d22c3a2e9120a4cc6 Mon Sep 17 00:00:00 2001 From: Alex Galarce Date: Fri, 20 Jun 2025 19:46:23 +0000 Subject: [PATCH] Merged in feature/rework-table-handling (pull request #575) Feature/rework table handling * test: add cases for configuration boundary * test: add table post-processing tests (cases to ensure preservation of post-table text and metadata during table operations) * test: enhance table handling tests for small and large tables, ensuring proper suffixing and splitting behavior * Fix and clarify intent in comments * fix: enhance handling of post-table text when combining tables * refactor: remove commented-out code in clean_tables function for clarity * test: update post-table text preservation assertions and add debug tracing * refactor: remove unused remove_table_from_page function and its test * refactor: enhance combine method documentation and clarify split_by_size_with_smart_tables logic * Merge remote-tracking branch 'origin/main' into feature/rework-table-handling * refactor: reduce large table threshold from 1000 to 50 rows for better table splitting * Add simple splitting * refactor: uncomment dynamic answer retrieval in run_one_to_n_prompts for improved functionality * Merge remote-tracking branch 'origin/main' into feature/rework-table-handling * refactor: remove unused table handling functions and related constants for cleaner code * refactor: remove unused test functions in preparation for new tests * refactor: streamline continuation table handling by utilizing existing flags * refactor: update table combination logic to include all tables from the main page * refactor: enhance table handling to automatically resolve column mismatches, respect row limits, and improve unit/integration test coverage * rename clean_tables_simple to clean_tables * Update docstrings * refactor: rename table handling functions * Fix indentation and add tests * isort, black * Merge remote-tracking branch 'origin/main' into feature/rework-table-handling Approved-by: Katon Minhas --- fieldExtraction/src/config.py | 3 +- .../src/investment/file_processing.py | 4 +- fieldExtraction/src/investment/preprocess.py | 86 +- fieldExtraction/src/investment/table_funcs.py | 721 ++++++++++---- .../tests/preprocessing_integration_test.py | 130 +++ fieldExtraction/tests/table_funcs_test.py | 878 +++++++++++++++++- 6 files changed, 1579 insertions(+), 243 deletions(-) create mode 100644 fieldExtraction/tests/preprocessing_integration_test.py diff --git a/fieldExtraction/src/config.py b/fieldExtraction/src/config.py index 7cdf0aa..ea22be3 100644 --- a/fieldExtraction/src/config.py +++ b/fieldExtraction/src/config.py @@ -253,8 +253,7 @@ TABLE_THRESHOLD = 2 COMPLEX_OUTPUT_PATH = "complex_flag_output" COMPLEX_OUTPUT_FILENAME = f"{TODAY}_complex_flag_test.csv" -TABLE_ROW_LIMIT = get_arg_value("table_row_limit", 10) - +TABLE_ROW_LIMIT = 10 ######################################## EC2 ######################################## diff --git a/fieldExtraction/src/investment/file_processing.py b/fieldExtraction/src/investment/file_processing.py index 748afbf..cce8c13 100644 --- a/fieldExtraction/src/investment/file_processing.py +++ b/fieldExtraction/src/investment/file_processing.py @@ -137,7 +137,7 @@ def run_one_to_n_prompts(filename, exhibit_dict, all_exhibit_headers, all_datase reimbursement_level_fields = FieldSet(relationship="one_to_n", field_type="reimbursement_level", file_path=FIELD_JSON_PATH) ################## GET REIMBURSEMENT TIN/NPI ################## - tin_npi_answers, reimbursement_level_fields = reimbursement_tin_npi(exhibit_text, reimbursement_level_fields, filename) + # tin_npi_answers, reimbursement_level_fields = reimbursement_tin_npi(exhibit_text, reimbursement_level_fields, filename) ################## GET EXHIBIT-LEVEL ANSWERS ################## exhibit_level_answers = one_to_n_funcs.get_exhibit_level_answers(exhibit_text, filename) @@ -172,7 +172,7 @@ def run_one_to_n_prompts(filename, exhibit_dict, all_exhibit_headers, all_datase reimbursement_level_answers = reimbursement_level(exhibit_text, filename, reimbursement_level_fields, all_dataset, exhibit_page, seen_pairs) # Return list of dictionaries ################# COMBINE ANSWERS ################## - full_answer_dict = combine_one_to_n_answers(exhibit_level_answers, reimbursement_level_answers, tin_npi_answers) + full_answer_dict = combine_one_to_n_answers(exhibit_level_answers, reimbursement_level_answers, tin_npi_answers={}) one_to_n_results += full_answer_dict ################## Crosswalk Fields ################## diff --git a/fieldExtraction/src/investment/preprocess.py b/fieldExtraction/src/investment/preprocess.py index bea2361..7ae3b50 100644 --- a/fieldExtraction/src/investment/preprocess.py +++ b/fieldExtraction/src/investment/preprocess.py @@ -1,12 +1,13 @@ -import src.investment.table_funcs as table_funcs -from src import config, keywords, preprocessing_funcs -import src.investment.smart_chunking_funcs as smart_chunking_funcs -from src import keywords, preprocessing_funcs -import src.prompts.investment_prompts as investment_prompts -from src.utils import llm_utils, string_utils -from src.enums.delimiters import Delimiter import logging +import src.investment.smart_chunking_funcs as smart_chunking_funcs +import src.investment.table_funcs as table_funcs +import src.prompts.investment_prompts as investment_prompts +from src import config, keywords, preprocessing_funcs +from src.enums.delimiters import Delimiter +from src.utils import llm_utils, string_utils + + def clean_text(contract_text): contract_text = preprocessing_funcs.remove_page_indicators(contract_text) @@ -19,8 +20,8 @@ def split_text(contract_text): """ Splits the contract text into separate pages and filters pages needing quick review. - This function processes the input contract text by splitting it into individual pages - and then categorizes these pages. Pages representing a Quick Review or Cover Page section are separated from + This function processes the input contract text by splitting it into individual pages + and then categorizes these pages. Pages representing a Quick Review or Cover Page section are separated from the main text. Parameters: @@ -33,9 +34,9 @@ def split_text(contract_text): - int: The total number of pages in the contract text. Notes: - - This function utilizes `preprocessing_funcs.split_text` to divide the contract text + - This function utilizes `preprocessing_funcs.split_text` to divide the contract text into a dictionary with page numbers as keys. - - It uses `preprocessing_funcs.filter_quick_review` to separate out pages that need quick + - It uses `preprocessing_funcs.filter_quick_review` to separate out pages that need quick review from the main text dictionary. """ text_dict = preprocessing_funcs.split_text( @@ -59,12 +60,21 @@ def one_to_n_exhibit_chunking(text_dict, filename) -> tuple[dict, dict]: - dict: A dictionary where keys are exhibit page numbers and values are the corresponding exhibit text chunks. - dict: A dictionary where keys are exhibit page numbers and values are the exhibit headers. """ - exhibit_pages, all_exhibit_headers = preprocessing_funcs.get_exhibit_pages(text_dict, filename) - exhibit_pages, all_exhibit_headers = preprocessing_funcs.link_exhibit_pages(all_exhibit_headers, filename) - exhibit_chunk_mapping = preprocessing_funcs.chunk_by_exhibit(text_dict, exhibit_pages) - exhibit_dict = preprocessing_funcs.get_exhibit_dict(text_dict, exhibit_chunk_mapping) + exhibit_pages, all_exhibit_headers = preprocessing_funcs.get_exhibit_pages( + text_dict, filename + ) + exhibit_pages, all_exhibit_headers = preprocessing_funcs.link_exhibit_pages( + all_exhibit_headers, filename + ) + exhibit_chunk_mapping = preprocessing_funcs.chunk_by_exhibit( + text_dict, exhibit_pages + ) + exhibit_dict = preprocessing_funcs.get_exhibit_dict( + text_dict, exhibit_chunk_mapping + ) return exhibit_dict, all_exhibit_headers + def one_to_one_smart_chunking(text_dict, contract_text, one_to_one_fields): """ In current state, this function is a wrapper for smart_chunk_ac. Call any new smart-chunking-related preprocessing funcs here @@ -77,24 +87,44 @@ def one_to_one_smart_chunking(text_dict, contract_text, one_to_one_fields): Returns: dict: ac_chunks """ - ac_chunks = smart_chunking_funcs.smart_chunk_one_to_one(text_dict, contract_text, one_to_one_fields) + ac_chunks = smart_chunking_funcs.smart_chunk_one_to_one( + text_dict, contract_text, one_to_one_fields + ) return ac_chunks -def clean_tables(text_dict: dict[str, str], filename): - - # Get Table Info +def clean_tables(text_dict: dict[str, str], filename: str) -> dict[str, str]: + """ + Process tables with simplified logic: combine continuations, then split by row count. + + Args: + text_dict (dict[str, str]): A dictionary where keys are page numbers and values are the text on those pages. + filename (str): The name of the file being processed. + + Returns: + dict[str, str]: A dictionary with the processed text for each page, with tables split when necessary. + """ + logging.info(f"Starting simple table processing for {filename}") + + # 1. Extract table info table_dict = table_funcs.get_table_info(text_dict) + logging.info(f"Found tables on pages: {list(table_dict.keys())}") - # Combine continuous tables - text_dict_combined = table_funcs.combine_continuous_tables(text_dict, table_dict) + # 2. Combine continuous tables (simple version) + text_dict_combined = table_funcs.combine_continuous_tables( + text_dict=text_dict, + table_dict=table_dict, + ) + logging.info(f"After combination, pages: {list(text_dict_combined.keys())}") - # Split tables - table_dict = table_funcs.get_table_info(text_dict_combined) - text_dict_split = table_funcs.split_long_tables(text_dict_combined, table_dict, config.TABLE_ROW_LIMIT) + # 3. Split tables by row count only + text_dict_final = table_funcs.split_tables_by_rows( + text_dict=text_dict_combined, + row_limit=config.TABLE_ROW_LIMIT, + ) + + logging.info( + f"Final pages after simple table processing: {list(text_dict_final.keys())}" + ) - # Align and Format - table_dict = table_funcs.get_table_info(text_dict_split) - text_dict_final = table_funcs.align_and_format_tables(text_dict_split, table_dict, filename) - return text_dict_final diff --git a/fieldExtraction/src/investment/table_funcs.py b/fieldExtraction/src/investment/table_funcs.py index 574e64d..6699758 100644 --- a/fieldExtraction/src/investment/table_funcs.py +++ b/fieldExtraction/src/investment/table_funcs.py @@ -1,11 +1,13 @@ import ast import copy +import logging import math import re from collections import defaultdict from typing import Any, Dict, List, Optional, Pattern, Tuple, Union import pandas as pd + import src.utils.llm_utils as llm_utils import src.utils.string_utils as string_utils from src import config @@ -21,73 +23,101 @@ CONTINUATION_PAGE_PATTERN: Pattern[str] = re.compile( START_PAGE_NUM_PATTERN: str = r"\n\d+\n?" CONTINUATION_PAGE_NUM_PATTERN: str = r"\d+\n" DICTIONARY_PATTERN: str = r"\{.*\}" +LIST_PATTERN: str = r"\[.*\]" # Markers for table start and end START_MARKER: str = "-------Table Start--------" END_MARKER: str = "-------Table End--------" -def manual_align(page_text, page_tables): +def convert_str_to_dataframe(text: str, is_continuation: bool = False) -> pd.DataFrame: """ - Align tables in the page text manually. + Converts a string representation of table data to a pandas DataFrame. + + Supports both old dict format (for backward compatibility) and new list-of-rows format. + + For list format, treats the first row as column headers and the rest as data rows. + For dict format, assumes the keys are column names and values are lists of column data. Args: - page_text (str): The text of the page. - page_tables (List[Table]): A list of Table objects. - - Returns: - str: The page text with tables aligned - """ - page_text = page_text[0:page_text.find(START_MARKER)] - page_clean = page_text.replace("\n", " ") # Clean tables - - for table in page_tables: - header_text = table.header - header_matches = list(re.finditer(re.escape(header_text), page_clean)) - if len(header_matches) == 1: - page_text = "\n".join([page_text[0:header_matches[0].span()[0]], - header_text, - START_MARKER, - str(table.table), - END_MARKER, - page_text[header_matches[0].span()[1]:]]) - else: - return False - return page_text - - -def align_and_format_tables(text_dict, table_dict, filename): - """ - Manual or prompt-based alignment of tables in the text_dict. - - Args: - text_dict (Dict[str, str]): A dictionary with keys as page numbers and values as page text. - table_dict (Dict[str, List[Table]]): A dictionary with keys as page numbers and values as a list of Table objects. - filename (str): The name of the file being processed. + text (str): The input string containing table data. + is_continuation (bool): Whether this is a continuation table (no header). Defaults to False. + If True, don't treat first row as headers. Returns: - Dict[str, str]: A dictionary with keys as page numbers and values as page text with tables aligned - and formatted. + pd.DataFrame: A DataFrame representation of the table data with string dtype. """ - for page_num, page_text in text_dict.items(): - - if "Table Start" in page_text and string_utils.contains_reimbursement( - page_text - ): - aligned_page = False - # Only attempt manual if 1 table on page - if page_text.count("Table Start") == 1: - aligned_page = manual_align(page_text, table_dict[page_num]) - - # If more than 1 table on page, or if manual_align failed - if not aligned_page: - prompt = preprocessing_prompts.ALIGN_TABLES(page_text) - aligned_page = llm_utils.invoke_claude( - prompt, config.MODEL_ID_CLAUDE35_SONNET, filename, 8192 + + # Try to parse as list first (new format) + list_match = re.search(LIST_PATTERN, text, re.DOTALL) + if list_match: + list_str = list_match.group(0) + try: + parsed_data = ast.literal_eval(list_str) + if isinstance(parsed_data, list): + if not parsed_data: + return pd.DataFrame() + + # Validate that all elements are lists (handle malformed data) + validated_rows = [] + for row in parsed_data: + if isinstance(row, list): + validated_rows.append( + [str(cell) for cell in row] + ) # Convert all to strings + else: + # Skip malformed rows or convert single values to single-item lists + if isinstance(row, str): + validated_rows.append([row]) + # Skip other types entirely + + if not validated_rows: + return pd.DataFrame() + + # Handle variable row lengths by padding with empty strings + max_cols = ( + max(len(row) for row in validated_rows) if validated_rows else 0 ) - - text_dict[page_num] = aligned_page - return text_dict + padded_rows = [ + row + [""] * (max_cols - len(row)) for row in validated_rows + ] + + if is_continuation: + # CONTINUATION: all rows are data, use generic column names + columns = [f"col_{i}" for i in range(max_cols)] + return pd.DataFrame(padded_rows, columns=columns, dtype=str) + else: + # REGULAR: treat first row as headers + if len(padded_rows) > 1: + headers = padded_rows[0] + data_rows = padded_rows[1:] + return pd.DataFrame(data_rows, columns=headers, dtype=str) + elif len(padded_rows) == 1: + # Only one row - treat as headers with no data + headers = padded_rows[0] + return pd.DataFrame(columns=headers, dtype=str) + else: + # Empty list - return empty DataFrame + return pd.DataFrame() + except (SyntaxError, ValueError): + print(f"Invalid list format: {list_str}") + + # Fall back to dict format (old format) - but warn about deprecation + dict_match = re.search(DICTIONARY_PATTERN, text) + if dict_match: + dict_str = dict_match.group(0) + try: + parsed_data = ast.literal_eval(dict_str) + if isinstance(parsed_data, dict): + logging.warning( + "Warning: old table (dict) format detected. Please update to new list-of-rows format by rerunning textract on this file." + ) + # No need to set the headers explicitly, DataFrame will use keys as column names + return pd.DataFrame(parsed_data, dtype=str) + except (SyntaxError, ValueError): + print(f"Invalid dictionary format: {dict_str}") + + return pd.DataFrame() # Return empty DataFrame if no valid format found def convert_str_to_dict(text: str) -> Dict[str, Any]: @@ -116,58 +146,268 @@ def convert_str_to_dict(text: str) -> Dict[str, Any]: class Table: """ - A class to represent a table extracted from a page. + Represents a table extracted from a document page. + + This class handles table data parsing, combination, and splitting operations + for tables found in OCR-processed documents. Attributes: - page_num (str): The page number the table is from. - metadata (str): The metadata of the table. - header (str): The header of the table. - table (Dict[str, Any]): The table data as a dictionary. - start_index (int): The starting index of the table in the page text. - end_index (int): The ending index of the table in the page text. - - Methods: - combine(other: "Table"): Combine another table into this one, keeping only the original keys. - split(row_limit: int) -> List["Table"]: Split the table into multiple tables each with at most row_limit rows. + page_num (str): The page number where the table is located. + metadata (str): Text appearing before the table on the page. + header (str): The table's title/header text. + table (pd.DataFrame): The table data with string dtype. + start_index (int): Starting position of the table in the page text. + end_index (int): Ending position of the table in the page text. + post_table_text (str): Text appearing after the table on the page. + is_continuation (bool): Whether this table continues from a previous page. + original_pages (List[str]): Pages this table spans (for tracking combinations). + + Key Methods: + combine(other): Combine another table into this one with column alignment. + split(row_limit): Split into multiple tables based on row count. + to_list_format(): Convert to list-of-lists format for serialization. """ - def __init__(self, page_num, metadata, header, table, start_index, end_index): + + def __init__( + self, + page_num: str, + metadata: str, + header: str, + table_data: Union[List[List[str]], pd.DataFrame, Dict[str, Any]], + start_index: int, + end_index: int, + post_table_text: str = "", + is_continuation: bool = False, + original_pages: Union[ + List[str], None + ] = None, # Track which pages this table spans + ): self.page_num = page_num self.metadata = metadata self.header = header - self.table = convert_str_to_dict(table) self.start_index = start_index self.end_index = end_index - + self.is_continuation = is_continuation + self.post_table_text: str = post_table_text # Text after the table in the page + self.original_pages = original_pages or [ + page_num + ] # if `original_pages` is None, initialize with current page + + # Handle a few different input types + if isinstance(table_data, str): + self.table = convert_str_to_dataframe( + table_data, is_continuation=is_continuation + ) + elif isinstance(table_data, pd.DataFrame): + self.table = table_data.astype(str) # Ensure DataFrame is of string dtype + elif isinstance(table_data, list): + # Direct list - check if continuation table + if is_continuation: + logging.debug(f"Detected continuation table on page {page_num}.") + # All rows are data rows, no headers + max_cols = max(len(row) for row in table_data) if table_data else 0 + padded_rows = [row + [""] * (max_cols - len(row)) for row in table_data] + columns = [f"col_{i}" for i in range(max_cols)] + self.table = pd.DataFrame(padded_rows, columns=columns, dtype=str) + else: + # Non-continuation table - treat first row as headers + if not table_data: + self.table = pd.DataFrame() + else: + # Handle variable row lengths by padding with empty strings + max_cols = max(len(row) for row in table_data) if table_data else 0 + padded_rows = [ + row + [""] * (max_cols - len(row)) for row in table_data + ] + + if len(padded_rows) > 1: + headers = padded_rows[0] + data_rows = padded_rows[1:] + self.table = pd.DataFrame(data_rows, columns=headers, dtype=str) + elif len(padded_rows) == 1: + # Only one row - treat as headers with no data + headers = padded_rows[0] + self.table = pd.DataFrame(columns=headers, dtype=str) + else: + # Empty list - return empty DataFrame + self.table = pd.DataFrame() + + elif isinstance(table_data, dict): # Old dict format - warn and convert + logging.warning( + "Warning: old table (dict) format detected. Please update to new list-of-rows format by rerunning textract on this file." + ) + self.table = pd.DataFrame(table_data, dtype=str) + else: + self.table = ( + pd.DataFrame() + ) # Default to empty DataFrame if no valid format found + + def get_page_span(self) -> int: + """ + Get the number of pages this table spans. + + Returns: + int: The number of pages this table spans. + """ + return len(set(self.original_pages)) + + def add_page_to_span(self, page_num: str): + """ + Add a page number to the original pages this table spans. + + Args: + page_num (str): The page number to add. + """ + if page_num not in self.original_pages: + self.original_pages.append(page_num) + + def to_list_format(self) -> List[List[str]]: + """ + Convert the DataFrame back to list-of-rows format for serialization, with headers + as the first row. + + Returns: + List[List[str]]: The table data as a list of lists. + """ + if self.table.empty: + return [] + + headers = self.table.columns.tolist() + data_rows = self.table.values.tolist() + + # If continuation table with generic column names, don't include headers + if self.is_continuation and all(col.startswith("col_") for col in headers): + return data_rows # Just the data, no headers + + # For all other cases (regular tables), include headers + return [headers] + data_rows + def combine(self, other: "Table"): - """Combine another table into this one, keeping only the original keys.""" - if len(self.table) != len(other.table): - raise ValueError("Tables must have the same number of columns to combine.") + """ + Combine another table into this one with intelligent column handling. - other_values = list(other.table.values()) # Get values from other table in order - self_keys = list(self.table.keys()) # Original keys order + Handles common OCR data issues like column count mismatches by automatically + padding or truncating the continuation table to match the main table. + + Business Rules: + - Empty tables: Handled gracefully, preserving post-table text + - Column mismatches: Continuation table adjusted to match main table + - Continuation tables: Column names aligned to main table headers + - Post-table text: Accumulated from both tables + + Args: + other: The table to combine into this one (consumed after operation). + + Note: + This method modifies the current table in-place. Column count mismatches + are resolved by padding short rows or truncating long rows in the + continuation table. + """ + + # Handle empty tables + if other.table.empty: + # if other table is empty, just update post_table_text if needed + if other.post_table_text: + self.post_table_text = other.post_table_text + return + + if self.table.empty: + # If this table is empty, just copy the other table + self.table = other.table.copy() + if other.post_table_text: + self.post_table_text = other.post_table_text + return + + # Handle column count mismatches + if self.table.shape[1] != other.table.shape[1]: + main_cols = self.table.shape[1] + other_cols = other.table.shape[1] + + logging.warning( + f"Column count mismatch when combining tables: main table has {main_cols} columns, " + f"continuation table has {other_cols} columns. Adjusting continuation table to match." + ) + + # Create a copy to modify + other_table_adjusted = other.table.copy() + + if other_cols < main_cols: + # Pad with empty columns + for i in range(main_cols - other_cols): + col_name = f"col_{other_cols + i}" + # Ensure unique column name + while col_name in other_table_adjusted.columns: + col_name = f"{col_name}_padded" + other_table_adjusted[col_name] = "" # Add empty column + else: + # Truncate to match main table column count + other_table_adjusted = other_table_adjusted.iloc[:, :main_cols] + + # Ensure column names match (this will be our "other" table for the rest of the method) + other_table_adjusted.columns = self.table.columns + + # Combine the adjusted table + self.table = pd.concat( + [self.table, other_table_adjusted], ignore_index=True + ) + else: + # Column counts match - proceed with normal logic + # If other table has generic column names (continuation table), rename them + if other.is_continuation or ( + list(other.table.columns) != list(self.table.columns) + ): + other_renamed = other.table.copy() + other_renamed.columns = self.table.columns # Use our column names + self.table = pd.concat([self.table, other_renamed], ignore_index=True) + else: + self.table = pd.concat([self.table, other.table], ignore_index=True) + + if other.post_table_text.strip(): # Only if the other table has post_table_text + if self.post_table_text.strip(): + if not other.post_table_text.strip() == self.post_table_text.strip(): + # If both tables have post_table_text AND they're not identical, append the other table's text + self.post_table_text = ( + f"{self.post_table_text}\n{other.post_table_text}" + ) + else: + # Only the other table has post_table_text, so just set it + self.post_table_text = other.post_table_text - for i, key in enumerate(self_keys): - self.table[key].extend(other_values[i]) # Append values in order - def split(self, row_limit: int) -> List["Table"]: """Split the table into multiple tables each with at most row_limit rows.""" + if len(self.table) <= row_limit: + return [self] + tables = [] - keys = list(self.table.keys()) - num_rows = len(self.table[keys[0]]) - for start_row in range(0, num_rows, row_limit): - end_row = min(start_row + row_limit, num_rows) - new_table_data = {key: self.table[key][start_row:end_row] for key in keys} - new_table_str = str(new_table_data) - tables.append(Table(self.page_num, self.metadata, self.header, new_table_str, self.start_index, self.end_index)) + for start_row in range(0, len(self.table), row_limit): + end_row = min(start_row + row_limit, len(self.table)) + subset_df = self.table.iloc[start_row:end_row].copy() + + # Create new table with DataFrame directly + new_table = Table( + page_num=self.page_num, + metadata=self.metadata, + header=self.header, + table_data=subset_df, + start_index=self.start_index, # Keep original start index + end_index=self.end_index, # Keep original end index + ) + tables.append(new_table) + + if tables: + # Only the LAST split table gets the post-table text + tables[-1].post_table_text = self.post_table_text + return tables - -def get_table_info(text_dict): + + +def get_table_info(text_dict: Dict[str, str]) -> Dict[str, List[Table]]: """ Extracts all of the Table objects from the text_dict. Args: text_dict (Dict[str, str]): A dictionary with keys as page numbers and values as page text. - + Returns: Dict[str, List[Table]]: A dictionary with keys as page numbers and values as a list of Table objects. """ @@ -189,143 +429,270 @@ def get_table_info(text_dict): metadata = page_text[:first_table_start_index].strip() header_start = start_index + len(START_MARKER) - header_end = page_text.find("{", header_start) - header = page_text[header_start:header_end].strip() - table = page_text[header_end:end_index].strip() - tables.append(Table(page_num, metadata, header, table, start_index, end_index + len(END_MARKER))) + # Look for either [ or { to find start of table data + list_start = page_text.find("[", header_start) + dict_start = page_text.find("{", header_start) + + if list_start != -1 and (dict_start == -1 or list_start < dict_start): + # New list format + # If list starts first, we assume it's a list of rows + header = page_text[header_start:list_start].strip() + table_data = page_text[list_start:end_index].strip() + + is_continuation = ( + header == "" # No header text + and len(page_text[header_start:list_start].strip()) + == 0 # Nothing between marker and `[` + ) + elif dict_start != -1: + # Old dict format + # If dict starts first, we assume it's a dictionary + header = page_text[header_start:dict_start].strip() + table_data = page_text[dict_start:end_index].strip() + + is_continuation = header == "" + else: + # No table data found + break + + tables.append( + Table( + page_num, + metadata, + header, + table_data, + start_index, + end_index + len(END_MARKER), + is_continuation=is_continuation, + ) + ) start_index = end_index + len(END_MARKER) + + if tables: + last_table_end = tables[-1].end_index + # If there's text after the last table, store it as post-table text + post_table_text = page_text[last_table_end:].strip() + + # Store post-table text in the last table + tables[-1].post_table_text = post_table_text table_dict[page_num] = tables return table_dict -def recreate_page(page_table_list): +def recreate_page(page_table_list: List[Table]) -> str: """ Puts together the tables in the page_table_list to recreate the page text. Args: page_table_list (List[Table]): A list of Table objects. - + Returns: str: The text of the page with the tables combined. """ page_text = "" for i in range(len(page_table_list)): t = page_table_list[i] + table_as_list = t.to_list_format() if i == 0: - page_text += "\n".join([t.metadata, START_MARKER, t.header, str(t.table), END_MARKER]) + page_text += "\n".join( + [t.metadata, START_MARKER, t.header, str(table_as_list), END_MARKER] + ) else: - page_text += "\n".join(["\n", START_MARKER, t.header, str(t.table), END_MARKER]) + page_text += "\n".join( + ["\n", START_MARKER, t.header, str(table_as_list), END_MARKER] + ) + + # Add post-table text if it exists + if t.post_table_text.strip(): + page_text += "\n" + t.post_table_text.strip() return page_text -def remove_table_from_page(page_text: str, t: Table): - """ - Remove the table from the page text. + +def split_large_table_with_headers( + table: Table, row_limit: int = config.TABLE_ROW_LIMIT +) -> List[str]: + """Split a large table into smaller chunks while preserving headers. Args: - page_text (str): The text of the page. - t (Table): The table to remove. - + table (Table): The table to split. + row_limit (int): The row limit to use. By default uses config.TABLE_ROW_LIMIT. + Returns: - str: The page text with the table removed. + List[str]: A list of text chunks representing the split table. """ - # Remove the table from the page text - table_text = page_text[t.start_index:t.end_index] - page_text = page_text.replace(table_text, "").strip() - return page_text + # Use existing table splitting logic + split_tables = split_table_by_rows(table, row_limit) -def combine_continuous_tables(text_dict, table_dict): - """ - Combine tables that are split across multiple pages. + chunks = [] + for i, split_table in enumerate(split_tables): + chunk_parts = [] + # add metadata if it exists + if split_table.metadata: + chunk_parts.append(split_table.metadata) + + # Add tables structure + chunk_parts.extend( + [ + START_MARKER, + split_table.header, # Header replicated in each chunk + str(split_table.to_list_format()), # Convert table to list format + END_MARKER, + ] + ) + + # Add post-table text only to the last chunk + if i == len(split_tables) - 1 and table.post_table_text: + chunk_parts.append(table.post_table_text) + + chunks.append("\n".join(chunk_parts).strip()) + + return chunks + + +def combine_continuous_tables( + text_dict: Dict[str, str], table_dict: Dict[str, List[Table]] +) -> Dict[str, str]: + """Combine continuous tables across pages. Args: text_dict (Dict[str, str]): A dictionary with keys as page numbers and values as page text. table_dict (Dict[str, List[Table]]): A dictionary with keys as page numbers and values as a list of Table objects. Returns: - Dict[str, str]: A dictionary with keys as page numbers and values as page text with tables combined if they are split - across multiple pages. + Dict[str, str]: Updated text_dict with combined tables. """ - text_dict_combined = {} - page_nums = sorted(text_dict.keys(), key=int) # Ensure page numbers are sorted + pages_to_remove = set() + updated_text_dict = text_dict.copy() - i = 0 - while i < len(page_nums): - page_num = page_nums[i] - current_page_tables = table_dict[page_num] - - # If no table on page, continue - if len(current_page_tables) == 0: - text_dict_combined[page_num] = text_dict[page_num] - i += 1 - continue - - combined_tables = current_page_tables[:] - next_page_num = str(int(page_num) + 1) - - while next_page_num in table_dict: - next_page_tables = table_dict[next_page_num] - if not next_page_tables or not string_utils.is_empty(next_page_tables[0].header) or len(combined_tables[-1].table.keys()) != len(next_page_tables[0].table.keys()): # Edge Case: could also add criteria that metadata is empty as well - break - - combined_tables[-1].combine(next_page_tables[0]) # Combine tables - text_dict[next_page_num] = remove_table_from_page(text_dict[next_page_num], next_page_tables[0]) - combined_tables.extend(next_page_tables[1:]) # Add remaining tables from next page - - i += 1 - next_page_num = str(int(next_page_num) + 1) - - new_page_text = recreate_page(combined_tables) - text_dict_combined[page_num] = new_page_text - i += 1 - - return text_dict_combined - - -def split_long_tables(text_dict, table_dict, row_limit): - """ - Split tables in the text_dict into multiple tables if they exceed the row limit. - - Args: - text_dict (Dict[str, str]): A dictionary with keys as page numbers and values as page text. - table_dict (Dict[str, List[Table]]): A dictionary with keys as page numbers and values as a list of Table objects. - row_limit (int): The maximum number of rows a table can have. - - Returns: - Dict[str, str]: A dictionary with keys as page numbers and values as page text with tables split if they exceed the row - limit. - """ - new_text_dict = {} for page_num, page_text in text_dict.items(): - tables = table_dict[page_num] - if not tables: - new_text_dict[page_num] = page_text + if page_num in pages_to_remove: continue - subpage_index = 0 - for table in tables: - if table.table and len(table.table[list(table.table.keys())[0]]) > row_limit: - split_tables = table.split(row_limit) - for split_table in split_tables: - new_page_num = f"{page_num}.{subpage_index}" - new_text_dict[new_page_num] = "\n".join([split_table.metadata, START_MARKER, split_table.header, str(split_table.table), END_MARKER]) - subpage_index += 1 + if page_num not in table_dict or not table_dict[page_num]: + # No tables on this page, nothing to combine + continue + + # Use the existing is_continuation flag from the table object + first_table = table_dict[page_num][0] + + if first_table.is_continuation: + # This is a continuation table, find the main table to combine with + main_table_page = None + for prev_page in sorted( + [p for p in table_dict.keys() if int(p) < int(page_num)], + key=string_utils.page_key_sort, + reverse=True, + ): + if prev_page in table_dict and table_dict[prev_page]: + main_table_page = prev_page + break + + if main_table_page and main_table_page not in pages_to_remove: + # Combine the tables + main_table = table_dict[main_table_page][ + -1 + ] # Last table on the main page + continuation_table = table_dict[page_num][ + 0 + ] # First table on current page + main_table.combine(continuation_table) + + # Update main page with ALL tables from that page + updated_text_dict[main_table_page] = recreate_page( + table_dict[main_table_page] + ) + + # Mark continuation page for removal + pages_to_remove.add(page_num) + # Remove continuation pages (we combined them into the main pages) + for page_num in pages_to_remove: + updated_text_dict.pop(page_num, None) + + return updated_text_dict + + +def split_tables_by_rows(text_dict: Dict[str, str], row_limit: int) -> Dict[str, str]: + """Split tables that exceed the row limit, add .0 suffix to pages with tables + + Args: + text_dict (Dict[str, str]): A dictionary with keys as page numbers and values as page text. + row_limit (int): The maximum number of rows allowed in a single table. + + Returns: + Dict[str, str]: Updated text_dict with split tables and .0 suffixes for table pages + """ + result = {} + + for page_num, page_text in text_dict.items(): + # get tables on this page + page_table_dict = get_table_info({page_num: page_text}) + + if page_num in page_table_dict and page_table_dict[page_num]: + # This page has tables + table = page_table_dict[page_num][0] + + if len(table.table) > row_limit: + # Large table - split it + logging.info( + f"Splitting table on page {page_num}: {len(table.table)} rows > row_limit ({row_limit})" + ) + split_chunks = split_large_table_with_headers(table, row_limit) + for i, chunk in enumerate(split_chunks): + # Use .0, .1, .2 suffixes for split tables + result[f"{page_num}.{i}"] = chunk.strip() else: - new_page_num = f"{page_num}.{subpage_index}" - new_text_dict[new_page_num] = "\n".join([table.metadata, START_MARKER, table.header, str(table.table), END_MARKER]) - subpage_index += 1 - - return new_text_dict - - + # Small table - keep it intact + logging.info( + f"Keeping table on page {page_num}: {len(table.table)} rows <= row_limit ({row_limit})" + ) + result[f"{page_num}.0"] = page_text + else: + # No tables on this page, keep it as is + result[page_num] = page_text + return result +def split_table_by_rows(table: Table, row_limit: int) -> List[Table]: + """Split a single table into multiple tables if it exceeds the row limit. + Args: + table (Table): The table to split. + row_limit (int): The maximum number of rows allowed in a single table. + Returns: + List[Table]: A list of tables, each with a maximum of row_limit rows. + """ + df = table.table + total_rows = len(df) + if total_rows <= row_limit: + return [table] + split_tables = [] + for start_idx in range(0, total_rows, row_limit): + end_idx = min(start_idx + row_limit, total_rows) + # Create subset DataFrame + subset_df = df.iloc[start_idx:end_idx].copy() + # Create new table object + split_table = Table( + page_num=table.page_num, + metadata=table.metadata, + header=table.header, + table_data=subset_df, + start_index=table.start_index, # Keep original start index + end_index=table.end_index, # Keep original end index + post_table_text="", # Only last split gets post-table text + is_continuation=False, + original_pages=table.original_pages.copy(), + ) + split_tables.append(split_table) + # Add post-table text to the last split + if split_tables: + split_tables[-1].post_table_text = table.post_table_text + return split_tables diff --git a/fieldExtraction/tests/preprocessing_integration_test.py b/fieldExtraction/tests/preprocessing_integration_test.py new file mode 100644 index 0000000..9474a2a --- /dev/null +++ b/fieldExtraction/tests/preprocessing_integration_test.py @@ -0,0 +1,130 @@ +import unittest +from unittest.mock import patch + +import pytest + +from src.investment.preprocess import clean_tables + + +class TestPreprocessingIntegration(unittest.TestCase): + + def test_clean_tables_simple_end_to_end(self): + """Test the complete simplified table cleaning pipeline.""" + text_dict = { + "1": "Metadata\n-------Table Start--------\nMain Table\n[['A', 'B']] + [[f'Row{i}', str(i)] for i in range(1, 8)]\n-------Table End--------", + "2": "-------Table Start--------\n\n[['Row8', '8'], ['Row9', '9']]\n-------Table End--------", # Continuation + "3": "Text without tables", + } + + # This should: 1) Combine continuations, 2) Split by row count + result = clean_tables(text_dict, "test_file.txt") + + # Should have combined page 2 into page 1, then split the long table + self.assertNotIn("2", result) # Continuation removed + self.assertIn("3", result) # Non-table page preserved + # Should have split pages like "1.0", "1.1", etc. + + def test_clean_tables_small_table_gets_dot_zero_suffix(self): + """Test that pages with small tables get .0 suffix but aren't split.""" + text_dict = { + "15": "Rate Table\n-------Table Start--------\nRate Table\n[['HCPC', 'Rate'], ['T1019', '$5.00']]\n-------Table End--------" + } + + result = clean_tables(text_dict, "test_file.pdf") + + # Should get .0 suffix because it contains tables + self.assertIn("15.0", result) + self.assertNotIn("15", result) + + # Content should be preserved + self.assertIn("T1019", result["15.0"]) + + # Should not be split (small table) + chunk_keys = [k for k in result.keys() if k.startswith("15.")] + self.assertEqual(len(chunk_keys), 1) + + def test_clean_tables_continuation_gets_combined_and_dot_zero(self): + """Test that continuations get combined and result gets .0 suffix.""" + text_dict = { + "1": "Rate Table\n-------Table Start--------\nRate Table Header\n[['HCPC', 'Rate'], ['T1019', '$5.00']]\n-------Table End--------", + "2": "-------Table Start--------\n\n[['T1021', '$4.00']]\n-------Table End--------", # Continuation + } + + result = clean_tables(text_dict, "test_file.pdf") + + # Should combine and get .0 suffix since contains tables + self.assertIn("1.0", result) + self.assertNotIn("1", result) + self.assertNotIn("2", result) # Absorbed into Page 1 + + # Verify combined content + self.assertIn("T1019", result["1.0"]) + self.assertIn("T1021", result["1.0"]) + + def test_clean_tables_large_table_row_based_splitting(self): + """Test that large tables get split based on row count.""" + # Create a table with many rows that will exceed TABLE_ROW_LIMIT + large_table_rows = [["HCPC", "Rate"]] + [ + [f"T{i:04d}", f"${i}.00"] for i in range(60) + ] # 60 data rows + text_dict = { + "31": f"EXHIBIT F\n-------Table Start--------\nPCAP Rates\n{large_table_rows}\n-------Table End--------\nPost-table text" + } + + result = clean_tables(text_dict, "test_file.pdf") + + # Should create multiple chunks due to row limit (default 50) + chunk_keys = [k for k in result.keys() if k.startswith("31.")] + self.assertGreater( + len(chunk_keys), + 1, + f"Expected multiple chunks for 60 rows, got: {chunk_keys}", + ) + + # Original page should be gone + self.assertNotIn("31", result) + + # All table data should be preserved somewhere + all_content = "".join(result.values()) + self.assertIn("T0000", all_content) + self.assertIn("T0059", all_content) + self.assertIn("EXHIBIT F", all_content) + + # Post-table text should be preserved on last chunk + last_chunk = max(chunk_keys) + self.assertIn("Post-table text", result[last_chunk]) + + def test_clean_tables_no_tables_unchanged(self): + """Test that pages without tables remain unchanged.""" + text_dict = { + "5": "This is just text content with no tables.", + "6": "More text content without any table markers.", + } + + result = clean_tables(text_dict, "test_file.pdf") + + # Should remain unchanged (no .0 suffix) + assert "5" in result + assert "6" in result + assert "5.0" not in result + assert "6.0" not in result + + # Content should be identical + assert result["5"] == text_dict["5"] + assert result["6"] == text_dict["6"] + + def test_clean_tables_multiple_tables_same_page(self): + """Test that multiple tables on same page get .0 suffix but stay together.""" + text_dict = { + "1": "Page metadata\n-------Table Start--------\nFirst Table\n[['Col1', 'Col2'], ['A', 'B']]\n-------Table End--------\n\nSome text\n-------Table Start--------\nSecond Table\n[['X', 'Y'], ['1', '2']]\n-------Table End--------" + } + + result = clean_tables(text_dict, "test_file.pdf") + + # Should get .0 suffix because it contains tables + assert "1.0" in result + assert "1" not in result + + # Both tables should be present in single chunk + assert "First Table" in result["1.0"] + assert "Second Table" in result["1.0"] diff --git a/fieldExtraction/tests/table_funcs_test.py b/fieldExtraction/tests/table_funcs_test.py index 9b2db26..77470df 100644 --- a/fieldExtraction/tests/table_funcs_test.py +++ b/fieldExtraction/tests/table_funcs_test.py @@ -1,49 +1,859 @@ +import logging import unittest -from src.investment.table_funcs import Table, get_table_info, recreate_page, remove_table_from_page, combine_continuous_tables, split_long_tables +from unittest.mock import patch + +import pandas as pd +from sympy import content + +from src.investment.table_funcs import (Table, combine_continuous_tables, + convert_str_to_dataframe, + get_table_info, recreate_page, + split_large_table_with_headers, + split_tables_by_rows) + +logging.basicConfig(level=logging.INFO) + class TestTableFunctions(unittest.TestCase): def setUp(self): - self.sample_text_dict = { - "1": "Some metadata\n-------Table Start--------\nTable 1\n{'A': ['1', '2'], 'B': ['3', '4']}\n-------Table End--------", - "2": "More metadata\n-------Table Start--------\n\n{'A': ['5', '6'], 'B': ['7', '8']}\n-------Table End--------" + # New list-of-rows format + self.sample_text_dict_new = { + "1": "Some metadata\n-------Table Start--------\nTable 1\n[['A', 'B'], ['1', '3'], ['2', '4']]\n-------Table End--------", + "2": "More metadata\n-------Table Start--------\n\n[['100', '200'], ['5', '7'], ['6', '8']]\n-------Table End--------", + } + # Old dict format for backward compatibility testing + self.sample_text_dict_old = { + "1": "Some metadata\n-------Table Start--------\nTable 1\n{'A': ['1', '2'], 'B': ['3', '4']}\n-------Table End--------", + "2": "More metadata\n-------Table Start--------\n\n{'A': ['5', '6'], 'B': ['7', '8']}\n-------Table End--------", + } + # Multiple tables on one page + self.multi_table_text = { + "1": "Page metadata\n-------Table Start--------\nFirst Table\n[['Col1', 'Col2'], ['A', 'B'], ['C', 'D']]\n-------Table End--------\n\nSome text\n-------Table Start--------\nSecond Table\n[['X', 'Y'], ['1', '2']]\n-------Table End--------" + } + # Long table for splitting tests + self.long_table_text = { + "1": "Long table metadata\n-------Table Start--------\nLong Table\n[['Name', 'Value'], ['Row1', '1'], ['Row2', '2'], ['Row3', '3'], ['Row4', '4'], ['Row5', '5']]\n-------Table End--------" } - - self.sample_table_dict = get_table_info(self.sample_text_dict) - def test_get_table_info(self): - self.assertIn("1", self.sample_table_dict) - self.assertEqual(len(self.sample_table_dict["1"]), 1) - self.assertEqual(self.sample_table_dict["1"][0].header, "Table 1") - + def test_convert_str_to_dataframe_new_format(self): + """Test conversion from new list-of-rows format.""" + table_str = "[['A', 'B'], ['1', '3'], ['2', '4']]" + df = convert_str_to_dataframe(table_str) + + self.assertIsInstance(df, pd.DataFrame) + self.assertEqual(df.shape, (2, 2)) # 2 rows, 2 columns + self.assertEqual(df.columns[0], "A") + self.assertEqual(df.iloc[0, 0], "1") + self.assertEqual(df.iloc[1, 0], "2") + self.assertEqual(df.iloc[0, 1], "3") + self.assertTrue( + df.dtypes.eq("object").all() + ) # All columns should be of type str (object) + + def test_convert_str_to_dataframe_old_format(self): + """Test conversion from old dict format with deprecation warning.""" + table_str = "{'A': ['1', '2'], 'B': ['3', '4']}" + with self.assertLogs(level="INFO") as log: + df = convert_str_to_dataframe(table_str) + + # Should warn about old format + self.assertTrue( + any("old table (dict) format detected" in message for message in log.output) + ) + + self.assertIsInstance(df, pd.DataFrame) + self.assertEqual(df.shape, (2, 2)) # 2 rows, 2 columns + self.assertTrue( + df.dtypes.eq("object").all() + ) # All columns should be of type str (object) + self.assertEqual(df.iloc[0, 0], "1") + self.assertEqual(df.iloc[0, 1], "3") + self.assertEqual(df.iloc[1, 0], "2") + self.assertEqual(df.iloc[1, 1], "4") + + def test_convert_str_to_dataframe_invalid_format(self): + """Test handling of invalid formats.""" + invalid_str = "not a valid table format" + df = convert_str_to_dataframe(invalid_str) + + self.assertIsInstance(df, pd.DataFrame) + self.assertTrue(df.empty) + + def test_table_init_new_format(self): + """Test Table initialization with new list-of-rows format.""" + table_data = [["A", "B"], ["1", "3"], ["2", "4"]] + table = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data=table_data, + start_index=0, + end_index=100, + ) + + self.assertEqual(table.page_num, "1") + self.assertEqual(table.metadata, "metadata") + self.assertEqual(table.header, "Table 1") + self.assertEqual(table.table.shape, (2, 2)) # 2 rows, 2 columns + self.assertEqual(list(table.table.columns), ["A", "B"]) + self.assertEqual(table.table.iloc[0, 0], "1") + + def test_table_init_dataframe(self): + """Test Table initialization with a DataFrame.""" + table_data = pd.DataFrame([["A", "B"], ["1", "3"]], dtype=str) + table = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data=table_data, + start_index=0, + end_index=100, + ) + self.assertEqual(table.table.shape, (2, 2)) # 2 rows, 2 columns + self.assertTrue( + table.table.dtypes.eq("object").all() + ) # All columns should be of type str (object) + + def test_table_init_column_mismatch_padding(self): + """Test table initialization handles column count mismatches by padding.""" + # Headers have 3 columns, but data rows have different counts + table_data = [ + ["A", "B", "C"], # 3 columns (header) + ["1", "2"], # 2 columns - should be padded + ["3", "4", "5", "6"], # 4 columns - determines max_cols = 4 + ["7"], # 1 column - should be padded + ] + table = Table( + page_num="1", + metadata="metadata", + header="Test Table", + table_data=table_data, + start_index=0, + end_index=100, + ) + + # Should have 4 columns (max from all rows) and 3 data rows + self.assertEqual(table.table.shape, (3, 4)) + self.assertEqual( + list(table.table.columns), ["A", "B", "C", ""] + ) # Header padded too + + # Check first row (padded to 4 columns) + self.assertEqual(table.table.iloc[0, 0], "1") + self.assertEqual(table.table.iloc[0, 1], "2") + self.assertEqual(table.table.iloc[0, 2], "") # Padded + self.assertEqual(table.table.iloc[0, 3], "") # Padded + + # Check second row (kept at 4 columns) + self.assertEqual(table.table.iloc[1, 0], "3") + self.assertEqual(table.table.iloc[1, 1], "4") + self.assertEqual(table.table.iloc[1, 2], "5") + self.assertEqual(table.table.iloc[1, 3], "6") # 4th column preserved + + # Check third row (heavily padded) + self.assertEqual(table.table.iloc[2, 0], "7") + self.assertEqual(table.table.iloc[2, 1], "") # Padded + self.assertEqual(table.table.iloc[2, 2], "") # Padded + self.assertEqual(table.table.iloc[2, 3], "") # Padded + + def test_table_init_continuation_column_handling(self): + """Test continuation table handles column mismatches with generic column names.""" + # Continuation table - no header row, just data with varying lengths + table_data = [ + ["1", "2"], # 2 columns + ["3", "4", "5"], # 3 columns + ["6"], # 1 column + ] + table = Table( + page_num="2", + metadata="", + header="", + table_data=table_data, + start_index=0, + end_index=100, + is_continuation=True, + ) + + # Should use max column count (3) and generate generic column names + self.assertEqual(table.table.shape, (3, 3)) + self.assertEqual(list(table.table.columns), ["col_0", "col_1", "col_2"]) + + # Check padding/truncation + self.assertEqual(table.table.iloc[0, 0], "1") + self.assertEqual(table.table.iloc[0, 1], "2") + self.assertEqual(table.table.iloc[0, 2], "") # Padded + + self.assertEqual(table.table.iloc[1, 0], "3") + self.assertEqual(table.table.iloc[1, 1], "4") + self.assertEqual(table.table.iloc[1, 2], "5") + + self.assertEqual(table.table.iloc[2, 0], "6") + self.assertEqual(table.table.iloc[2, 1], "") # Padded + self.assertEqual(table.table.iloc[2, 2], "") # Padded + + def test_combine_continuous_tables_simple(self): + """Test the simplified continuous table combination logic.""" + # Create test data with continuation table + text_dict = { + "1": "Metadata 1\n-------Table Start--------\nMain Table\n[['Name', 'Value'], ['A', '1'], ['B', '2']]\n-------Table End--------", + "2": "-------Table Start--------\n\n[['C', '3'], ['D', '4']]\n-------Table End--------", # Continuation (empty header) + "3": "Metadata 3\n-------Table Start--------\nAnother Table\n[['X', 'Y'], ['Z', '5']]\n-------Table End--------", + } + + # Debug: Let's see what get_table_info thinks about these tables + table_dict = get_table_info(text_dict) + print(f"Table dict keys: {list(table_dict.keys())}") + for page_num, tables in table_dict.items(): + for i, table in enumerate(tables): + print( + f"Page {page_num}, Table {i}: is_continuation={table.is_continuation}, header='{table.header}'" + ) + + result = combine_continuous_tables(text_dict, table_dict) + print(f"Result keys: {list(result.keys())}") + + # Page 2 should be removed (combined into page 1) + self.assertNotIn("2", result) + self.assertIn("1", result) + self.assertIn("3", result) + + # Page 1 should now contain the combined table + combined_table_dict = get_table_info({"1": result["1"]}) + combined_table = combined_table_dict["1"][0] + + # Should have 4 rows total (2 from main + 2 from continuation) + self.assertEqual(combined_table.table.shape, (4, 2)) + self.assertEqual(combined_table.table.iloc[0, 0], "A") # First table data + self.assertEqual(combined_table.table.iloc[2, 0], "C") # Continuation data + self.assertEqual(combined_table.table.iloc[3, 0], "D") # Continuation data + + def test_combine_continuous_tables_multiple_tables_per_page(self): + """Test that continuation tables combine with the LAST table on previous page, not the first.""" + # Page 1 has TWO tables: Table A (complete) and Table B (continues on page 2) + text_dict = { + "1": ( + "Metadata 1\n" + "-------Table Start--------\nTable A (Complete)\n[['Col1', 'Col2'], ['A1', 'A2'], ['A3', 'A4']]\n-------Table End--------\n" + "Some text between tables\n" + "-------Table Start--------\nTable B (Continues)\n[['Name', 'Value'], ['B1', '1'], ['B2', '2']]\n-------Table End--------" + ), + "2": ( + "-------Table Start--------\n\n[['B3', '3'], ['B4', '4']]\n-------Table End--------" # Continuation of Table B + ), + } + + table_dict = get_table_info(text_dict) + + # Verify setup: Page 1 should have 2 tables, Page 2 should have 1 continuation table + self.assertEqual(len(table_dict["1"]), 2) + self.assertEqual(len(table_dict["2"]), 1) + self.assertFalse( + table_dict["1"][0].is_continuation + ) # Table A - not continuation + self.assertFalse( + table_dict["1"][1].is_continuation + ) # Table B - not continuation + self.assertTrue( + table_dict["2"][0].is_continuation + ) # Table B continuation - is continuation + + # Get original row counts before combining + table_a_original_rows = table_dict["1"][0].table.shape[0] # Should be 2 rows + table_b_original_rows = table_dict["1"][1].table.shape[0] # Should be 2 rows + continuation_rows = table_dict["2"][0].table.shape[0] # Should be 2 rows + + # Run the combination + result = combine_continuous_tables(text_dict, table_dict) + + # Page 2 should be removed (combined into page 1) + self.assertNotIn("2", result) + self.assertIn("1", result) + + # Parse the combined result to verify tables + combined_table_dict = get_table_info({"1": result["1"]}) + combined_tables = combined_table_dict["1"] + + # Should still have 2 tables on page 1 + self.assertEqual(len(combined_tables), 2) + + # Table A (first table) should be unchanged + table_a_after = combined_tables[0] + self.assertEqual( + table_a_after.table.shape[0], table_a_original_rows + ) # Same row count + self.assertEqual(table_a_after.header, "Table A (Complete)") + self.assertEqual( + table_a_after.table.iloc[0, 0], "A1" + ) # Original data preserved + + # Table B (second table) should now include the continuation data + table_b_after = combined_tables[1] + expected_combined_rows = table_b_original_rows + continuation_rows # 2 + 2 = 4 + self.assertEqual(table_b_after.table.shape[0], expected_combined_rows) + self.assertEqual(table_b_after.header, "Table B (Continues)") + + # Verify Table B has both original and continuation data + self.assertEqual(table_b_after.table.iloc[0, 0], "B1") # Original data + self.assertEqual(table_b_after.table.iloc[1, 0], "B2") # Original data + self.assertEqual(table_b_after.table.iloc[2, 0], "B3") # Continuation data + self.assertEqual(table_b_after.table.iloc[3, 0], "B4") # Continuation data + + def test_combine_continuous_tables_wrong_combination_scenario(self): + """Test edge case: what if we accidentally used [0] instead of [-1]?""" + # This test documents what WOULD happen with the wrong indexing + # Page 1: Table A + Table B (continues) + # Page 2: Table B continuation + # If we used [0], continuation would wrongly combine with Table A + + text_dict = { + "1": ( + "-------Table Start--------\nTable A\n[['X', 'Y'], ['X1', 'Y1']]\n-------Table End--------\n" + "-------Table Start--------\nTable B\n[['Name', 'Value'], ['B1', '1']]\n-------Table End--------" + ), + "2": "-------Table Start--------\n\n[['B2', '2']]\n-------Table End--------", # Should combine with Table B, not Table A + } + + table_dict = get_table_info(text_dict) + result = combine_continuous_tables(text_dict, table_dict) + + # Get the result and verify correct combination + combined_table_dict = get_table_info({"1": result["1"]}) + + # Table A (index 0) should have 1 row still - NOT combined with continuation + table_a = combined_table_dict["1"][0] + self.assertEqual(table_a.table.shape[0], 1) # Still just 1 row + self.assertEqual(table_a.table.iloc[0, 0], "X1") # Original data only + + # Table B (index 1, which is -1) should have 2 rows - correctly combined + table_b = combined_table_dict["1"][1] + self.assertEqual(table_b.table.shape[0], 2) # 1 original + 1 continuation = 2 + self.assertEqual(table_b.table.iloc[0, 0], "B1") # Original + self.assertEqual(table_b.table.iloc[1, 0], "B2") # Continuation + + def test_table_combine_same_header(self): + """Test combining two tables with the same header.""" + table1 = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data=[["Column 1", "Column 2"], ["A", "B"], ["1", "2"]], + start_index=0, + end_index=100, + ) + table2 = Table( + page_num="1", + metadata="metadata", + header="Table 2", + table_data=[["Column 1", "Column 2"], ["C", "D"], ["3", "4"]], + start_index=0, + end_index=100, + ) + + original_shape = table1.table.shape + table1.combine(table2) + self.assertEqual( + table1.table.shape[0], original_shape[0] + table2.table.shape[0] + ) # Rows should be combined + self.assertEqual(table1.table.shape, (4, 2)) # 4 rows (2 + 2 rows), 2 columns + self.assertEqual(table1.table.iloc[0, 0], "A") + self.assertEqual( + table1.table.iloc[2, 0], "C" + ) # Check if second table's data is appended correctly + self.assertEqual(table1.table.iloc[3, 1], "4") # Last row of second table + + def test_table_combine_different_headers(self): + """Test combining two tables with different headers (the second table is assumed to be a continuation in this case). + We will pass in table_data as STRING for both we use convert_str_to_dataframe and + check that `is_continuation` is getting instantiated correctly. + """ + + table1 = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data="[['Column 1', 'Column 2'], ['A', 'B'], ['1', '2']]", + start_index=0, + end_index=100, + ) + + # Continuation table with empty header + table2 = Table( + page_num="1", + metadata="metadata", + header="", + table_data="[['C', 'D'], ['3', '4'], ['E', 'F']]", # Pass as string! + start_index=0, + end_index=100, + is_continuation=True, # Mark as continuation + ) + + original_shape = table1.table.shape + table1.combine(table2) + + print(table1.table) + self.assertEqual( + table1.table.shape[0], original_shape[0] + table2.table.shape[0] + ) # Rows should be combined + self.assertEqual(table1.table.shape, (5, 2)) # 5 rows (2 + 3 rows), 2 columns + self.assertEqual(table1.table.iloc[0, 0], "A") + self.assertEqual( + table1.table.iloc[2, 0], "C" + ) # Check if second table's data is appended correctly + self.assertEqual(table1.table.iloc[3, 1], "4") # Last row of second table + self.assertEqual( + list(table1.table.columns), ["Column 1", "Column 2"] + ) # Should keep the first table's header + + def test_table_combine_different_columns(self): + """Test that combining tables with different column counts handles mismatches gracefully.""" + table1 = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data=[["A", "B"], ["1", "2"]], + start_index=0, + end_index=100, + ) + table2 = Table( + page_num="1", + metadata="metadata", + header="Table 2", + table_data=[["C", "D", "E"], ["3", "4", "5"]], # 3 columns vs 2 columns + start_index=0, + end_index=100, + ) + + # Should not raise an error anymore - should handle column mismatch + table1.combine(table2) + + # Verify the result + # Should have 2 rows total (1 from each table) + self.assertEqual(table1.table.shape[0], 2) + + # Should have 2 columns (from the main table) + self.assertEqual(table1.table.shape[1], 2) + + # Column names should match the main table + self.assertEqual(list(table1.table.columns), ["A", "B"]) + + # First row should be from table1 + self.assertEqual(table1.table.iloc[0, 0], "1") + self.assertEqual(table1.table.iloc[0, 1], "2") + + # Second row should be from table2, but truncated to 2 columns + self.assertEqual(table1.table.iloc[1, 0], "3") + self.assertEqual(table1.table.iloc[1, 1], "4") + # The "5" from column "E" should be truncated/lost + + def test_table_combine_column_padding(self): + """Test that combining tables pads shorter tables with empty columns.""" + table1 = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data=[["A", "B", "C"], ["1", "2", "3"]], # 3 columns + start_index=0, + end_index=100, + ) + table2 = Table( + page_num="1", + metadata="metadata", + header="Table 2", + table_data=[["X"], ["4"]], # Only 1 column + start_index=0, + end_index=100, + ) + + table1.combine(table2) + + # Should have 2 rows total + self.assertEqual(table1.table.shape[0], 2) + + # Should have 3 columns (from the main table) + self.assertEqual(table1.table.shape[1], 3) + + # Column names should match the main table + self.assertEqual(list(table1.table.columns), ["A", "B", "C"]) + + # First row should be unchanged + self.assertEqual(table1.table.iloc[0, 0], "1") + self.assertEqual(table1.table.iloc[0, 1], "2") + self.assertEqual(table1.table.iloc[0, 2], "3") + + # Second row should be from table2, padded with empty strings + self.assertEqual(table1.table.iloc[1, 0], "4") + self.assertEqual(table1.table.iloc[1, 1], "") # Padded + self.assertEqual(table1.table.iloc[1, 2], "") # Padded + + def test_table_combine_empty_tables(self): + """Test combining with empty tables.""" + table1 = Table( + page_num="1", + metadata="meta", + header="Table 1", + table_data=[["A", "B"]], + start_index=0, + end_index=100, + ) + empty_table = Table( + page_num="2", + metadata="meta", + header="", + table_data=[], + start_index=0, + end_index=100, + is_continuation=True, + ) + + original_shape = table1.table.shape + table1.combine(empty_table) + self.assertEqual(table1.table.shape, original_shape) # Should remain unchanged + + def test_convert_str_to_dataframe_malformed_data(self): + """Test handling of malformed table data.""" + malformed_inputs = [ + "[['A', 'B'], ['1']]", # Inconsistent row lengths + "[['A'], [], ['B']]", # Empty rows + "[[]]", # Empty nested list + "['not', 'nested']", # Not nested list + ] + + for malformed_input in malformed_inputs: + with self.subTest(input=malformed_input): + df = convert_str_to_dataframe(malformed_input) + self.assertIsInstance(df, pd.DataFrame) + + def test_table_post_table_text_preservation(self): + """Test that post-table text is properly preserved through operations.""" + text_dict = { + "1": "Metadata\n-------Table Start--------\nTable\n[['A', 'B'], ['1', '2']]\n-------Table End--------\nImportant footnote\nSignature line" + } + + table_dict = get_table_info(text_dict) + table = table_dict["1"][0] + + # Should capture post-table text + self.assertIn("Important footnote", table.post_table_text) + self.assertIn("Signature line", table.post_table_text) + + def test_split_large_table_preserves_metadata_and_post_text(self): + """Test that table splitting preserves all important information.""" + table = Table( + page_num="31", + metadata="EXHIBIT F", + header="Rate Table", + table_data=[["HCPC", "Rate"]] + + [[f"T{i:04d}", f"${i}.00"] for i in range(15)], + start_index=100, + end_index=2000, + post_table_text="*Rates effective 2024", + ) + + with patch("src.config.TABLE_ROW_LIMIT", 5): + chunks = split_large_table_with_headers(table) + + # All chunks should have metadata and header + for chunk in chunks: + self.assertIn("EXHIBIT F", chunk) + self.assertIn("Rate Table", chunk) + + # Only last chunk should have post-table text + self.assertIn("*Rates effective 2024", chunks[-1]) + for chunk in chunks[:-1]: + self.assertNotIn("*Rates effective 2024", chunk) + + def test_table_split(self): + """Test splitting a table.""" + long_data = [["Name", "Value"]] + [ + [f"Row{i}", str(i)] for i in range(1, 6) + ] # This table will have 5 rows + long_table = Table( + page_num="1", + metadata="metadata", + header="Long Table", + table_data=long_data, + start_index=0, + end_index=100, + ) + split_tables = long_table.split(row_limit=2) + + self.assertEqual( + len(split_tables), 3 + ) # Should split into 3 tables (5 rows split by 2 = 3) + self.assertEqual( + split_tables[0].table.shape, (2, 2) + ) # First split should have 2 rows + self.assertEqual( + split_tables[1].table.shape, (2, 2) + ) # Second split should have 2 rows + self.assertEqual( + split_tables[2].table.shape, (1, 2) + ) # Last split should have 1 row + + # Check that original metadata is preserved + for split_table in split_tables: + self.assertEqual(split_table.metadata, "metadata") + self.assertEqual(split_table.page_num, "1") + self.assertEqual(split_table.header, "Long Table") + + def test_table_split_no_split_needed(self): + """Test splitting a table that does not need splitting. It should return the original table.""" + table = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data=[["A", "B"], ["1", "2"], ["3", "4"]], + start_index=0, + end_index=100, + ) + split_tables = table.split(row_limit=10) # Row limit larger than table size + self.assertEqual(len(split_tables), 1) # Should return the original table + self.assertEqual( + split_tables[0].table.shape, (2, 2) + ) # Should be the same shape as original + self.assertEqual(split_tables[0], table) # Should be the same object + + def test_table_to_list_format(self): + """Test conversion back to list format.""" + table_data = [["A", "B"], ["1", "3"], ["2", "4"]] + table = Table( + page_num="1", + metadata="metadata", + header="Table 1", + table_data=table_data, + start_index=0, + end_index=100, + ) + list_format = table.to_list_format() + + self.assertEqual(list_format, table_data) + self.assertIsInstance(list_format, list) + self.assertIsInstance(list_format[0], list) # Should be a list of lists + + def test_table_info_new_format(self): + """Test getting table info from new list-of-rows format.""" + table_dict = get_table_info(self.sample_text_dict_new) + + self.assertIn("1", table_dict) + self.assertEqual(len(table_dict["1"]), 1) + self.assertEqual(table_dict["1"][0].header, "Table 1") + self.assertEqual(table_dict["1"][0].metadata, "Some metadata") + self.assertEqual( + table_dict["1"][0].table.shape, (2, 2) + ) # 2 rows, 2 columns in first table + self.assertEqual(table_dict["2"][0].header, "") # Second table has no header + self.assertEqual( + table_dict["2"][0].table.shape, (3, 2) + ) # 3 rows, 2 columns in second table (empty header means all data rows) + + def test_table_info_old_format(self): + """Test getting table info from old dict format.""" + with self.assertLogs(level="INFO") as log: # Expect a deprecation warning + table_dict = get_table_info(self.sample_text_dict_old) + + self.assertIn("1", table_dict) + self.assertEqual(len(table_dict["1"]), 1) + self.assertEqual(table_dict["1"][0].header, "Table 1") + self.assertEqual(table_dict["1"][0].metadata, "Some metadata") + self.assertEqual(table_dict["1"][0].table.shape, (2, 2)) + + def test_get_table_info_multiple_tables(self): + """Test getting table info from a page with multiple tables.""" + table_dict = get_table_info(self.multi_table_text) + + self.assertIn("1", table_dict) + self.assertEqual(len(table_dict["1"]), 2) + self.assertEqual(table_dict["1"][0].header, "First Table") + def test_recreate_page(self): - page_tables = self.sample_table_dict["1"] + """Test recreating page text from tables.""" + table_dict = get_table_info(self.sample_text_dict_new) + page_tables = table_dict["1"] recreated_text = recreate_page(page_tables) + self.assertIn("Table 1", recreated_text) self.assertIn("-------Table Start--------", recreated_text) self.assertIn("-------Table End--------", recreated_text) - - def test_remove_table_from_page(self): - table = self.sample_table_dict["1"][0] - updated_text = remove_table_from_page(self.sample_text_dict["1"], table) - self.assertNotIn("Table 1", updated_text) - self.assertNotIn("-------Table Start--------", updated_text) - - def test_combine_continuous_tables(self): - combined_text_dict = combine_continuous_tables(self.sample_text_dict, self.sample_table_dict) - self.assertIn("1", combined_text_dict) - self.assertNotIn("2", combined_text_dict) - self.assertIn("5", combined_text_dict["1"]) # Ensuring second table is merged - - def test_split_long_tables(self): - long_table_dict = { - "1": "Some metadata\n-------Table Start--------\nTable 1\n{'A': ['1', '2', '3', '4', '5'], 'B': ['6', '7', '8', '9', '10']}\n-------Table End--------" + self.assertIn("Some metadata", recreated_text) + + def test_split_tables_by_row_count(self): + """Test splitting tables that exceed row limit and adding .0 suffix.""" + # Create a long table that will need splitting + long_table_data = [["Name", "Value"]] + [ + [f"Row{i}", str(i)] for i in range(1, 8) + ] # 7 data rows + text_dict = { + "5": f"Table metadata\n-------Table Start--------\nLong Table\n{long_table_data}\n-------Table End--------\nPost-table text" } - table_dict = get_table_info(long_table_dict) - split_text_dict = split_long_tables(long_table_dict, table_dict, row_limit=2) - self.assertGreater(len(split_text_dict), 1) - for key, text in split_text_dict.items(): - self.assertIn("-------Table Start--------", text) - self.assertIn("-------Table End--------", text) + + result = split_tables_by_rows(text_dict, row_limit=3) + + # Should create multiple pages with .0, .1, .2 suffixes + expected_pages = ["5.0", "5.1", "5.2"] + for page in expected_pages: + self.assertIn(page, result) + + # Original page should not exist + self.assertNotIn("5", result) + + # Last page should have post-table text + self.assertIn("Post-table text", result["5.2"]) + + def test_split_tables_by_row_count_no_split_needed(self): + """Test that small tables get .0 suffix but aren't split.""" + text_dict = { + "3": "Metadata\n-------Table Start--------\nSmall Table\n[['A', 'B'], ['1', '2']]\n-------Table End--------" + } + + result = split_tables_by_rows(text_dict, row_limit=10) + + # Should have .0 suffix but no additional splits + self.assertIn("3.0", result) + self.assertNotIn("3", result) + self.assertNotIn("3.1", result) + + def test_split_tables_by_row_count_no_tables(self): + """Test that pages without tables keep original page numbers.""" + text_dict = { + "1": "Just some regular text without any tables.", + "2": "More text, no tables here either.", + } + + result = split_tables_by_rows(text_dict, row_limit=5) + + # Pages without tables should keep original numbers + self.assertIn("1", result) + self.assertIn("2", result) + self.assertNotIn("1.0", result) + self.assertNotIn("2.0", result) + + def test_combine_continuous_tables_no_continuation(self): + """Test that non-continuation tables are not combined.""" + text_dict = { + "1": "Metadata 1\n-------Table Start--------\nTable 1\n[['A', 'B'], ['1', '2']]\n-------Table End--------", + "2": "Metadata 2\n-------Table Start--------\nTable 2\n[['C', 'D'], ['3', '4']]\n-------Table End--------", + } + + table_dict = get_table_info(text_dict) + result = combine_continuous_tables(text_dict, table_dict) + + # Both pages should remain (no continuation tables to combine) + self.assertIn("1", result) + self.assertIn("2", result) + self.assertEqual(len(result), 2) + + def test_table_combine_post_table_text_accumulation(self): + """Test that post-table text is properly accumulated when combining tables.""" + table1 = Table( + page_num="1", + metadata="meta", + header="Main Table", + table_data=[["A", "B"], ["1", "2"]], + start_index=0, + end_index=100, + post_table_text="First footnote", + ) + + table2 = Table( + page_num="2", + metadata="", + header="", + table_data=[["3", "4"]], + start_index=0, + end_index=100, + is_continuation=True, + post_table_text="Second footnote", + ) + + table1.combine(table2) + + # Both footnotes should be preserved + self.assertIn("First footnote", table1.post_table_text) + self.assertIn("Second footnote", table1.post_table_text) + + def test_table_combine_column_mismatch_preserves_post_table_text(self): + """Test that post-table text is preserved when combining tables with column mismatches.""" + table1 = Table( + page_num="1", + metadata="metadata", + header="Main Table", + table_data=[["A", "B"], ["1", "2"], ["x", "y"]], # Header + 2 data rows + start_index=0, + end_index=100, + post_table_text="Main table footnote", + ) + + table2 = Table( + page_num="2", + metadata="", + header="", + table_data=[["3", "4", "5"]], # Just 1 data row (3 columns - mismatch!) + start_index=0, + end_index=100, + is_continuation=True, + post_table_text="Continuation footnote", + ) + + table1.combine(table2) + + # Verify both post-table texts are preserved despite column mismatch + self.assertIn("Main table footnote", table1.post_table_text) + self.assertIn("Continuation footnote", table1.post_table_text) + + # Verify the table combination worked correctly + self.assertEqual( + table1.table.shape[0], 3 + ) # 2 original + 1 continuation = 3 total rows + self.assertEqual(table1.table.shape[1], 2) # 2 columns (from main table) + + # Verify data integrity + self.assertEqual(table1.table.iloc[0, 0], "1") # Original data row 1 + self.assertEqual(table1.table.iloc[1, 0], "x") # Original data row 2 + self.assertEqual(table1.table.iloc[2, 0], "3") # Continuation data (truncated) + + def test_table_combine_column_mismatch_padding_preserves_post_table_text(self): + """Test post-table text preservation when padding columns.""" + table1 = Table( + page_num="1", + metadata="metadata", + header="Main Table", + table_data=[ + ["A", "B", "C"], + ["1", "2", "3"], + ["x", "y", "z"], + ], # Header + 2 data rows + start_index=0, + end_index=100, + post_table_text="Main footnote", + ) + + table2 = Table( + page_num="2", + metadata="", + header="", + table_data=[["4"]], # Just 1 data row (1 column - needs padding) + start_index=0, + end_index=100, + is_continuation=True, + post_table_text="Padded footnote", + ) + + table1.combine(table2) + + # Both footnotes should be preserved even with column padding + self.assertIn("Main footnote", table1.post_table_text) + self.assertIn("Padded footnote", table1.post_table_text) + + # Verify padding worked + self.assertEqual( + table1.table.shape, (3, 3) + ) # 2 original + 1 continuation = 3 rows, 3 columns + self.assertEqual(table1.table.iloc[0, 0], "1") # Original data row 1 + self.assertEqual(table1.table.iloc[1, 0], "x") # Original data row 2 + self.assertEqual(table1.table.iloc[2, 0], "4") # Continuation data + self.assertEqual(table1.table.iloc[2, 1], "") # Padded + self.assertEqual(table1.table.iloc[2, 2], "") # Padded + if __name__ == "__main__": unittest.main()