diff --git a/fieldExtraction/src/investment/tin_npi_funcs.py b/fieldExtraction/src/investment/tin_npi_funcs.py index d8d9ecc..6d3c510 100644 --- a/fieldExtraction/src/investment/tin_npi_funcs.py +++ b/fieldExtraction/src/investment/tin_npi_funcs.py @@ -62,7 +62,7 @@ def clean_provider_info(provider_info: list[dict]) -> list[dict]: cleaned_providers = [] for provider in provider_info: # Create a new provider dict to hold the cleaned data - cleaned_provider = {"ON_SIGNATURE_PAGE" : provider["ON_SIGNATURE_PAGE"], "ON_INTRO_PAGE" : provider["ON_INTRO_PAGE"]} + cleaned_provider = {"IS_GROUP" : provider["IS_GROUP"]} # Clean TIN - remove hyphens and dots, ensure it's 9 digits if "TIN" in provider and provider["TIN"]: @@ -116,7 +116,7 @@ def merge_provider_info(one_to_one_results: dict, provider_info: list[dict]) -> other_names = [] for provider in provider_info: - if provider.get("ON_SIGNATURE_PAGE", False) == "Y" or provider.get("ON_INTRO_PAGE", False) == "Y" : + if provider.get("IS_GROUP", False) == "Y": if provider.get("TIN"): group_tins.add(provider["TIN"]) if provider.get("NPI"): @@ -332,7 +332,6 @@ def get_provider_info(text_dict: dict, page_num : str, filename: str) -> list: Raises: ValueError: If the response from the language model cannot be parsed as valid JSON. """ - chunk = text_dict[page_num] provider_fields = FieldSet(file_path=config.FIELD_JSON_PATH, field_type="provider_info") @@ -342,7 +341,6 @@ def get_provider_info(text_dict: dict, page_num : str, filename: str) -> list: # Defensive programming for any json parsing errors try: providers = string_utils.universal_json_load(claude_answer_raw) - # Ensure we have a list of providers (not a dict or other type) if isinstance(providers, dict): # Single provider returned as dict - wrap in list @@ -350,18 +348,23 @@ def get_provider_info(text_dict: dict, page_num : str, filename: str) -> list: elif not isinstance(providers, list): # If not a list or dict, create a default list with error info print(f"Warning: Unexpected format in provider info response: {type(providers)}") - providers = [{"TIN": "UNKNOWN", "NPI": "UNKNOWN", "NAME": "PARSING_ERROR", "ON_SIGNATURE_PAGE": "N", "ON_INTRO_PAGE" : "N"}] - - # Set ON_INTRO_PAGE flag based on page number - for provider in providers: - provider["ON_INTRO_PAGE"] = "Y" if int(page_num) <= 2 else "N" + providers = [{"TIN": "UNKNOWN", "NPI": "UNKNOWN", "NAME": "PARSING_ERROR", "ON_SIGNATURE_PAGE": "N"}] + # Add IS_GROUP flag based on ON_SIGNATURE_PAGE and page_num + for provider in providers: + if int(page_num) <= 2 or provider.get("ON_SIGNATURE_PAGE", "N") == "Y": + provider["IS_GROUP"] = "Y" + else: + provider["IS_GROUP"] = "N" + if "ON_SIGNATURE_PAGE" in provider: + del provider["ON_SIGNATURE_PAGE"] + return providers except ValueError as e: # Handle JSON parsing error print(f"Error parsing JSON response from LLM: {str(e)}") print(f"Raw response: {claude_answer_raw}") # Return a default value - return [{"TIN": "UNKNOWN", "NPI": "UNKNOWN", "NAME": "PARSING_ERROR", "ON_SIGNATURE_PAGE": "N", "ON_INTRO_PAGE" : "N"}] + return [{"TIN": "UNKNOWN", "NPI": "UNKNOWN", "NAME": "PARSING_ERROR", "IS_GROUP": "N"}] def run_provider_info_fields(contract_text: str, @@ -388,7 +391,7 @@ def run_provider_info_fields(contract_text: str, # If there are no identifiers, return early with default values if not all_tins and not all_npis: - deduplicated_provider_info = [{"TIN": "UNKNOWN", "NPI": "UNKNOWN", "NAME": "NO_IDENTIFIERS_FOUND", "ON_SIGNATURE_PAGE": "N", "ON_INTRO_PAGE" : "N"}] + deduplicated_provider_info = [{"TIN": "UNKNOWN", "NPI": "UNKNOWN", "NAME": "NO_IDENTIFIERS_FOUND", "IS_GROUP" : "N"}] else: # Find pages with TINs or NPIs relevant_pages = [page_num for page_num, page_text in text_dict.items() if any(tin in page_text for tin in all_tins) or any(npi in page_text for npi in all_npis)] diff --git a/fieldExtraction/tests/test_tin_npi_funcs.py b/fieldExtraction/tests/test_tin_npi_funcs.py index 077a478..ee78715 100644 --- a/fieldExtraction/tests/test_tin_npi_funcs.py +++ b/fieldExtraction/tests/test_tin_npi_funcs.py @@ -45,15 +45,13 @@ class TestTinNpiFuncs: "TIN": "12-345.6789", "NPI": "12.34567890", "NAME": "Test Provider", - "ON_SIGNATURE_PAGE": "Y", - "ON_INTRO_PAGE": "N" + "IS_GROUP": "Y" }, { "TIN": "", "NPI": "invalid", "NAME": "", - "ON_SIGNATURE_PAGE": "N", - "ON_INTRO_PAGE": "Y" + "IS_GROUP": "N", } ] result = tin_npi_funcs.clean_provider_info(providers) @@ -73,13 +71,13 @@ class TestTinNpiFuncs: "TIN": "123456789", "NPI": "1234567890", "NAME": "Group Provider", - "ON_SIGNATURE_PAGE": "Y" + "IS_GROUP": "Y" }, { "TIN": "987654321", "NPI": "9876543210", "NAME": "Other Provider", - "ON_SIGNATURE_PAGE": "N" + "IS_GROUP": "N" } ] @@ -125,7 +123,7 @@ class TestTinNpiFuncs: result = tin_npi_funcs.get_provider_info(sample_text_dict, "1", "test.pdf") assert len(result) == 1 - assert result[0]["ON_INTRO_PAGE"] == "Y" # Page 1 should be marked as intro + assert result[0]["IS_GROUP"] == "Y" # Page 1 should be marked as intro @patch('src.utils.llm_utils.invoke_claude') def test_run_provider_info_fields(self, mock_invoke_claude, sample_text_dict):