From 50a4549c020a580b940f42e3b24ec7d37858d1a7 Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Fri, 11 Jul 2025 17:17:50 +0100 Subject: [PATCH] Added more LLMs --- tests/unit/test_text_completion/conftest.py | 106 +++- .../test_llamafile_processor.py | 454 ++++++++++++++++++ 2 files changed, 559 insertions(+), 1 deletion(-) create mode 100644 tests/unit/test_text_completion/test_llamafile_processor.py diff --git a/tests/unit/test_text_completion/conftest.py b/tests/unit/test_text_completion/conftest.py index 497aed07..c444ebbb 100644 --- a/tests/unit/test_text_completion/conftest.py +++ b/tests/unit/test_text_completion/conftest.py @@ -392,4 +392,108 @@ def mock_vllm_error_response(): """Mock vLLM error response""" mock_response = MagicMock() mock_response.status = 500 - return mock_response \ No newline at end of file + return mock_response + + +# === Cohere Specific Fixtures === + +@pytest.fixture +def cohere_processor_config(base_processor_config): + """Default configuration for Cohere processor""" + config = base_processor_config.copy() + config.update({ + 'model': 'c4ai-aya-23-8b', + 'api_key': 'test-api-key', + 'temperature': 0.0 + }) + return config + + +@pytest.fixture +def mock_cohere_client(): + """Mock Cohere client""" + mock_client = MagicMock() + + # Mock the response structure + mock_output = MagicMock() + mock_output.text = "Test response from Cohere" + mock_output.meta.billed_units.input_tokens = 18 + mock_output.meta.billed_units.output_tokens = 10 + + mock_client.chat.return_value = mock_output + return mock_client + + +@pytest.fixture +def mock_cohere_rate_limit_error(): + """Mock Cohere rate limit error""" + import cohere + return cohere.TooManyRequestsError("Rate limit exceeded") + + +# === Google AI Studio Specific Fixtures === + +@pytest.fixture +def googleaistudio_processor_config(base_processor_config): + """Default configuration for Google AI Studio processor""" + config = base_processor_config.copy() + config.update({ + 'model': 'gemini-2.0-flash-001', + 'api_key': 'test-api-key', + 'temperature': 0.0, + 'max_output': 8192 + }) + return config + + +@pytest.fixture +def mock_googleaistudio_client(): + """Mock Google AI Studio client""" + mock_client = MagicMock() + + # Mock the response structure + mock_response = MagicMock() + mock_response.text = "Test response from Google AI Studio" + mock_response.usage_metadata.prompt_token_count = 20 + mock_response.usage_metadata.candidates_token_count = 12 + + mock_client.models.generate_content.return_value = mock_response + return mock_client + + +@pytest.fixture +def mock_googleaistudio_rate_limit_error(): + """Mock Google AI Studio rate limit error""" + from google.api_core.exceptions import ResourceExhausted + return ResourceExhausted("Rate limit exceeded") + + +# === LlamaFile Specific Fixtures === + +@pytest.fixture +def llamafile_processor_config(base_processor_config): + """Default configuration for LlamaFile processor""" + config = base_processor_config.copy() + config.update({ + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096 + }) + return config + + +@pytest.fixture +def mock_llamafile_client(): + """Mock OpenAI client for LlamaFile""" + mock_client = MagicMock() + + # Mock the response structure + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Test response from LlamaFile" + mock_response.usage.prompt_tokens = 14 + mock_response.usage.completion_tokens = 8 + + mock_client.chat.completions.create.return_value = mock_response + return mock_client \ No newline at end of file diff --git a/tests/unit/test_text_completion/test_llamafile_processor.py b/tests/unit/test_text_completion/test_llamafile_processor.py new file mode 100644 index 00000000..bae1a4bb --- /dev/null +++ b/tests/unit/test_text_completion/test_llamafile_processor.py @@ -0,0 +1,454 @@ +""" +Unit tests for trustgraph.model.text_completion.llamafile +Following the same successful pattern as previous tests +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch +from unittest import IsolatedAsyncioTestCase + +# Import the service under test +from trustgraph.model.text_completion.llamafile.llm import Processor +from trustgraph.base import LlmResult +from trustgraph.exceptions import TooManyRequests + + +class TestLlamaFileProcessorSimple(IsolatedAsyncioTestCase): + """Test LlamaFile processor functionality""" + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_processor_initialization_basic(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test basic processor initialization""" + # Arrange + mock_openai_client = MagicMock() + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + # Act + processor = Processor(**config) + + # Assert + assert processor.model == 'LLaMA_CPP' + assert processor.llamafile == 'http://localhost:8080/v1' + assert processor.temperature == 0.0 + assert processor.max_output == 4096 + assert hasattr(processor, 'openai') + mock_openai_class.assert_called_once_with( + base_url='http://localhost:8080/v1', + api_key='sk-no-key-required' + ) + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_success(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test successful content generation""" + # Arrange + mock_openai_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Generated response from LlamaFile" + mock_response.usage.prompt_tokens = 20 + mock_response.usage.completion_tokens = 12 + + mock_openai_client.chat.completions.create.return_value = mock_response + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act + result = await processor.generate_content("System prompt", "User prompt") + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Generated response from LlamaFile" + assert result.in_token == 20 + assert result.out_token == 12 + assert result.model == 'llama.cpp' # Note: model in result is hardcoded to 'llama.cpp' + + # Verify the OpenAI API call structure + mock_openai_client.chat.completions.create.assert_called_once_with( + model='LLaMA_CPP', + messages=[{ + "role": "user", + "content": "System prompt\n\nUser prompt" + }] + ) + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_generic_exception(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test handling of generic exceptions""" + # Arrange + mock_openai_client = MagicMock() + mock_openai_client.chat.completions.create.side_effect = Exception("Connection error") + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act & Assert + with pytest.raises(Exception, match="Connection error"): + await processor.generate_content("System prompt", "User prompt") + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_processor_initialization_with_custom_parameters(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test processor initialization with custom parameters""" + # Arrange + mock_openai_client = MagicMock() + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'custom-llama', + 'llamafile': 'http://custom-host:8080/v1', + 'temperature': 0.7, + 'max_output': 2048, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + # Act + processor = Processor(**config) + + # Assert + assert processor.model == 'custom-llama' + assert processor.llamafile == 'http://custom-host:8080/v1' + assert processor.temperature == 0.7 + assert processor.max_output == 2048 + mock_openai_class.assert_called_once_with( + base_url='http://custom-host:8080/v1', + api_key='sk-no-key-required' + ) + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_processor_initialization_with_defaults(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test processor initialization with default values""" + # Arrange + mock_openai_client = MagicMock() + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + # Only provide required fields, should use defaults + config = { + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + # Act + processor = Processor(**config) + + # Assert + assert processor.model == 'LLaMA_CPP' # default_model + assert processor.llamafile == 'http://localhost:8080/v1' # default_llamafile + assert processor.temperature == 0.0 # default_temperature + assert processor.max_output == 4096 # default_max_output + mock_openai_class.assert_called_once_with( + base_url='http://localhost:8080/v1', + api_key='sk-no-key-required' + ) + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_empty_prompts(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test content generation with empty prompts""" + # Arrange + mock_openai_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Default response" + mock_response.usage.prompt_tokens = 2 + mock_response.usage.completion_tokens = 3 + + mock_openai_client.chat.completions.create.return_value = mock_response + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act + result = await processor.generate_content("", "") + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Default response" + assert result.in_token == 2 + assert result.out_token == 3 + assert result.model == 'llama.cpp' + + # Verify the combined prompt is sent correctly + call_args = mock_openai_client.chat.completions.create.call_args + expected_prompt = "\n\n" # Empty system + "\n\n" + empty user + assert call_args[1]['messages'][0]['content'] == expected_prompt + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_message_structure(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test that LlamaFile messages are structured correctly""" + # Arrange + mock_openai_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Response with proper structure" + mock_response.usage.prompt_tokens = 25 + mock_response.usage.completion_tokens = 15 + + mock_openai_client.chat.completions.create.return_value = mock_response + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act + result = await processor.generate_content("You are a helpful assistant", "What is AI?") + + # Assert + assert result.text == "Response with proper structure" + assert result.in_token == 25 + assert result.out_token == 15 + + # Verify the message structure + call_args = mock_openai_client.chat.completions.create.call_args + messages = call_args[1]['messages'] + + assert len(messages) == 1 + assert messages[0]['role'] == 'user' + assert messages[0]['content'] == "You are a helpful assistant\n\nWhat is AI?" + + # Verify model parameter + assert call_args[1]['model'] == 'LLaMA_CPP' + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_openai_client_initialization(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test that OpenAI client is initialized correctly for LlamaFile""" + # Arrange + mock_openai_client = MagicMock() + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'llama-custom', + 'llamafile': 'http://llamafile-server:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + # Act + processor = Processor(**config) + + # Assert + # Verify OpenAI client was called with correct parameters + mock_openai_class.assert_called_once_with( + base_url='http://llamafile-server:8080/v1', + api_key='sk-no-key-required' + ) + + # Verify processor has the client + assert processor.openai == mock_openai_client + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_prompt_construction(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test prompt construction with system and user prompts""" + # Arrange + mock_openai_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Response with system instructions" + mock_response.usage.prompt_tokens = 30 + mock_response.usage.completion_tokens = 20 + + mock_openai_client.chat.completions.create.return_value = mock_response + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act + result = await processor.generate_content("You are a helpful assistant", "What is machine learning?") + + # Assert + assert result.text == "Response with system instructions" + assert result.in_token == 30 + assert result.out_token == 20 + + # Verify the combined prompt + call_args = mock_openai_client.chat.completions.create.call_args + expected_prompt = "You are a helpful assistant\n\nWhat is machine learning?" + assert call_args[1]['messages'][0]['content'] == expected_prompt + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_hardcoded_model_response(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test that response model is hardcoded to 'llama.cpp'""" + # Arrange + mock_openai_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Test response" + mock_response.usage.prompt_tokens = 15 + mock_response.usage.completion_tokens = 10 + + mock_openai_client.chat.completions.create.return_value = mock_response + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'custom-model-name', # This should be ignored in response + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act + result = await processor.generate_content("System", "User") + + # Assert + assert result.model == 'llama.cpp' # Should always be 'llama.cpp', not 'custom-model-name' + assert processor.model == 'custom-model-name' # But processor.model should still be custom + + @patch('trustgraph.model.text_completion.llamafile.llm.OpenAI') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_no_rate_limiting(self, mock_llm_init, mock_async_init, mock_openai_class): + """Test that no rate limiting is implemented (SLM assumption)""" + # Arrange + mock_openai_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "No rate limiting test" + mock_response.usage.prompt_tokens = 10 + mock_response.usage.completion_tokens = 5 + + mock_openai_client.chat.completions.create.return_value = mock_response + mock_openai_class.return_value = mock_openai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'LLaMA_CPP', + 'llamafile': 'http://localhost:8080/v1', + 'temperature': 0.0, + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act + result = await processor.generate_content("System", "User") + + # Assert + assert result.text == "No rate limiting test" + # No specific rate limit error handling tested since SLM presumably has no rate limits + + +if __name__ == '__main__': + pytest.main([__file__]) \ No newline at end of file