Merged in hotfix/prov-info-json (pull request #607)

derive IS_GROUP

* derive IS_GROUP

* Update unit tests


Approved-by: Alex Galarce
This commit is contained in:
Katon Minhas
2025-07-09 21:01:15 +00:00
parent 4d3bfe0256
commit 45db295d17
2 changed files with 19 additions and 18 deletions
+14 -11
View File
@@ -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)]
+5 -7
View File
@@ -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):