Files
doczyai-pipelines/fieldExtraction/tests/test_tin_npi_funcs.py
T
Alex Galarce 674add9a3e Merged in feature/black-and-isort (pull request #717)
Feature/black and isort

* black and isort constants

* black and isort codes

* black and isort crosswalks

* black and isort investment

* black and isort prompts

* isort testbed

* isort tracking

* black and isort utils

* poetry and black tests


Approved-by: Katon Minhas
2025-09-30 20:56:02 +00:00

345 lines
12 KiB
Python

import json
import re
from unittest.mock import MagicMock, patch
import constants.regex_patterns as regex_patterns
import pytest
import src.investment.tin_npi_funcs as tin_npi_funcs
from src.prompts.fieldset import FieldSet
class TestTinNpiFuncs:
@pytest.fixture
def sample_text_dict(self):
return {
"1": "First page with TIN 12-3456789",
"2": "Second page with NPI 1234567890",
"2.1": "Continuation with same TIN 12-3456789",
"3": "Page with both TIN 98-7654321 and NPI 9876543210",
}
def test_get_all_matches(self):
text = "TIN: 12-3456789 and another TIN: 98-7654321"
pattern = r"\d{2}[-\s]?\d{7}"
result = tin_npi_funcs.get_all_matches(text, pattern)
assert sorted(result) == sorted(["12-3456789", "98-7654321"])
# Test empty pattern
assert tin_npi_funcs.get_all_matches(text, "") == []
# Test no matches
assert tin_npi_funcs.get_all_matches("No TINs here", pattern) == []
def test_chunk_on_matches(self):
text_dict = {
"1": "Page with match1",
"2": "Page without matches",
"3": "Page with match2",
}
matches = ["match1", "match2"]
result = tin_npi_funcs.chunk_on_matches(matches, text_dict)
assert "match1" in result
assert "match2" in result
assert "without matches" not in result
def test_clean_provider_info(self):
providers = [
{
"TIN": "12-345.6789",
"NPI": "12.34567890",
"NAME": "Test Provider",
"IS_GROUP": "Y",
},
{
"TIN": "",
"NPI": "invalid",
"NAME": "",
"IS_GROUP": "N",
},
]
result = tin_npi_funcs.clean_provider_info(providers)
assert result[0]["TIN"] == "123456789"
assert result[0]["NPI"] == "1234567890"
assert result[0]["NAME"] == "Test Provider"
assert result[1]["TIN"] == "UNKNOWN"
assert result[1]["NPI"] == "UNKNOWN"
assert result[1]["NAME"] == "UNKNOWN"
def test_merge_provider_info(self):
one_to_one_results = {}
provider_info = [
{
"TIN": "123456789",
"NPI": "1234567890",
"NAME": "Group Provider",
"IS_GROUP": "Y",
},
{
"TIN": "987654321",
"NPI": "9876543210",
"NAME": "Other Provider",
"IS_GROUP": "N",
},
]
result = tin_npi_funcs.merge_provider_info(one_to_one_results, provider_info)
assert result["PROV_GROUP_TIN"] == "123456789"
assert result["PROV_OTHER_TIN"] == "987654321"
assert "Group Provider" in result["PROV_GROUP_NAME_FULL"]
assert "Other Provider" in result["PROV_OTHER_NAME_FULL"]
def test_deduplicate_providers(self):
providers = [
{"TIN": "123456789", "NPI": "1234567890", "NAME": "Provider A"},
{"TIN": "123456789", "NPI": "1234567890", "NAME": "Provider A"},
{"TIN": "UNKNOWN", "NPI": "UNKNOWN", "NAME": "UNKNOWN"},
]
result = tin_npi_funcs.deduplicate_providers(providers)
assert len(result) == 1
assert result[0]["TIN"] == "123456789"
@patch("src.utils.llm_utils.invoke_claude")
def test_get_provider_info(self, mock_invoke_claude, sample_text_dict):
mock_invoke_claude.return_value = json.dumps(
[
{
"TIN": "123456789",
"NPI": "1234567890",
"NAME": "Test Provider",
"ON_SIGNATURE_PAGE": "N",
}
]
)
result = tin_npi_funcs.get_provider_info(sample_text_dict, "1", "test.pdf")
assert len(result) == 1
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):
mock_invoke_claude.return_value = json.dumps(
[
{
"TIN": "123456789",
"NPI": "1234567890",
"NAME": "Test Provider",
"ON_SIGNATURE_PAGE": "N",
}
]
)
one_to_one_fields = FieldSet(relationship="one_to_one")
contract_text = "Contract with TIN 12-3456789"
results, fields = tin_npi_funcs.run_provider_info_fields(
contract_text, one_to_one_fields, sample_text_dict, "test.pdf"
)
assert "PROV_INFO_JSON" in results
assert "PROV_INFO_JSON_FORMATTED" in results
assert isinstance(results["PROV_INFO_JSON"], str)
def test_get_all_matches_with_ocr_exact_matches(self):
"""Test that exact matches are found and marked correctly"""
text = "TIN: 123456789 and NPI: 1234567890"
tin_results = tin_npi_funcs.get_all_matches_with_ocr(
text, r"\b\d{9}\b", r"\b[O0lI1S5G6B8gbo\d]{9}\b", "TIN"
)
assert len(tin_results) == 1
assert tin_results[0] == ("123456789", "EXACT")
def test_get_all_matches_with_ocr_corrected_matches(self):
"""Test that OCR-corrected matches work properly"""
text = "TIN: 12345678O and Provider ID: l234567890"
tin_results = tin_npi_funcs.get_all_matches_with_ocr(
text, r"\b\d{9}\b", r"\b[O0lI1S5G6B8gbo\d]{9}\b", "TIN"
)
npi_results = tin_npi_funcs.get_all_matches_with_ocr(
text, r"\b\d{10}\b", r"\b[O0lI1S5G6B8gbo\d]{10}\b", "NPI"
)
assert len(tin_results) == 1
assert tin_results[0] == ("123456780", "OCR_CORRECTED")
assert len(npi_results) == 1
assert npi_results[0] == ("1234567890", "OCR_CORRECTED")
def test_get_all_matches_with_ocr_mixed_matches(self):
"""Test that both exact and OCR matches are found in same text"""
text = "Clean TIN: 123456789 and corrupted TIN: 98765432O"
results = tin_npi_funcs.get_all_matches_with_ocr(
text, r"\b\d{9}\b", r"\b[O0lI1S5G6B8gbo\d]{9}\b", "TIN"
)
assert len(results) == 2
# Sort results to ensure consistent order
results.sort(key=lambda x: x[0])
assert results[0] == ("123456789", "EXACT")
assert results[1] == ("987654320", "OCR_CORRECTED")
def test_get_all_matches_with_ocr_no_duplicates(self):
"""Test that OCR correction doesn't create duplicates of exact matches"""
text = "TIN: 123456789 and same TIN with OCR error: 12345G789"
results = tin_npi_funcs.get_all_matches_with_ocr(
text, r"\b\d{9}\b", r"\b[O0lI1S5G6B8gbo\d]{9}\b", "TIN"
)
# Should only get one result (the exact match wins)
assert len(results) == 1
assert results[0] == ("123456789", "EXACT")
def test_attempt_ocr_correction_tin_success(self):
"""Test successful TIN OCR correction"""
result, is_valid = tin_npi_funcs.attempt_ocr_correction("12345678O", "TIN")
assert is_valid is True
assert result == "123456780"
result, is_valid = tin_npi_funcs.attempt_ocr_correction("l23456789", "TIN")
assert is_valid is True
assert result == "123456789"
def test_attempt_ocr_correction_npi_success(self):
"""Test successful NPI OCR correction"""
result, is_valid = tin_npi_funcs.attempt_ocr_correction("123456789O", "NPI")
assert is_valid is True
assert result == "1234567890"
result, is_valid = tin_npi_funcs.attempt_ocr_correction("l234567890", "NPI")
assert is_valid is True
assert result == "1234567890"
def test_attempt_ocr_correction_failure_cases(self):
"""Test cases where OCR correction should fail"""
# Wrong length after correction
result, is_valid = tin_npi_funcs.attempt_ocr_correction("12345678", "TIN")
assert is_valid is False
assert result == "12345678"
# Invalid characters that aren't in OCR substitution map
result, is_valid = tin_npi_funcs.attempt_ocr_correction("123ABC789", "TIN")
assert is_valid is False
assert result == "123ABC789"
# Unsupported identifier type
result, is_valid = tin_npi_funcs.attempt_ocr_correction("123456789", "UNKNOWN")
assert is_valid is False
assert result == "123456789"
def test_attempt_ocr_correction_formatted_tins(self):
"""Test OCR correction on formatted TINs"""
result, is_valid = tin_npi_funcs.attempt_ocr_correction("12-345678O", "TIN")
assert is_valid is True
assert result == "123456780"
result, is_valid = tin_npi_funcs.attempt_ocr_correction("123-4S-6789", "TIN")
assert is_valid is True
assert result == "123456789"
def test_attempt_ocr_correction_multiple_substitutions(self):
"""Test OCR correction with multiple character substitutions"""
result, is_valid = tin_npi_funcs.attempt_ocr_correction("l2345678O", "TIN")
assert is_valid is True
assert result == "123456780"
result, is_valid = tin_npi_funcs.attempt_ocr_correction("lS345678O", "TIN")
assert is_valid is True
assert result == "153456780"
@patch("src.utils.llm_utils.invoke_claude")
def test_run_provider_info_fields_with_ocr_quality_tracking(
self, mock_invoke_claude, sample_text_dict
):
"""Test that quality tracking fields are added to results"""
mock_invoke_claude.return_value = json.dumps(
[
{
"TIN": "123456789",
"NPI": "1234567890",
"NAME": "Test Provider",
"ON_SIGNATURE_PAGE": "N",
}
]
)
one_to_one_fields = FieldSet(relationship="one_to_one")
# Include both exact and OCR-correctable TINs
contract_text = "Contract with TIN 12-3456789 and corrupted TIN 98765432O"
results, fields = tin_npi_funcs.run_provider_info_fields(
contract_text, one_to_one_fields, sample_text_dict, "test.pdf"
)
# Check that quality tracking fields are present
assert "TIN_EXTRACTION_QUALITY" in results
assert "NPI_EXTRACTION_QUALITY" in results
# Check format of quality strings
assert "exact" in results["TIN_EXTRACTION_QUALITY"]
assert "OCR-corrected" in results["TIN_EXTRACTION_QUALITY"]
def test_ocr_regex_patterns(self):
"""Test that OCR regex patterns match expected cases"""
# Test TIN OCR pattern
tin_ocr_cases = [
"12345678O", # O at end
"O23456789", # O at start
"1234O6789", # O in middle
"12-345678O", # Formatted with O
"123-4S-6789", # Formatted with S
]
for case in tin_ocr_cases:
matches = re.findall(regex_patterns.TIN_OCR_PATTERN, case)
assert len(matches) == 1, f"TIN OCR pattern should match '{case}'"
# Test NPI OCR pattern
npi_ocr_cases = [
"123456789O", # O at end
"l234567890", # l at start
"12345O7890", # O in middle
]
for case in npi_ocr_cases:
matches = re.findall(regex_patterns.NPI_OCR_PATTERN, case)
assert len(matches) == 1, f"NPI OCR pattern should match '{case}'"
@patch("src.utils.llm_utils.invoke_claude")
def test_run_provider_info_fields_no_exact_matches_but_ocr_matches(
self, mock_invoke_claude, sample_text_dict
):
"""Test scenario where no exact matches found but OCR matches are"""
mock_invoke_claude.return_value = json.dumps(
[
{
"TIN": "123456780", # Corrected from 12345678O
"NPI": "1234567890", # Corrected from 123456789O
"NAME": "Test Provider",
"ON_SIGNATURE_PAGE": "N",
}
]
)
one_to_one_fields = FieldSet(relationship="one_to_one")
# Only OCR-corrupted identifiers, no clean ones
contract_text = "Contract with corrupted TIN 12345678O and NPI 123456789O"
results, fields = tin_npi_funcs.run_provider_info_fields(
contract_text, one_to_one_fields, sample_text_dict, "test.pdf"
)
# Should find OCR-corrected values
assert "0 exact, 1 OCR-corrected" in results["TIN_EXTRACTION_QUALITY"]
assert "0 exact, 1 OCR-corrected" in results["NPI_EXTRACTION_QUALITY"]