diff --git a/fieldExtraction/tests/llm_utils_test.py b/fieldExtraction/tests/llm_utils_test.py index b6bbc85..4365eb2 100644 --- a/fieldExtraction/tests/llm_utils_test.py +++ b/fieldExtraction/tests/llm_utils_test.py @@ -1,166 +1,64 @@ -import pytest -from src.utils.llm_utils import ( - count_tokens, - create_cost_log_entry, - update_stats, - invoke_claude, -) +import unittest +from unittest.mock import patch, mock_open +import hashlib +import csv +from collections import OrderedDict +import src.utils.llm_utils as llm_utils from src import config -class TestLLMUtils: - # Fixture to reset global stats before each test - @pytest.fixture(autouse=True) - def reset_global_stats(self): - """Fixture to reset global stats before each test.""" - config.MODEL_STATS = { - "file1.txt": {"requests": 0, "tokens_sent": 0, "tokens_received": 0, "total_time": 0}, - "file2.txt": {"requests": 0, "tokens_sent": 0, "tokens_received": 0, "total_time": 0}, - "file3.txt": {"requests": 0, "tokens_sent": 0, "tokens_received": 0, "total_time": 0}, - "file4.txt": {"requests": 0, "tokens_sent": 0, "tokens_received": 0, "total_time": 0}, - } - config.GLOBAL_STATS = { - "total_requests": 0, - "total_tokens_sent": 0, - "total_tokens_received": 0, - "total_time": 0, - } +class TestLLMUtils(unittest.TestCase): + def setUp(self): + self.filename = "test_file.txt" + self.prompt = "This is a test prompt." + self.response = "This is a test response." + self.model_type = "test_model" + self.input_cost_per_1k = 0.01 + self.output_cost_per_1k = 0.02 + self.mock_stats = {self.filename: {"requests": 0, "tokens_sent": 0, "tokens_received": 0, "total_time": 0}} + config.MODEL_STATS = self.mock_stats + config.GLOBAL_STATS = {} - # Test cases for count_tokens - @pytest.mark.parametrize("input_str, expected_tokens", [ - ("hello world", 2), # 11 characters / 6 = 1.83 → ceil to 2 - ("", 0), # Empty string - ("a" * 12, 2), # 12 characters / 6 = 2 - ("a" * 13, 3), # 13 characters / 6 = 2.16 → ceil to 3 - ("😊", 1), # Unicode characters - ]) - def test_count_tokens(self, input_str, expected_tokens): - """Test the count_tokens function.""" - assert count_tokens(input_str) == expected_tokens - - # Test cases for create_cost_log_entry - @pytest.mark.parametrize("filename, caller_name, prompt, response, input_cost_per_1k, output_cost_per_1k, expected_output", [ - ( - "file1.txt", "func1", "hello", "world", 0.1, 0.2, - { - "Filename": "file1.txt", - "Caller Function": "func1", - "Input Prompt Length": 5, - "Input Prompt Tokens": 1, # 5 chars / 6 = 0.83 → ceil to 1 - "Input Cost": 0.1 * 1 / 1000, - "Output Response Length": 5, - "Output Response Tokens": 1, # 5 chars / 6 = 0.83 → ceil to 1 - "Output Cost": 0.2 * 1 / 1000, - "Total Cost": (0.1 * 1 / 1000) + (0.2 * 1 / 1000), - } - ), - ( - "file2.txt", "func2", "", "", 0.1, 0.2, - { - "Filename": "file2.txt", - "Caller Function": "func2", - "Input Prompt Length": 0, - "Input Prompt Tokens": 0, # Empty string - "Input Cost": 0.1 * 0 / 1000, - "Output Response Length": 0, - "Output Response Tokens": 0, # Empty string - "Output Cost": 0.2 * 0 / 1000, - "Total Cost": (0.1 * 0 / 1000) + (0.2 * 0 / 1000), - } - ), - ( - "file3.txt", "func3", "a" * 12, "b" * 13, 0.1, 0.2, - { - "Filename": "file3.txt", - "Caller Function": "func3", - "Input Prompt Length": 12, - "Input Prompt Tokens": 2, # 12 chars / 6 = 2 - "Input Cost": 0.1 * 2 / 1000, - "Output Response Length": 13, - "Output Response Tokens": 3, # 13 chars / 6 = 2.16 → ceil to 3 - "Output Cost": 0.2 * 3 / 1000, - "Total Cost": (0.1 * 2 / 1000) + (0.2 * 3 / 1000), - } - ), - ( - "file4.txt", "func4", "😊", "😊😊", 0.1, 0.2, - { - "Filename": "file4.txt", - "Caller Function": "func4", - "Input Prompt Length": 1, - "Input Prompt Tokens": 1, # Unicode character - "Input Cost": 0.1 * 1 / 1000, - "Output Response Length": 2, - "Output Response Tokens": 1, # Unicode characters - "Output Cost": 0.2 * 1 / 1000, - "Total Cost": (0.1 * 1 / 1000) + (0.2 * 1 / 1000), - } - ), - ]) - - def test_create_cost_log_entry(self, filename, caller_name, prompt, response, input_cost_per_1k, output_cost_per_1k, expected_output): - """Test the create_cost_log_entry function.""" - result = create_cost_log_entry(filename, caller_name, prompt, response, input_cost_per_1k, output_cost_per_1k) - assert result == expected_output - - # Test cases for update_stats def test_update_stats(self): - """Test the update_stats function.""" - # Initialize the global state to zeros - config.MODEL_STATS = {"file1.txt": {"requests": 0, "tokens_sent": 0, "tokens_received": 0, "total_time": 0}} - config.GLOBAL_STATS = {"total_requests": 0, "total_tokens_sent": 0, "total_tokens_received": 0, "total_time": 0} + llm_utils.update_stats(self.filename, 10, 20, 5, self.model_type) + self.assertEqual(config.MODEL_STATS[self.filename]["requests"], 1) + self.assertEqual(config.MODEL_STATS[self.filename]["tokens_sent"], 10) + self.assertEqual(config.MODEL_STATS[self.filename]["tokens_received"], 20) + self.assertEqual(config.MODEL_STATS[self.filename]["total_time"], 5) + self.assertEqual(config.GLOBAL_STATS["total_requests"], 1) + self.assertEqual(config.GLOBAL_STATS["total_tokens_sent"], 10) + self.assertEqual(config.GLOBAL_STATS["total_tokens_received"], 20) + self.assertEqual(config.GLOBAL_STATS["total_time"], 5) - # Verify initial state - assert config.MODEL_STATS["file1.txt"]["requests"] == 0 - assert config.GLOBAL_STATS["total_requests"] == 0 + def test_count_tokens(self): + self.assertEqual(llm_utils.count_tokens("123456"), 1) + self.assertEqual(llm_utils.count_tokens("1234567"), 2) + self.assertEqual(llm_utils.count_tokens(""), 0) - # Update stats - update_stats("file1.txt", 10, 20, 5.0, "claude2") + def test_create_cost_log_entry(self): + entry = llm_utils.create_cost_log_entry( + self.filename, "caller_func", self.prompt, self.response, self.input_cost_per_1k, self.output_cost_per_1k + ) + self.assertEqual(entry["Filename"], self.filename) + self.assertEqual(entry["Caller Function"], "caller_func") + self.assertEqual(entry["Input Prompt Length"], len(self.prompt)) + self.assertEqual(entry["Output Response Length"], len(self.response)) + self.assertIn("Total Cost", entry) - # Check updated state - assert config.MODEL_STATS["file1.txt"]["requests"] == 1 - assert config.MODEL_STATS["file1.txt"]["tokens_sent"] == 10 - assert config.MODEL_STATS["file1.txt"]["tokens_received"] == 20 - assert config.MODEL_STATS["file1.txt"]["total_time"] == 5.0 - assert config.GLOBAL_STATS["total_requests"] == 1 - assert config.GLOBAL_STATS["total_tokens_sent"] == 10 - assert config.GLOBAL_STATS["total_tokens_received"] == 20 - assert config.GLOBAL_STATS["total_time"] == 5.0 + def test_get_cache_key(self): + model_id = "test_model" + expected_hash = hashlib.sha256(f"{model_id}:{self.prompt}".encode()).hexdigest() + self.assertEqual(llm_utils.get_cache_key(self.prompt, model_id), expected_hash) - # Test cases for invoke_claude - @pytest.mark.parametrize("run_mode, model_id, expected_response", [ - ("local", config.MODEL_ID_CLAUDE2, "mocked response from claude 2 (local)"), - ("local", config.MODEL_ID_CLAUDE3_HAIKU, "mocked response from claude 3 (local)"), - ("local", config.MODEL_ID_CLAUDE35_SONNET, "mocked response from claude 3.5 (local)"), - ("ec2", config.MODEL_ID_CLAUDE2, "mocked response from claude 2 (ec2)"), - ("ec2", config.MODEL_ID_CLAUDE3_HAIKU, "mocked response from claude 3 (ec2)"), - ("ec2", config.MODEL_ID_CLAUDE35_SONNET, "mocked response from claude 3.5 (ec2)"), - ]) - - def test_invoke_claude(self, run_mode, model_id, expected_response, mocker): - """Test the invoke_claude function.""" - # Map model IDs to their corresponding mock functions - mock_functions = { - ("local", config.MODEL_ID_CLAUDE2): "src.utils.llm_utils.local_claude_2", - ("local", config.MODEL_ID_CLAUDE3_HAIKU): "src.utils.llm_utils.local_claude_3", - ("local", config.MODEL_ID_CLAUDE35_SONNET): "src.utils.llm_utils.local_claude_3", - ("ec2", config.MODEL_ID_CLAUDE2): "src.utils.llm_utils.ec2_claude2", - ("ec2", config.MODEL_ID_CLAUDE3_HAIKU): "src.utils.llm_utils.ec2_claude_3", - ("ec2", config.MODEL_ID_CLAUDE35_SONNET): "src.utils.llm_utils.ec2_claude_35", - } + @patch("builtins.open", new_callable=mock_open) + @patch("csv.DictWriter.writerow") + def test_invoke_claude_local_mode(self, mock_csv_writer, mock_open_file): + config.RUN_MODE = "local" + model_id = config.MODEL_ID_CLAUDE2 + + with patch("src.utils.llm_utils.local_claude_2", return_value=self.response): + response = llm_utils.invoke_claude(self.prompt, model_id, self.filename) + self.assertEqual(response, self.response) + mock_csv_writer.assert_called() - # Mock the appropriate function based on run_mode and model_id - mock_function_path = mock_functions.get((run_mode, model_id)) - if mock_function_path: - mocker.patch(mock_function_path, return_value=expected_response) - - # Mock config.RUN_MODE to simulate the run mode - mocker.patch("src.utils.llm_utils.config.RUN_MODE", run_mode) - - # Mock CSV writer to avoid writing to disk - mocker.patch("csv.DictWriter.writerow") - - # Call the function - response = invoke_claude("test prompt", model_id, "file1.txt") - - # Assertions - assert response == expected_response +if __name__ == "__main__": + unittest.main()