import hashlib import threading import unittest from unittest.mock import mock_open, patch import src.utils.llm_utils as llm_utils from src import config 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." # Clear cache before each test llm_utils.claude_cache.clear() def tearDown(self): # Clear cache after each test llm_utils.claude_cache.clear() def test_get_cache_key_basic(self): """Test basic cache key generation (backward compatible).""" 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) def test_get_cache_key_with_instruction(self): """Test cache key includes instruction parameter.""" model_id = "test_model" instruction = "You are a helpful assistant." key1 = llm_utils.get_cache_key(self.prompt, model_id, instruction=instruction) key2 = llm_utils.get_cache_key(self.prompt, model_id, instruction="Different instruction") # Keys should be different when instruction differs self.assertNotEqual(key1, key2) # Same instruction should produce same key key3 = llm_utils.get_cache_key(self.prompt, model_id, instruction=instruction) self.assertEqual(key1, key3) def test_get_cache_key_with_max_tokens(self): """Test cache key includes max_tokens parameter.""" model_id = "test_model" key1 = llm_utils.get_cache_key(self.prompt, model_id, max_tokens=100) key2 = llm_utils.get_cache_key(self.prompt, model_id, max_tokens=200) # Keys should be different when max_tokens differs self.assertNotEqual(key1, key2) # Same max_tokens should produce same key key3 = llm_utils.get_cache_key(self.prompt, model_id, max_tokens=100) self.assertEqual(key1, key3) def test_get_cache_key_with_cache_flag(self): """Test cache key includes cache flag.""" model_id = "test_model" key1 = llm_utils.get_cache_key(self.prompt, model_id, cache=False) key2 = llm_utils.get_cache_key(self.prompt, model_id, cache=True) # Keys should be different when cache flag differs self.assertNotEqual(key1, key2) def test_get_cache_key_all_parameters(self): """Test cache key with all parameters.""" model_id = "test_model" instruction = "You are helpful" max_tokens = 100 key1 = llm_utils.get_cache_key( self.prompt, model_id, instruction=instruction, max_tokens=max_tokens, cache=True ) key2 = llm_utils.get_cache_key( self.prompt, model_id, instruction=instruction, max_tokens=max_tokens, cache=True ) # Same parameters should produce same key self.assertEqual(key1, key2) # Different parameters should produce different keys key3 = llm_utils.get_cache_key( self.prompt, model_id, instruction="Different", max_tokens=max_tokens, cache=True ) self.assertNotEqual(key1, key3) def test_get_cache_key_with_image_array(self): """Test cache key generation for image arrays.""" model_id = "test_model" image_list = ["base64_image1", "base64_image2"] key1 = llm_utils.get_cache_key_with_image_array( self.prompt, model_id, image_list, max_tokens=100 ) key2 = llm_utils.get_cache_key_with_image_array( self.prompt, model_id, image_list, max_tokens=200 ) # Keys should be different when max_tokens differs self.assertNotEqual(key1, key2) # Same parameters should produce same key key3 = llm_utils.get_cache_key_with_image_array( self.prompt, model_id, image_list, max_tokens=100 ) self.assertEqual(key1, key3) def test_cache_thread_safety_concurrent_reads(self): """Test thread safety with concurrent cache reads.""" model_id = "test_model" cache_key = llm_utils.get_cache_key(self.prompt, model_id) # Pre-populate cache llm_utils.claude_cache[cache_key] = self.response results = [] errors = [] def read_cache(): try: # Simulate cache read with llm_utils._cache_lock: if cache_key in llm_utils.claude_cache: result = llm_utils.claude_cache[cache_key] results.append(result) except Exception as e: errors.append(e) # Create multiple threads reading from cache threads = [] for _ in range(10): thread = threading.Thread(target=read_cache) threads.append(thread) thread.start() for thread in threads: thread.join() # Verify no errors and all reads succeeded self.assertEqual(len(errors), 0) self.assertEqual(len(results), 10) self.assertTrue(all(r == self.response for r in results)) def test_cache_thread_safety_concurrent_writes(self): """Test thread safety with concurrent cache writes.""" model_id = "test_model" num_threads = 10 def write_to_cache(thread_id): prompt = f"prompt_{thread_id}" response = f"response_{thread_id}" cache_key = llm_utils.get_cache_key(prompt, model_id) # Simulate cache write with lock with llm_utils._cache_lock: if len(llm_utils.claude_cache) >= llm_utils.CACHE_LIMIT: llm_utils.claude_cache.popitem(last=False) llm_utils.claude_cache[cache_key] = response llm_utils.claude_cache.move_to_end(cache_key) threads = [] for i in range(num_threads): thread = threading.Thread(target=write_to_cache, args=(i,)) threads.append(thread) thread.start() for thread in threads: thread.join() # Verify all writes succeeded (cache should have entries) self.assertGreater(len(llm_utils.claude_cache), 0) self.assertLessEqual(len(llm_utils.claude_cache), num_threads) def test_cache_thread_safety_mixed_operations(self): """Test thread safety with mixed read/write operations.""" model_id = "test_model" initial_key = llm_utils.get_cache_key("initial_prompt", model_id) llm_utils.claude_cache[initial_key] = "initial_response" read_count = [0] write_count = [0] def read_operation(): with llm_utils._cache_lock: if initial_key in llm_utils.claude_cache: _ = llm_utils.claude_cache[initial_key] read_count[0] += 1 def write_operation(thread_id): prompt = f"write_prompt_{thread_id}" response = f"write_response_{thread_id}" cache_key = llm_utils.get_cache_key(prompt, model_id) with llm_utils._cache_lock: if len(llm_utils.claude_cache) >= llm_utils.CACHE_LIMIT: llm_utils.claude_cache.popitem(last=False) llm_utils.claude_cache[cache_key] = response write_count[0] += 1 threads = [] # Mix of read and write threads for i in range(5): threads.append(threading.Thread(target=read_operation)) for i in range(5): threads.append(threading.Thread(target=write_operation, args=(i,))) for thread in threads: thread.start() for thread in threads: thread.join() # Verify operations completed without errors self.assertEqual(read_count[0], 5) self.assertEqual(write_count[0], 5) @patch("builtins.open", new_callable=mock_open) def test_invoke_claude_local_mode(self, mock_open_file): """Test invoke_claude in local mode.""" original_run_mode = config.RUN_MODE try: config.RUN_MODE = "local" model_id = config.MODEL_ID_CLAUDE35_SONNET with patch("src.utils.llm_utils.local_claude_3_and_up", return_value=self.response): response = llm_utils.invoke_claude(self.prompt, model_id, self.filename) self.assertEqual(response, self.response) finally: config.RUN_MODE = original_run_mode @patch("builtins.open", new_callable=mock_open) def test_invoke_claude_uses_enhanced_cache_key(self, mock_open_file): """Test that invoke_claude uses enhanced cache key with all parameters.""" original_run_mode = config.RUN_MODE try: config.RUN_MODE = "local" model_id = config.MODEL_ID_CLAUDE35_SONNET instruction = "Test instruction" max_tokens = 100 cache = True with patch("src.utils.llm_utils.local_claude_3_and_up", return_value=self.response): # First call should not use cache (cache is empty) response1 = llm_utils.invoke_claude( self.prompt, model_id, self.filename, max_tokens=max_tokens, cache=cache, instruction=instruction ) # Second call with same parameters should use cache with patch("src.utils.llm_utils.local_claude_3_and_up") as mock_local: response2 = llm_utils.invoke_claude( self.prompt, model_id, self.filename, max_tokens=max_tokens, cache=cache, instruction=instruction ) # Should not call local_claude_3_and_up again (cache hit) mock_local.assert_not_called() self.assertEqual(response1, response2) finally: config.RUN_MODE = original_run_mode llm_utils.claude_cache.clear() def test_token_tracking_uses_exact_api_counts(self): """Test that token tracking uses exact API response counts, not estimates.""" # This test verifies that we're using usage_tracking.extract_usage_from_response() # which gets exact counts from API, not estimated counts like len(prompt.split()) # Mock API response with exact token counts mock_response_body = { "usage": { "input_tokens": 150, "output_tokens": 75, "cache_creation_input_tokens": 20, "cache_read_input_tokens": 30, } } # Verify extract_usage_from_response extracts exact counts import src.utils.usage_tracking as usage_tracking input_tokens, output_tokens, cache_creation, cache_read = ( usage_tracking.extract_usage_from_response(mock_response_body) ) # Should match exact API response, not estimates self.assertEqual(input_tokens, 150) self.assertEqual(output_tokens, 75) self.assertEqual(cache_creation, 20) self.assertEqual(cache_read, 30) # Verify these are not estimated (e.g., len(prompt.split())) # Exact counts from API are integers, not based on string length self.assertIsInstance(input_tokens, int) self.assertIsInstance(output_tokens, int) def test_validate_model_support_deprecated_claude_2(self): """Test that deprecated Claude 2 models raise ValueError.""" with self.assertRaises(ValueError) as context: llm_utils._validate_model_support( "anthropic.claude-instant-v1", "anthropic.claude-instant-v1" ) self.assertIn("Claude 2 is no longer supported", str(context.exception)) def test_validate_model_support_claude_3_haiku_supported(self): """Test that Claude 3 Haiku models are supported.""" # Should not raise an exception try: llm_utils._validate_model_support( config.MODEL_ID_CLAUDE45_HAIKU, config.MODEL_ID_CLAUDE45_HAIKU ) except ValueError: self.fail("_validate_model_support raised ValueError for supported Claude 3 Haiku model") def test_validate_model_support_supported_model(self): """Test that supported models pass validation.""" # Should not raise an exception try: llm_utils._validate_model_support( config.MODEL_ID_CLAUDE35_SONNET, config.MODEL_ID_CLAUDE35_SONNET ) except ValueError: self.fail("_validate_model_support raised ValueError for supported model") def test_validate_model_support_unsupported_model(self): """Test that unsupported models raise ValueError.""" with self.assertRaises(ValueError) as context: llm_utils._validate_model_support( "unknown.model.id", "unknown.model.id" ) self.assertIn("Only Claude 3+ models are supported", str(context.exception)) self.assertIn("Haiku", str(context.exception)) def test_validate_model_support_with_context(self): """Test that context parameter is included in error messages.""" with self.assertRaises(ValueError) as context: llm_utils._validate_model_support( "unknown.model.id", "unknown.model.id", context="for image array" ) self.assertIn("for image array", str(context.exception)) def test_supports_prompt_cache_supported_models(self): """Test that models supporting prompt cache return True.""" # Test Claude 3.7 Sonnet if config.MODEL_ID_CLAUDE37_SONNET: self.assertTrue(llm_utils._supports_prompt_cache(config.MODEL_ID_CLAUDE37_SONNET)) # Test Claude 4 Sonnet if config.MODEL_ID_CLAUDE4_SONNET: self.assertTrue(llm_utils._supports_prompt_cache(config.MODEL_ID_CLAUDE4_SONNET)) # Test Claude 4.5 Sonnet if config.MODEL_ID_CLAUDE45_SONNET: self.assertTrue(llm_utils._supports_prompt_cache(config.MODEL_ID_CLAUDE45_SONNET)) def test_supports_prompt_cache_unsupported_models(self): """Test that models not supporting prompt cache return False.""" # Claude 3.5 Sonnet v1 doesn't support caching (only v2+) if config.MODEL_ID_CLAUDE35_SONNET: self.assertFalse(llm_utils._supports_prompt_cache(config.MODEL_ID_CLAUDE35_SONNET)) # Unknown model self.assertFalse(llm_utils._supports_prompt_cache("unknown.model.id")) def test_build_claude_3_request_body_basic(self): """Test building request body without instruction.""" body = llm_utils._build_claude_3_request_body( "test prompt", max_tokens=100 ) self.assertEqual(body["anthropic_version"], "bedrock-2023-05-31") self.assertEqual(body["max_tokens"], 100) self.assertEqual(body["temperature"], 0.0) self.assertEqual(len(body["messages"]), 1) self.assertEqual(body["messages"][0]["role"], "user") self.assertNotIn("system", body) def test_build_claude_3_request_body_with_instruction_no_cache(self): """Test building request body with instruction but no caching.""" instruction = "You are a helpful assistant." body = llm_utils._build_claude_3_request_body( "test prompt", max_tokens=100, instruction=instruction, cache=False ) self.assertEqual(body["system"], instruction) self.assertNotIsInstance(body["system"], list) def test_build_claude_3_request_body_with_cache_supported_model(self): """Test building request body with caching enabled for supported model.""" instruction = "You are a helpful assistant." model_id = config.MODEL_ID_CLAUDE37_SONNET or config.MODEL_ID_CLAUDE4_SONNET if model_id and llm_utils._supports_prompt_cache(model_id): body = llm_utils._build_claude_3_request_body( "test prompt", max_tokens=100, instruction=instruction, cache=True, model_id=model_id ) # Should use array format with cache_control self.assertIsInstance(body["system"], list) self.assertEqual(len(body["system"]), 1) self.assertEqual(body["system"][0]["type"], "text") self.assertEqual(body["system"][0]["text"], instruction) self.assertEqual(body["system"][0]["cache_control"]["type"], "ephemeral") def test_build_claude_3_request_body_with_cache_unsupported_model(self): """Test building request body with caching enabled for unsupported model.""" instruction = "You are a helpful assistant." model_id = config.MODEL_ID_CLAUDE35_SONNET # v1 doesn't support caching if model_id: body = llm_utils._build_claude_3_request_body( "test prompt", max_tokens=100, instruction=instruction, cache=True, model_id=model_id ) # Should use simple string format (no cache_control) self.assertEqual(body["system"], instruction) self.assertNotIsInstance(body["system"], list) if __name__ == "__main__": unittest.main()