Files
doczyai-pipelines/fieldExtraction/tests/investment_postprocess_test.py
T
Katon Minhas 1c15489370 Merged in feature/rename-proc-cd (pull request #645)
rename columns function

* rename columns function

* update columns


Approved-by: Alex Galarce
2025-08-04 16:10:26 +00:00

447 lines
19 KiB
Python

from unittest.mock import patch, MagicMock
import unittest
from unittest.mock import patch
import pytest
import pandas as pd
from src.investment.investment_postprocessing_funcs import (
normalize_indicator_field,
format_rate_fields_with_commas,
remove_hyphens,
flatten_singleton_string_list,
validate_and_reformat_date,
date_postprocess,
normalize_auto_renewal_term,
normalize_cpt_fields,
process_patient_age_range,
remove_redundant_reimb_info,
deduplicate_provider_columns,
rename_columns
)
class TestPostprocessFunctions(unittest.TestCase):
def test_normalize_indicator_field(self):
self.assertEqual(normalize_indicator_field("Y"), "Y")
self.assertEqual(normalize_indicator_field("y"), "Y")
self.assertEqual(normalize_indicator_field("N"), "N")
self.assertEqual(normalize_indicator_field(""), "N")
self.assertEqual(normalize_indicator_field(None), "N")
def test_format_rate_fields_with_commas(self):
# Test numeric values
self.assertEqual(format_rate_fields_with_commas("1234.567"), "1,234.57")
self.assertEqual(format_rate_fields_with_commas("1000"), "1,000.00")
self.assertEqual(format_rate_fields_with_commas(1234.567), "1,234.57")
self.assertEqual(format_rate_fields_with_commas(1000), "1,000.00")
def test_remove_hyphens(self):
self.assertEqual(remove_hyphens("123-45-6789"), "123456789")
self.assertEqual(remove_hyphens("123-456"), "123456")
self.assertEqual(remove_hyphens(None), "")
self.assertEqual(remove_hyphens(""), "")
def test_flatten_singleton_string_list(self):
self.assertEqual(flatten_singleton_string_list("['123']"), "123")
self.assertEqual(flatten_singleton_string_list("['123', '456']"), "123, 456")
self.assertEqual(flatten_singleton_string_list("11"), "11")
self.assertEqual(flatten_singleton_string_list("invalid"), "invalid")
self.assertEqual(flatten_singleton_string_list(None), "")
def test_rename_columns(self):
"""Tests the rename_columns function to ensure it correctly renames specified columns.
Tests:
1. Basic column renaming from PROCEDURE_CD to CPT4_PROC_CD
2. Multiple columns being renamed
3. Handling of columns not in the rename map
4. Empty DataFrame
"""
# Test case 1: Basic column renaming
input_df1 = pd.DataFrame({
"PROCEDURE_CD": ["12345", "67890"],
"PROCEDURE_CD_DESC": ["Test Procedure", "Another Procedure"],
"OTHER_COLUMN": ["value1", "value2"]
})
expected_df1 = pd.DataFrame({
"CPT4_PROC_CD": ["12345", "67890"],
"CPT4_PROC_CD_DESC": ["Test Procedure", "Another Procedure"],
"OTHER_COLUMN": ["value1", "value2"]
})
result_df1 = rename_columns(input_df1)
pd.testing.assert_frame_equal(result_df1, expected_df1)
# Test case 2: Only some columns need renaming
input_df2 = pd.DataFrame({
"PROCEDURE_CD": ["12345", "67890"],
"OTHER_COLUMN": ["value1", "value2"]
})
expected_df2 = pd.DataFrame({
"CPT4_PROC_CD": ["12345", "67890"],
"OTHER_COLUMN": ["value1", "value2"]
})
result_df2 = rename_columns(input_df2)
pd.testing.assert_frame_equal(result_df2, expected_df2)
# Test case 3: None of the columns need renaming
input_df3 = pd.DataFrame({
"COLUMN_A": ["a", "b"],
"COLUMN_B": ["c", "d"]
})
result_df3 = rename_columns(input_df3)
pd.testing.assert_frame_equal(result_df3, input_df3) # Should be unchanged
# Test case 4: Empty DataFrame
empty_df = pd.DataFrame()
result_empty_df = rename_columns(empty_df)
pd.testing.assert_frame_equal(result_empty_df, empty_df) # Should be unchanged
@patch("src.utils.llm_utils.invoke_claude")
@patch("src.investment.investment_postprocessing_funcs.invoke_derived_term_date")
@patch("src.prompts.investment_prompts.FieldSet.load_from_file") # Patch this specific method
def test_date_postprocess(self, mock_load_from_file, mock_derived_term_date, mock_invoke_claude):
"""Tests the date_postprocess function with mocked dependencies.
This test verifies that:
1. Date fields are standardized to YYYYMMDD format using the LLM (mocked)
2. The derived termination date is properly calculated and added to the DataFrame
The test uses multiple patches to isolate the function from external dependencies:
- Mocks `invoke_claude` to return a fixed date string ("|20230101|")
- Mocks `invoke_derived_term_date` to return a fixed termination date ("20231231")
- Mocks `FieldSet.load_from_file` and `FieldSet.filter` to avoid file operations
and provide controlled field definitions
The test checks that the date fields in the DataFrame are correctly formatted and
that the derived termination date is added as expected.
"""
# Setup field mock
field_mocks = [
MagicMock(field_name="CONTRACT_EFFECTIVE_DT"),
MagicMock(field_name="CONTRACT_TERMINATION_DT"),
MagicMock(field_name="TERMINATION_DT")
]
# Configure the mock methods
mock_invoke_claude.return_value = "|20230101|"
mock_derived_term_date.return_value = "20231231"
df = pd.DataFrame({
"CONTRACT_EFFECTIVE_DT": ["2023-01-01", "01/01/2023", "invalid"],
"CONTRACT_TERMINATION_DT": ["2023-12-31", "12/31/2023", "invalid"],
"TERMINATION_DT": ["2023-12-31", "12/31/2023", "invalid"],
"AARETE_DERIVED_EFFECTIVE_DT": ["2023-01-01", "2023-01-01", "2023-01-01"]
})
with patch("src.prompts.investment_prompts.FieldSet.filter") as mock_filter:
# Configure the filter to return a mock object with fields attribute
mock_filter.return_value = MagicMock(fields=field_mocks)
result_df = date_postprocess(df, "mock_path")
assert result_df["CONTRACT_EFFECTIVE_DT"].tolist() == ["20230101", "20230101", "20230101"]
assert result_df["AARETE_DERIVED_TERMINATION_DT"].tolist() == ["20231231", "20231231", "20231231"]
def test_normalize_auto_renewal_term(self):
self.assertEqual(normalize_auto_renewal_term("12 months"), "1 year")
self.assertEqual(normalize_auto_renewal_term("one year"), "1 year")
self.assertEqual(normalize_auto_renewal_term("month to month"), "1 month")
self.assertEqual(normalize_auto_renewal_term("(12) 12 months"), "1 year")
self.assertEqual(normalize_auto_renewal_term(None), "")
def test_normalize_cpt_fields(self):
self.assertEqual(normalize_cpt_fields("[123, 456]"), "['123', '456']")
self.assertEqual(normalize_cpt_fields("123-456"), "['123-456']")
self.assertEqual(normalize_cpt_fields("['T0000-T9999, S0000-S9999']"), "['T0000-T9999', 'S0000-S9999']")
self.assertEqual(normalize_cpt_fields("123"), "['123']")
self.assertEqual(normalize_cpt_fields(None), "")
self.assertEqual(normalize_cpt_fields(123), "['123']")
def test_process_patient_age_range(self):
"""Tests the process_patient_age_range function with various age range formats.
Tests:
1. Standard hyphenated ranges (e.g., "0-18")
2. Single age values
3. Text descriptions with "to"
4. Text descriptions with "and under"
5. Special cases like "newborn"
6. Empty/None values
7. Invalid formats
8. Missing PATIENT_AGE_RANGE column
"""
# Test case 1: DataFrame with PATIENT_AGE_RANGE column
input_df = pd.DataFrame({
"PATIENT_AGE_RANGE": [
"0-18", # Standard hyphenated range
"21", # Single age
]
})
result_df = process_patient_age_range(input_df)
# Verify columns
self.assertIn("PATIENT_AGE_MIN", result_df.columns)
self.assertIn("PATIENT_AGE_MAX", result_df.columns)
self.assertNotIn("PATIENT_AGE_RANGE", result_df.columns)
# Expected values
expected_min = ["0", "21"]
expected_max = ["18", "21"]
# Check transformations
pd.testing.assert_series_equal(
result_df["PATIENT_AGE_MIN"],
pd.Series(expected_min, name="PATIENT_AGE_MIN"),
check_dtype=False
)
pd.testing.assert_series_equal(
result_df["PATIENT_AGE_MAX"],
pd.Series(expected_max, name="PATIENT_AGE_MAX"),
check_dtype=False
)
# Test case 2: DataFrame without PATIENT_AGE_RANGE column
input_df_no_age = pd.DataFrame({
"OTHER_COLUMN": ["value1", "value2"]
})
result_df_no_age = process_patient_age_range(input_df_no_age)
# Verify the DataFrame is unchanged
pd.testing.assert_frame_equal(input_df_no_age, result_df_no_age)
# Test case 3: Empty DataFrame
empty_df = pd.DataFrame()
result_empty_df = process_patient_age_range(empty_df)
# Verify empty DataFrame is unchanged
pd.testing.assert_frame_equal(empty_df, result_empty_df)
def test_validate_and_reformat_date(self):
"""Tests validate_and_reformat_date with various date formats.
Tests:
1. Date already in YYYY/MM/DD format
2. Common alternative formats (YYYY-MM-DD, MM/DD/YYYY, etc.)
3. Invalid date formats
4. None and non-string values
"""
# Date already in correct format
self.assertEqual(validate_and_reformat_date("2023/01/15"), "2023/01/15")
# Test various date formats that should be reformatted
self.assertEqual(validate_and_reformat_date("2023-01-15"), "2023/01/15")
self.assertEqual(validate_and_reformat_date("01/15/2023"), "2023/01/15")
self.assertEqual(validate_and_reformat_date("15-Jan-2023"), "2023/01/15")
# Test invalid formats - should return the original string
self.assertEqual(validate_and_reformat_date("Invalid date"), "Invalid date")
self.assertEqual(validate_and_reformat_date("01-15"), "01-15")
# Test None and non-string values
self.assertEqual(validate_and_reformat_date(None), None)
self.assertEqual(validate_and_reformat_date(12345), 12345)
def test_remove_redundant_reimb_info(self):
"""Tests remove_redundant_reimb_date function.
Tests:
1. When reimbursement dates match derived dates (should remove)
2. When reimbursement dates differ from derived dates (should keep)
3. When only some rows match (should remove only matching rows)
4. When columns are missing (should return unchanged DataFrame)
"""
# Test case 1: When dates match (should remove)
input_df1 = pd.DataFrame({
"REIMB_EFFECTIVE_DT": ["2023/01/01", "2023/02/01"],
"REIMB_TERMINATION_DT": ["2023/12/31", "2023/12/31"],
"AARETE_DERIVED_EFFECTIVE_DT": ["2023/01/01", "2023/02/01"],
"AARETE_DERIVED_TERMINATION_DT": ["2023/12/31", "2023/12/31"]
})
expected_df1 = pd.DataFrame({
"REIMB_EFFECTIVE_DT": ["", ""],
"REIMB_TERMINATION_DT": ["", ""],
"AARETE_DERIVED_EFFECTIVE_DT": ["2023/01/01", "2023/02/01"],
"AARETE_DERIVED_TERMINATION_DT": ["2023/12/31", "2023/12/31"]
})
result_df1 = remove_redundant_reimb_info(input_df1)
pd.testing.assert_frame_equal(result_df1, expected_df1)
# Test case 2: When dates differ (should keep)
input_df2 = pd.DataFrame({
"REIMB_EFFECTIVE_DT": ["2023/01/15", "2023/02/15"],
"REIMB_TERMINATION_DT": ["2023/12/15", "2023/12/15"],
"AARETE_DERIVED_EFFECTIVE_DT": ["2023/01/01", "2023/02/01"],
"AARETE_DERIVED_TERMINATION_DT": ["2023/12/31", "2023/12/31"]
})
# Result should be unchanged
result_df2 = remove_redundant_reimb_info(input_df2)
pd.testing.assert_frame_equal(result_df2, input_df2)
# Test case 3: Mixed case - some match, some don't
input_df3 = pd.DataFrame({
"REIMB_EFFECTIVE_DT": ["2023/01/01", "2023/02/15"],
"REIMB_TERMINATION_DT": ["2023/12/31", "2023/12/15"],
"AARETE_DERIVED_EFFECTIVE_DT": ["2023/01/01", "2023/02/01"],
"AARETE_DERIVED_TERMINATION_DT": ["2023/12/31", "2023/12/31"]
})
expected_df3 = pd.DataFrame({
"REIMB_EFFECTIVE_DT": ["", "2023/02/15"],
"REIMB_TERMINATION_DT": ["", "2023/12/15"],
"AARETE_DERIVED_EFFECTIVE_DT": ["2023/01/01", "2023/02/01"],
"AARETE_DERIVED_TERMINATION_DT": ["2023/12/31", "2023/12/31"]
})
result_df3 = remove_redundant_reimb_info(input_df3)
pd.testing.assert_frame_equal(result_df3, expected_df3)
# Test case 4: Missing columns
input_df4 = pd.DataFrame({
"REIMB_EFFECTIVE_DT": ["2023/01/01", "2023/02/01"],
"OTHER_COLUMN": ["value1", "value2"]
})
# Result should be unchanged
result_df4 = remove_redundant_reimb_info(input_df4)
pd.testing.assert_frame_equal(result_df4, input_df4)
def test_deduplicate_provider_columns(self):
"""Tests deduplicate_provider_columns function.
Tests:
1. Basic deduplication - removes GROUP values from OTHER fields
2. Internal deduplication - removes duplicates within OTHER fields
3. Combined scenario - both GROUP removal and internal deduplication
4. Empty/missing values handling
5. Missing columns - should return unchanged DataFrame
6. Empty DataFrame
"""
# Test case 1: Basic GROUP removal
input_df1 = pd.DataFrame({
"PROV_GROUP_TIN": ["123456789"],
"PROV_GROUP_NPI": ["1234567890"],
"PROV_GROUP_NAME_FULL": ["Main Hospital"],
"PROV_OTHER_TIN": ["123456789|987654321"],
"PROV_OTHER_NPI": ["1234567890|0987654321"],
"PROV_OTHER_NAME_FULL": ["Main Hospital|Other Clinic"]
})
expected_df1 = pd.DataFrame({
"PROV_GROUP_TIN": ["123456789"],
"PROV_GROUP_NPI": ["1234567890"],
"PROV_GROUP_NAME_FULL": ["Main Hospital"],
"PROV_OTHER_TIN": ["987654321"],
"PROV_OTHER_NPI": ["0987654321"],
"PROV_OTHER_NAME_FULL": ["Other Clinic"]
})
result_df1 = deduplicate_provider_columns(input_df1)
pd.testing.assert_frame_equal(result_df1, expected_df1)
# Test case 2: Internal deduplication (your original example)
input_df2 = pd.DataFrame({
"PROV_GROUP_TIN": ["061798267"],
"PROV_GROUP_NPI": ["1111111111"],
"PROV_GROUP_NAME_FULL": ["Group Practice"],
"PROV_OTHER_TIN": ["061798267|061798267|UNKNOWN|UNKNOWN|061992277|061798267"],
"PROV_OTHER_NPI": ["1111111111|2222222222|2222222222|UNKNOWN"],
"PROV_OTHER_NAME_FULL": ["Group Practice|Other Practice|Other Practice|UNKNOWN"]
})
expected_df2 = pd.DataFrame({
"PROV_GROUP_TIN": ["061798267"],
"PROV_GROUP_NPI": ["1111111111"],
"PROV_GROUP_NAME_FULL": ["Group Practice"],
"PROV_OTHER_TIN": ["061992277"],
"PROV_OTHER_NPI": ["2222222222"],
"PROV_OTHER_NAME_FULL": ["Other Practice"]
})
result_df2 = deduplicate_provider_columns(input_df2)
pd.testing.assert_frame_equal(result_df2, expected_df2)
# Test case 3: Empty OTHER fields after deduplication
input_df3 = pd.DataFrame({
"PROV_GROUP_TIN": ["123456789"],
"PROV_GROUP_NPI": ["1234567890"],
"PROV_GROUP_NAME_FULL": ["Main Hospital"],
"PROV_OTHER_TIN": ["123456789|123456789|UNKNOWN"],
"PROV_OTHER_NPI": ["1234567890|UNKNOWN|UNKNOWN"],
"PROV_OTHER_NAME_FULL": ["Main Hospital|UNKNOWN"]
})
expected_df3 = pd.DataFrame({
"PROV_GROUP_TIN": ["123456789"],
"PROV_GROUP_NPI": ["1234567890"],
"PROV_GROUP_NAME_FULL": ["Main Hospital"],
"PROV_OTHER_TIN": [""],
"PROV_OTHER_NPI": [""],
"PROV_OTHER_NAME_FULL": [""]
})
result_df3 = deduplicate_provider_columns(input_df3)
pd.testing.assert_frame_equal(result_df3, expected_df3)
# Test case 4: Multiple rows
input_df4 = pd.DataFrame({
"PROV_GROUP_TIN": ["111111111", "222222222"],
"PROV_GROUP_NPI": ["1111111111", "2222222222"],
"PROV_GROUP_NAME_FULL": ["Hospital A", "Hospital B"],
"PROV_OTHER_TIN": ["111111111|333333333", "444444444|222222222|444444444"],
"PROV_OTHER_NPI": ["3333333333|1111111111", "4444444444|2222222222"],
"PROV_OTHER_NAME_FULL": ["Clinic C|Hospital A", "Clinic D|Hospital B"]
})
expected_df4 = pd.DataFrame({
"PROV_GROUP_TIN": ["111111111", "222222222"],
"PROV_GROUP_NPI": ["1111111111", "2222222222"],
"PROV_GROUP_NAME_FULL": ["Hospital A", "Hospital B"],
"PROV_OTHER_TIN": ["333333333", "444444444"],
"PROV_OTHER_NPI": ["3333333333", "4444444444"],
"PROV_OTHER_NAME_FULL": ["Clinic C", "Clinic D"]
})
result_df4 = deduplicate_provider_columns(input_df4)
pd.testing.assert_frame_equal(result_df4, expected_df4)
# Test case 5: Missing columns - should return unchanged
input_df5 = pd.DataFrame({
"PROV_GROUP_TIN": ["123456789"],
"OTHER_COLUMN": ["value1"]
})
result_df5 = deduplicate_provider_columns(input_df5)
pd.testing.assert_frame_equal(result_df5, input_df5)
# Test case 6: Empty DataFrame
empty_df = pd.DataFrame()
result_empty_df = deduplicate_provider_columns(empty_df)
pd.testing.assert_frame_equal(empty_df, result_empty_df)
# Test case 7: Empty OTHER fields (already empty strings)
input_df7 = pd.DataFrame({
"PROV_GROUP_TIN": ["123456789"],
"PROV_GROUP_NPI": ["1234567890"],
"PROV_GROUP_NAME_FULL": ["Main Hospital"],
"PROV_OTHER_TIN": [""],
"PROV_OTHER_NPI": [""],
"PROV_OTHER_NAME_FULL": [""]
})
result_df7 = deduplicate_provider_columns(input_df7)
pd.testing.assert_frame_equal(result_df7, input_df7)
if __name__ == "__main__":
unittest.main()