mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-24 12:41:02 +02:00
More tests
This commit is contained in:
parent
5cb51e282d
commit
c5be657835
4 changed files with 861 additions and 0 deletions
280
tests/unit/test_text_completion/test_bedrock_processor.py
Normal file
280
tests/unit/test_text_completion/test_bedrock_processor.py
Normal file
|
|
@ -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__])
|
||||||
|
|
@ -442,6 +442,162 @@ class TestCohereProcessorSimple(IsolatedAsyncioTestCase):
|
||||||
assert call_args[1]['prompt_truncation'] == 'auto'
|
assert call_args[1]['prompt_truncation'] == 'auto'
|
||||||
assert call_args[1]['connectors'] == []
|
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__':
|
if __name__ == '__main__':
|
||||||
pytest.main([__file__])
|
pytest.main([__file__])
|
||||||
|
|
@ -477,6 +477,156 @@ class TestGoogleAIStudioProcessorSimple(IsolatedAsyncioTestCase):
|
||||||
# The system instruction should be in the config object
|
# The system instruction should be in the config object
|
||||||
assert call_args[1]['contents'] == "Explain quantum computing"
|
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__':
|
if __name__ == '__main__':
|
||||||
pytest.main([__file__])
|
pytest.main([__file__])
|
||||||
275
tests/unit/test_text_completion/test_mistral_processor.py
Normal file
275
tests/unit/test_text_completion/test_mistral_processor.py
Normal file
|
|
@ -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__])
|
||||||
Loading…
Add table
Add a link
Reference in a new issue