Merged in hotfix/pipeline-fix (pull request #393)

Draft: update unit tests

* update unit tests

* Update unit tests

* Update unit tests


Approved-by: Alex Galarce
This commit is contained in:
Katon Minhas
2025-02-12 18:59:59 +00:00
parent 2535400124
commit 7724a7ddc4
+55 -157
View File
@@ -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()