724e31b15d
Feature/filter reimbursements * First-pass implementation for reimbursement filtering * Enhance reimbursement level processing by filtering out services without reimbursement terms and logging when no valid pairs are found. * Update mock return values in reimbursement level tests to reflect new reimbursement terms * Merge remote-tracking branch 'origin/main' into feature/split-and-filter-reimbursements * Refine reimbursement filtering by adding 'billed' to indicators and ensuring REIMB_TERM is a string before processing. * Refactor is_empty function signature to support multiple input types and clarify return values * adding LLM review for ambiguous reimbursement methodology cases * more stringent keywords for stage 1 of reimbursement filtering * refining rate and cost patterns for reimbursement filtering * Enhance LLM response handling in reimbursement methodology check to improve accuracy and logging for ambiguous cases. * Merge remote-tracking branch 'origin/main' into feature/split-and-filter-reimbursements * Update reimbursement test cases to reflect accurate reimbursement terms and improve deduplication logic * move VALIDATE_REIMBURSEMENTS_PROMPT to investment_prompts.py * black, isort formatting * Remove IDENTIFY_REIMBURSEMENT_EXHIBITS_PROMPT function (unused) * Add MODEL_ID_CLAUDE4_SONNET for new model integration (doesn't work right now) * pipe delimited-output and parsing * Merge remote-tracking branch 'origin/main' into feature/split-and-filter-reimbursements * Refactor reimbursement filtering by moving clear indicators to constants Approved-by: Katon Minhas
160 lines
6.4 KiB
Python
160 lines
6.4 KiB
Python
import unittest
|
|
from unittest.mock import patch, MagicMock
|
|
|
|
import src.investment.one_to_n_funcs as one_to_n_funcs
|
|
from src.prompts.investment_prompts import FieldSet, Field
|
|
import src.config as config
|
|
|
|
class TestOneToN(unittest.TestCase):
|
|
def setUp(self):
|
|
self.exhibit_text = "Sample exhibit text with some content"
|
|
self.filename = "test_file.pdf"
|
|
|
|
def test_get_exhibit_level_answers(self):
|
|
# Save original functions to restore later
|
|
original_fieldset = one_to_n_funcs.FieldSet
|
|
original_exhibit_level = one_to_n_funcs.EXHIBIT_LEVEL
|
|
original_invoke_claude = one_to_n_funcs.llm_utils.invoke_claude
|
|
original_json_load = one_to_n_funcs.string_utils.universal_json_load
|
|
|
|
try:
|
|
# Create mocks
|
|
mock_fieldset = MagicMock()
|
|
mock_fieldset_instance = MagicMock()
|
|
mock_fieldset.return_value = mock_fieldset_instance
|
|
mock_exhibit_level = MagicMock(return_value="test prompt")
|
|
mock_invoke_claude = MagicMock(return_value="raw json")
|
|
mock_json_load = MagicMock(return_value={"field1": "value1", "field2": "value2"})
|
|
|
|
# Replace functions with mocks
|
|
one_to_n_funcs.FieldSet = mock_fieldset
|
|
one_to_n_funcs.EXHIBIT_LEVEL = mock_exhibit_level
|
|
one_to_n_funcs.llm_utils.invoke_claude = mock_invoke_claude
|
|
one_to_n_funcs.string_utils.universal_json_load = mock_json_load
|
|
|
|
# Test Case 1: contains_fields() returns False
|
|
mock_fieldset_instance.contains_fields.return_value = False
|
|
result = one_to_n_funcs.get_exhibit_level_answers(self.exhibit_text, self.filename)
|
|
self.assertEqual(result, {})
|
|
|
|
# Test Case 2: contains_fields() returns True
|
|
mock_fieldset_instance.contains_fields.return_value = True
|
|
result = one_to_n_funcs.get_exhibit_level_answers(self.exhibit_text, self.filename)
|
|
self.assertEqual(result, {"field1": "value1", "field2": "value2"})
|
|
|
|
finally:
|
|
# Restore original functions
|
|
one_to_n_funcs.FieldSet = original_fieldset
|
|
one_to_n_funcs.EXHIBIT_LEVEL = original_exhibit_level
|
|
one_to_n_funcs.llm_utils.invoke_claude = original_invoke_claude
|
|
one_to_n_funcs.string_utils.universal_json_load = original_json_load
|
|
|
|
|
|
@patch('src.utils.llm_utils.invoke_claude')
|
|
def test_get_reimbursement_primary(self, mock_invoke_claude):
|
|
# Setup
|
|
mock_fields = MagicMock()
|
|
mock_fields.get_prompt_dict.return_value = {"field1": "prompt1"}
|
|
mock_invoke_claude.return_value = '{"answer1": "value1"}'
|
|
|
|
# Execute
|
|
result = one_to_n_funcs.get_reimbursement_primary(
|
|
mock_fields, self.exhibit_text, self.filename
|
|
)
|
|
|
|
# Assert
|
|
self.assertEqual(result, {"answer1": "value1"})
|
|
mock_invoke_claude.assert_called_once()
|
|
|
|
@patch('src.utils.llm_utils.invoke_claude')
|
|
def test_get_special_cases(self, mock_invoke_claude):
|
|
# Setup
|
|
test_answers = [{
|
|
"SERVICE_TERM": "service1",
|
|
"REIMB_TERM": "reimbursement1"
|
|
}]
|
|
mock_invoke_claude.return_value = "|N/A|"
|
|
|
|
# Execute
|
|
result = one_to_n_funcs.get_special_cases(test_answers, self.filename)
|
|
|
|
# Assert
|
|
self.assertEqual(len(result), 1)
|
|
self.assertEqual(result[0]["CARVEOUT_CD"], "N/A")
|
|
self.assertEqual(result[0]["CARVEOUT_IND"], "N")
|
|
self.assertEqual(result[0]["DEFAULT_IND"], "N")
|
|
|
|
@patch('src.utils.llm_utils.invoke_claude')
|
|
def test_methodology_breakout_single_row(self, mock_invoke_claude):
|
|
# Setup
|
|
test_answer = {
|
|
"SERVICE_TERM": "service1",
|
|
"REIMB_TERM": "Fee Schedule"
|
|
}
|
|
mock_invoke_claude.side_effect = [
|
|
'{"AARETE_DERIVED_REIMB_METHOD": "Fee Schedule", "REIMB_PCT_RATE": "80%"}',
|
|
'{"FEE_SCHEDULE": "TEST", "AARETE_DERIVED_FEE_SCHEDULE": "TEST", "AARETE_DERIVED_FEE_SCHEDULE_VERSION": "1.0"}'
|
|
]
|
|
|
|
# Execute
|
|
result = one_to_n_funcs.methodology_breakout_single_row(
|
|
test_answer,
|
|
{"methodology": "questions"},
|
|
{"fs": "questions"},
|
|
self.filename
|
|
)
|
|
|
|
# Assert
|
|
self.assertEqual(len(result), 1)
|
|
self.assertEqual(result[0]["AARETE_DERIVED_REIMB_METHOD"], "Fee Schedule")
|
|
self.assertEqual(result[0]["FEE_SCHEDULE"], "TEST")
|
|
|
|
def test_combine_one_to_n_answers(self):
|
|
# Setup
|
|
exhibit_level = {"EXHIBIT_TITLE": "Test Exhibit"}
|
|
reimbursement_level = [
|
|
{"REIMB_TERM": "term1"},
|
|
{"REIMB_TERM": "term2"}
|
|
]
|
|
tin_npi = {"TIN": "123456789"}
|
|
|
|
# Execute
|
|
result = one_to_n_funcs.combine_one_to_n_answers(
|
|
exhibit_level, reimbursement_level, tin_npi
|
|
)
|
|
|
|
# Assert
|
|
self.assertEqual(len(result), 2)
|
|
self.assertTrue(all("EXHIBIT_TITLE" in d for d in result))
|
|
self.assertTrue(all("TIN" in d for d in result))
|
|
self.assertTrue(all("REIMB_TERM" in d for d in result))
|
|
|
|
@patch('src.investment.one_to_n_funcs.get_reimbursement_primary')
|
|
@patch('src.investment.one_to_n_funcs.get_special_cases')
|
|
@patch('src.investment.code_funcs.get_code_breakout')
|
|
def test_reimbursement_level(self, mock_code_breakout,
|
|
mock_special_cases, mock_reimbursement_primary):
|
|
# Setup
|
|
mock_reimbursement_primary.return_value = [{"REIMB_TERM": "term1$"}]
|
|
mock_special_cases.return_value = [{"REIMB_TERM": "term1$", "CARVEOUT_IND": "N"}]
|
|
mock_code_breakout.return_value = [{"REIMB_TERM": "term1$", "CODE": "code1"}]
|
|
|
|
mock_fields = MagicMock()
|
|
mock_dataset = {}
|
|
exhibit_page = "1"
|
|
seen_pairs = set()
|
|
|
|
# Execute
|
|
result = one_to_n_funcs.reimbursement_level(
|
|
self.exhibit_text, self.filename, mock_fields, mock_dataset, exhibit_page, seen_pairs
|
|
)
|
|
|
|
# Assert
|
|
self.assertEqual(len(result), 1)
|
|
self.assertEqual(result[0]["CODE"], "code1")
|
|
mock_reimbursement_primary.assert_called_once()
|
|
mock_special_cases.assert_called_once()
|
|
mock_code_breakout.assert_called_once()
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main() |