diff --git a/tests/unit/test_text_completion/test_bedrock_processor.py b/tests/unit/test_text_completion/test_bedrock_processor.py new file mode 100644 index 00000000..b8ead9a6 --- /dev/null +++ b/tests/unit/test_text_completion/test_bedrock_processor.py @@ -0,0 +1,280 @@ +""" +Unit tests for trustgraph.model.text_completion.bedrock +Following the same successful pattern as other processor tests +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch +from unittest import IsolatedAsyncioTestCase +import json + +# Import the service under test +from trustgraph.model.text_completion.bedrock.llm import Processor, Mistral, Anthropic +from trustgraph.base import LlmResult + + +class TestBedrockProcessorSimple(IsolatedAsyncioTestCase): + """Test Bedrock processor functionality""" + + @patch('trustgraph.model.text_completion.bedrock.llm.boto3.Session') + @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_session_class): + """Test basic processor initialization""" + # Arrange + mock_session = MagicMock() + mock_bedrock = MagicMock() + mock_session.client.return_value = mock_bedrock + mock_session_class.return_value = mock_session + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'mistral.mistral-large-2407-v1:0', + 'temperature': 0.1, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + # Act + processor = Processor(**config) + + # Assert + assert processor.default_model == 'mistral.mistral-large-2407-v1:0' + assert processor.temperature == 0.1 + assert hasattr(processor, 'bedrock') + mock_session_class.assert_called_once() + + @patch('trustgraph.model.text_completion.bedrock.llm.boto3.Session') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_success_mistral(self, mock_llm_init, mock_async_init, mock_session_class): + """Test successful content generation with Mistral model""" + # Arrange + mock_session = MagicMock() + mock_bedrock = MagicMock() + mock_session.client.return_value = mock_bedrock + mock_session_class.return_value = mock_session + + mock_response = { + 'body': MagicMock(), + 'ResponseMetadata': { + 'HTTPHeaders': { + 'x-amzn-bedrock-input-token-count': '15', + 'x-amzn-bedrock-output-token-count': '8' + } + } + } + mock_response['body'].read.return_value = json.dumps({ + 'outputs': [{'text': 'Generated response from Bedrock'}] + }) + mock_bedrock.invoke_model.return_value = mock_response + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'mistral.mistral-large-2407-v1:0', + 'temperature': 0.0, + '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 Bedrock" + assert result.in_token == 15 + assert result.out_token == 8 + assert result.model == 'mistral.mistral-large-2407-v1:0' + mock_bedrock.invoke_model.assert_called_once() + + @patch('trustgraph.model.text_completion.bedrock.llm.boto3.Session') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_temperature_override(self, mock_llm_init, mock_async_init, mock_session_class): + """Test temperature parameter override functionality""" + # Arrange + mock_session = MagicMock() + mock_bedrock = MagicMock() + mock_session.client.return_value = mock_bedrock + mock_session_class.return_value = mock_session + + mock_response = { + 'body': MagicMock(), + 'ResponseMetadata': { + 'HTTPHeaders': { + 'x-amzn-bedrock-input-token-count': '20', + 'x-amzn-bedrock-output-token-count': '12' + } + } + } + mock_response['body'].read.return_value = json.dumps({ + 'outputs': [{'text': 'Response with custom temperature'}] + }) + mock_bedrock.invoke_model.return_value = mock_response + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'mistral.mistral-large-2407-v1:0', + 'temperature': 0.0, # Default temperature + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override temperature at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model=None, # Use default model + temperature=0.8 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom temperature" + + # Verify the model variant was created with overridden temperature + # The cache key should include the temperature + cache_key = f"mistral.mistral-large-2407-v1:0:0.8" + assert cache_key in processor.model_variants + variant = processor.model_variants[cache_key] + assert variant.temperature == 0.8 + + @patch('trustgraph.model.text_completion.bedrock.llm.boto3.Session') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_model_override(self, mock_llm_init, mock_async_init, mock_session_class): + """Test model parameter override functionality""" + # Arrange + mock_session = MagicMock() + mock_bedrock = MagicMock() + mock_session.client.return_value = mock_bedrock + mock_session_class.return_value = mock_session + + mock_response = { + 'body': MagicMock(), + 'ResponseMetadata': { + 'HTTPHeaders': { + 'x-amzn-bedrock-input-token-count': '18', + 'x-amzn-bedrock-output-token-count': '14' + } + } + } + mock_response['body'].read.return_value = json.dumps({ + 'content': [{'text': 'Response with custom model'}] + }) + mock_bedrock.invoke_model.return_value = mock_response + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'mistral.mistral-large-2407-v1:0', # Default model + 'temperature': 0.1, # Default temperature + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override model at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="anthropic.claude-3-sonnet-20240229-v1:0", # Override model + temperature=None # Use default temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom model" + + # Verify Bedrock API was called with overridden model + mock_bedrock.invoke_model.assert_called_once() + call_args = mock_bedrock.invoke_model.call_args + assert call_args[1]['modelId'] == "anthropic.claude-3-sonnet-20240229-v1:0" + + # Verify the correct model variant (Anthropic) was used + cache_key = f"anthropic.claude-3-sonnet-20240229-v1:0:0.1" + assert cache_key in processor.model_variants + variant = processor.model_variants[cache_key] + assert isinstance(variant, Anthropic) + + @patch('trustgraph.model.text_completion.bedrock.llm.boto3.Session') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_both_parameters_override(self, mock_llm_init, mock_async_init, mock_session_class): + """Test overriding both model and temperature parameters simultaneously""" + # Arrange + mock_session = MagicMock() + mock_bedrock = MagicMock() + mock_session.client.return_value = mock_bedrock + mock_session_class.return_value = mock_session + + mock_response = { + 'body': MagicMock(), + 'ResponseMetadata': { + 'HTTPHeaders': { + 'x-amzn-bedrock-input-token-count': '22', + 'x-amzn-bedrock-output-token-count': '16' + } + } + } + mock_response['body'].read.return_value = json.dumps({ + 'generation': 'Response with both overrides' + }) + mock_bedrock.invoke_model.return_value = mock_response + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'mistral.mistral-large-2407-v1:0', # Default model + 'temperature': 0.0, # Default temperature + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override both parameters at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="meta.llama3-70b-instruct-v1:0", # Override model (Meta/Llama) + temperature=0.9 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with both overrides" + + # Verify Bedrock API was called with both overrides + mock_bedrock.invoke_model.assert_called_once() + call_args = mock_bedrock.invoke_model.call_args + assert call_args[1]['modelId'] == "meta.llama3-70b-instruct-v1:0" + + # Verify the correct model variant (Meta) was used with correct temperature + cache_key = f"meta.llama3-70b-instruct-v1:0:0.9" + assert cache_key in processor.model_variants + variant = processor.model_variants[cache_key] + assert variant.temperature == 0.9 + + +if __name__ == '__main__': + pytest.main([__file__]) \ No newline at end of file diff --git a/tests/unit/test_text_completion/test_cohere_processor.py b/tests/unit/test_text_completion/test_cohere_processor.py index 9e4397bc..6201f95c 100644 --- a/tests/unit/test_text_completion/test_cohere_processor.py +++ b/tests/unit/test_text_completion/test_cohere_processor.py @@ -442,6 +442,162 @@ class TestCohereProcessorSimple(IsolatedAsyncioTestCase): assert call_args[1]['prompt_truncation'] == 'auto' assert call_args[1]['connectors'] == [] + @patch('trustgraph.model.text_completion.cohere.llm.cohere.Client') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_temperature_override(self, mock_llm_init, mock_async_init, mock_cohere_class): + """Test temperature parameter override functionality""" + # Arrange + mock_cohere_client = MagicMock() + mock_output = MagicMock() + mock_output.text = 'Response with custom temperature' + mock_output.meta.billed_units.input_tokens = 20 + mock_output.meta.billed_units.output_tokens = 12 + + mock_cohere_client.chat.return_value = mock_output + mock_cohere_class.return_value = mock_cohere_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'c4ai-aya-23-8b', + 'api_key': 'test-api-key', + 'temperature': 0.0, # Default temperature + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override temperature at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model=None, # Use default model + temperature=0.8 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom temperature" + + # Verify Cohere API was called with overridden temperature + mock_cohere_client.chat.assert_called_once_with( + model='c4ai-aya-23-8b', + message='User prompt', + preamble='System prompt', + temperature=0.8, # Should use runtime override + chat_history=[], + prompt_truncation='auto', + connectors=[] + ) + + @patch('trustgraph.model.text_completion.cohere.llm.cohere.Client') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_model_override(self, mock_llm_init, mock_async_init, mock_cohere_class): + """Test model parameter override functionality""" + # Arrange + mock_cohere_client = MagicMock() + mock_output = MagicMock() + mock_output.text = 'Response with custom model' + mock_output.meta.billed_units.input_tokens = 18 + mock_output.meta.billed_units.output_tokens = 14 + + mock_cohere_client.chat.return_value = mock_output + mock_cohere_class.return_value = mock_cohere_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'c4ai-aya-23-8b', # Default model + 'api_key': 'test-api-key', + 'temperature': 0.1, # Default temperature + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override model at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="command-r-plus", # Override model + temperature=None # Use default temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom model" + + # Verify Cohere API was called with overridden model + mock_cohere_client.chat.assert_called_once_with( + model='command-r-plus', # Should use runtime override + message='User prompt', + preamble='System prompt', + temperature=0.1, # Should use processor default + chat_history=[], + prompt_truncation='auto', + connectors=[] + ) + + @patch('trustgraph.model.text_completion.cohere.llm.cohere.Client') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_both_parameters_override(self, mock_llm_init, mock_async_init, mock_cohere_class): + """Test overriding both model and temperature parameters simultaneously""" + # Arrange + mock_cohere_client = MagicMock() + mock_output = MagicMock() + mock_output.text = 'Response with both overrides' + mock_output.meta.billed_units.input_tokens = 22 + mock_output.meta.billed_units.output_tokens = 16 + + mock_cohere_client.chat.return_value = mock_output + mock_cohere_class.return_value = mock_cohere_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'c4ai-aya-23-8b', # Default model + 'api_key': 'test-api-key', + 'temperature': 0.0, # Default temperature + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override both parameters at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="command-r", # Override model + temperature=0.9 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with both overrides" + + # Verify Cohere API was called with both overrides + mock_cohere_client.chat.assert_called_once_with( + model='command-r', # Should use runtime override + message='User prompt', + preamble='System prompt', + temperature=0.9, # Should use runtime override + chat_history=[], + prompt_truncation='auto', + connectors=[] + ) + if __name__ == '__main__': pytest.main([__file__]) \ No newline at end of file diff --git a/tests/unit/test_text_completion/test_googleaistudio_processor.py b/tests/unit/test_text_completion/test_googleaistudio_processor.py index f31715d2..c54b3928 100644 --- a/tests/unit/test_text_completion/test_googleaistudio_processor.py +++ b/tests/unit/test_text_completion/test_googleaistudio_processor.py @@ -477,6 +477,156 @@ class TestGoogleAIStudioProcessorSimple(IsolatedAsyncioTestCase): # The system instruction should be in the config object assert call_args[1]['contents'] == "Explain quantum computing" + @patch('trustgraph.model.text_completion.googleaistudio.llm.genai.Client') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_temperature_override(self, mock_llm_init, mock_async_init, mock_genai_class): + """Test temperature parameter override functionality""" + # Arrange + mock_genai_client = MagicMock() + mock_response = MagicMock() + mock_response.text = 'Response with custom temperature' + mock_response.usage_metadata.prompt_token_count = 20 + mock_response.usage_metadata.candidates_token_count = 12 + + mock_genai_client.models.generate_content.return_value = mock_response + mock_genai_class.return_value = mock_genai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'gemini-2.0-flash-001', + 'api_key': 'test-api-key', + 'temperature': 0.0, # Default temperature + 'max_output': 8192, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override temperature at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model=None, # Use default model + temperature=0.8 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom temperature" + + # Verify the generation config was created with overridden temperature + cache_key = f"gemini-2.0-flash-001:0.8" + assert cache_key in processor.generation_configs + config_obj = processor.generation_configs[cache_key] + assert config_obj.temperature == 0.8 + + @patch('trustgraph.model.text_completion.googleaistudio.llm.genai.Client') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_model_override(self, mock_llm_init, mock_async_init, mock_genai_class): + """Test model parameter override functionality""" + # Arrange + mock_genai_client = MagicMock() + mock_response = MagicMock() + mock_response.text = 'Response with custom model' + mock_response.usage_metadata.prompt_token_count = 18 + mock_response.usage_metadata.candidates_token_count = 14 + + mock_genai_client.models.generate_content.return_value = mock_response + mock_genai_class.return_value = mock_genai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'gemini-2.0-flash-001', # Default model + 'api_key': 'test-api-key', + 'temperature': 0.1, # Default temperature + 'max_output': 8192, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override model at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="gemini-1.5-pro", # Override model + temperature=None # Use default temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom model" + + # Verify Google AI Studio API was called with overridden model + call_args = mock_genai_client.models.generate_content.call_args + assert call_args[1]['model'] == 'gemini-1.5-pro' # Should use runtime override + + # Verify the generation config was created for the correct model + cache_key = f"gemini-1.5-pro:0.1" + assert cache_key in processor.generation_configs + + @patch('trustgraph.model.text_completion.googleaistudio.llm.genai.Client') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_both_parameters_override(self, mock_llm_init, mock_async_init, mock_genai_class): + """Test overriding both model and temperature parameters simultaneously""" + # Arrange + mock_genai_client = MagicMock() + mock_response = MagicMock() + mock_response.text = 'Response with both overrides' + mock_response.usage_metadata.prompt_token_count = 22 + mock_response.usage_metadata.candidates_token_count = 16 + + mock_genai_client.models.generate_content.return_value = mock_response + mock_genai_class.return_value = mock_genai_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'gemini-2.0-flash-001', # Default model + 'api_key': 'test-api-key', + 'temperature': 0.0, # Default temperature + 'max_output': 8192, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override both parameters at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="gemini-1.5-flash", # Override model + temperature=0.9 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with both overrides" + + # Verify Google AI Studio API was called with both overrides + call_args = mock_genai_client.models.generate_content.call_args + assert call_args[1]['model'] == 'gemini-1.5-flash' # Should use runtime override + + # Verify the generation config was created with both overrides + cache_key = f"gemini-1.5-flash:0.9" + assert cache_key in processor.generation_configs + config_obj = processor.generation_configs[cache_key] + assert config_obj.temperature == 0.9 + if __name__ == '__main__': pytest.main([__file__]) \ No newline at end of file diff --git a/tests/unit/test_text_completion/test_mistral_processor.py b/tests/unit/test_text_completion/test_mistral_processor.py new file mode 100644 index 00000000..a40cca70 --- /dev/null +++ b/tests/unit/test_text_completion/test_mistral_processor.py @@ -0,0 +1,275 @@ +""" +Unit tests for trustgraph.model.text_completion.mistral +Following the same successful pattern as other processor tests +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch +from unittest import IsolatedAsyncioTestCase + +# Import the service under test +from trustgraph.model.text_completion.mistral.llm import Processor +from trustgraph.base import LlmResult + + +class TestMistralProcessorSimple(IsolatedAsyncioTestCase): + """Test Mistral processor functionality""" + + @patch('trustgraph.model.text_completion.mistral.llm.Mistral') + @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_mistral_class): + """Test basic processor initialization""" + # Arrange + mock_mistral_client = MagicMock() + mock_mistral_class.return_value = mock_mistral_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'ministral-8b-latest', + 'api_key': 'test-api-key', + 'temperature': 0.1, + 'max_output': 2048, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + # Act + processor = Processor(**config) + + # Assert + assert processor.default_model == 'ministral-8b-latest' + assert processor.temperature == 0.1 + assert processor.max_output == 2048 + assert hasattr(processor, 'mistral') + mock_mistral_class.assert_called_once_with(api_key='test-api-key') + + @patch('trustgraph.model.text_completion.mistral.llm.Mistral') + @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_mistral_class): + """Test successful content generation""" + # Arrange + mock_mistral_client = MagicMock() + mock_response = MagicMock() + mock_response.choices[0].message.content = 'Generated response from Mistral' + mock_response.usage.prompt_tokens = 15 + mock_response.usage.completion_tokens = 8 + mock_mistral_client.chat.complete.return_value = mock_response + mock_mistral_class.return_value = mock_mistral_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'ministral-8b-latest', + 'api_key': 'test-api-key', + '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 Mistral" + assert result.in_token == 15 + assert result.out_token == 8 + assert result.model == 'ministral-8b-latest' + mock_mistral_client.chat.complete.assert_called_once() + + @patch('trustgraph.model.text_completion.mistral.llm.Mistral') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_temperature_override(self, mock_llm_init, mock_async_init, mock_mistral_class): + """Test temperature parameter override functionality""" + # Arrange + mock_mistral_client = MagicMock() + mock_response = MagicMock() + mock_response.choices[0].message.content = 'Response with custom temperature' + mock_response.usage.prompt_tokens = 20 + mock_response.usage.completion_tokens = 12 + mock_mistral_client.chat.complete.return_value = mock_response + mock_mistral_class.return_value = mock_mistral_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'ministral-8b-latest', + 'api_key': 'test-api-key', + 'temperature': 0.0, # Default temperature + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override temperature at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model=None, # Use default model + temperature=0.8 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom temperature" + + # Verify Mistral API was called with overridden temperature + call_args = mock_mistral_client.chat.complete.call_args + assert call_args[1]['temperature'] == 0.8 # Should use runtime override + assert call_args[1]['model'] == 'ministral-8b-latest' + + @patch('trustgraph.model.text_completion.mistral.llm.Mistral') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_model_override(self, mock_llm_init, mock_async_init, mock_mistral_class): + """Test model parameter override functionality""" + # Arrange + mock_mistral_client = MagicMock() + mock_response = MagicMock() + mock_response.choices[0].message.content = 'Response with custom model' + mock_response.usage.prompt_tokens = 18 + mock_response.usage.completion_tokens = 14 + mock_mistral_client.chat.complete.return_value = mock_response + mock_mistral_class.return_value = mock_mistral_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'ministral-8b-latest', # Default model + 'api_key': 'test-api-key', + 'temperature': 0.1, # Default temperature + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override model at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="mistral-large-latest", # Override model + temperature=None # Use default temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with custom model" + + # Verify Mistral API was called with overridden model + call_args = mock_mistral_client.chat.complete.call_args + assert call_args[1]['model'] == 'mistral-large-latest' # Should use runtime override + assert call_args[1]['temperature'] == 0.1 # Should use processor default + + @patch('trustgraph.model.text_completion.mistral.llm.Mistral') + @patch('trustgraph.base.async_processor.AsyncProcessor.__init__') + @patch('trustgraph.base.llm_service.LlmService.__init__') + async def test_generate_content_both_parameters_override(self, mock_llm_init, mock_async_init, mock_mistral_class): + """Test overriding both model and temperature parameters simultaneously""" + # Arrange + mock_mistral_client = MagicMock() + mock_response = MagicMock() + mock_response.choices[0].message.content = 'Response with both overrides' + mock_response.usage.prompt_tokens = 22 + mock_response.usage.completion_tokens = 16 + mock_mistral_client.chat.complete.return_value = mock_response + mock_mistral_class.return_value = mock_mistral_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'ministral-8b-latest', # Default model + 'api_key': 'test-api-key', + 'temperature': 0.0, # Default temperature + 'max_output': 4096, + 'concurrency': 1, + 'taskgroup': AsyncMock(), + 'id': 'test-processor' + } + + processor = Processor(**config) + + # Act - Override both parameters at runtime + result = await processor.generate_content( + "System prompt", + "User prompt", + model="mistral-large-latest", # Override model + temperature=0.9 # Override temperature + ) + + # Assert + assert isinstance(result, LlmResult) + assert result.text == "Response with both overrides" + + # Verify Mistral API was called with both overrides + call_args = mock_mistral_client.chat.complete.call_args + assert call_args[1]['model'] == 'mistral-large-latest' # Should use runtime override + assert call_args[1]['temperature'] == 0.9 # Should use runtime override + + @patch('trustgraph.model.text_completion.mistral.llm.Mistral') + @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_mistral_class): + """Test prompt construction with system and user prompts""" + # Arrange + mock_mistral_client = MagicMock() + mock_response = MagicMock() + mock_response.choices[0].message.content = 'Response with system instructions' + mock_response.usage.prompt_tokens = 25 + mock_response.usage.completion_tokens = 15 + mock_mistral_client.chat.complete.return_value = mock_response + mock_mistral_class.return_value = mock_mistral_client + + mock_async_init.return_value = None + mock_llm_init.return_value = None + + config = { + 'model': 'ministral-8b-latest', + 'api_key': 'test-api-key', + '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 system instructions" + assert result.in_token == 25 + assert result.out_token == 15 + + # Verify the combined prompt structure + call_args = mock_mistral_client.chat.complete.call_args + messages = call_args[1]['messages'] + assert len(messages) == 1 + assert messages[0]['role'] == 'user' + assert messages[0]['content'][0]['type'] == 'text' + assert messages[0]['content'][0]['text'] == "You are a helpful assistant\n\nWhat is AI?" + + +if __name__ == '__main__': + pytest.main([__file__]) \ No newline at end of file