Files
doczyai-pipelines/fieldExtraction/tests/test_dynamic.py
T
Katon Minhas c4e519894b Merged in refactor/one-time-io (pull request #648)
Refactor/one time io

* Delete client work

* Streamline imports

* Fix list class

* Update tests

* Clear code funcs test

* Fix unit tests

* Single read code funcs

* Single-read model

* Single-load embeddings

* Successful E2E test

* update unit tests

* Update importsg

* Reload poetry.lock

* remove old exhibit header function

* Remove aarete_derived generic function

* remove align and format tables

* Remove strings to dict

* references

* Clear code

* Move preprocessing_funcs

* remove keywords

* refactor postprocessing_funcs

* Pass unit test - remove qa_qc directory

* Black and isort

* remove print


Approved-by: Alex Galarce
2025-08-05 20:46:19 +00:00

176 lines
6.5 KiB
Python

import unittest
from unittest.mock import MagicMock, patch
import src.investment.one_to_n_funcs as one_to_n_funcs
from constants.constants import Constants
class TestOneToN(unittest.TestCase):
def setUp(self):
self.exhibit_text = "Sample exhibit text with some content"
self.filename = "test_file.pdf"
self.constants = Constants()
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.constants.VALID_CARVEOUTS,
self.constants.VALID_SPECIAL,
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.one_to_n_funcs.validate_reimbursements_for_llm")
def test_reimbursement_level(
self, mock_validate_llm, mock_special_cases, mock_reimbursement_primary
):
# Setup
mock_reimbursement_primary.return_value = [
{"SERVICE_TERM": "Test Service", "REIMB_TERM": "term1$"}
]
mock_special_cases.return_value = [
{
"SERVICE_TERM": "Test Service",
"REIMB_TERM": "term1$",
"CARVEOUT_IND": "N",
}
]
mock_validate_llm.return_value = True # Mock all entries pass LLM validation
mock_fields = MagicMock()
exhibit_page = "1"
seen_pairs = set()
# Execute
result = one_to_n_funcs.reimbursement_level(
self.exhibit_text,
self.filename,
mock_fields,
exhibit_page,
seen_pairs,
exhibit_lesser_of="N/A",
constants=self.constants,
)
# Assert
self.assertEqual(len(result), 1)
mock_reimbursement_primary.assert_called_once()
mock_special_cases.assert_called_once()
if __name__ == "__main__":
unittest.main()