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" # Execute result = one_to_n_funcs.reimbursement_level( self.exhibit_text, self.filename, mock_fields, mock_dataset, exhibit_page ) # 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()