mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-24 12:41:02 +02:00
Fix tests
This commit is contained in:
parent
43c8d25092
commit
27563966fb
4 changed files with 416 additions and 1 deletions
|
|
@ -74,7 +74,7 @@ class TestFlowParameterSpecs(IsolatedAsyncioTestCase):
|
||||||
# Create different spec types
|
# Create different spec types
|
||||||
param_spec = ParameterSpec(name="model")
|
param_spec = ParameterSpec(name="model")
|
||||||
consumer_spec = ConsumerSpec(name="input", schema=MagicMock(), handler=MagicMock())
|
consumer_spec = ConsumerSpec(name="input", schema=MagicMock(), handler=MagicMock())
|
||||||
producer_spec = ProducerSpec(name="output")
|
producer_spec = ProducerSpec(name="output", schema=MagicMock())
|
||||||
|
|
||||||
# Act
|
# Act
|
||||||
processor.register_specification(param_spec)
|
processor.register_specification(param_spec)
|
||||||
|
|
|
||||||
|
|
@ -459,5 +459,150 @@ class TestAzureProcessorSimple(IsolatedAsyncioTestCase):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@patch('trustgraph.model.text_completion.azure.llm.requests')
|
||||||
|
@patch('trustgraph.base.async_processor.AsyncProcessor.__init__')
|
||||||
|
@patch('trustgraph.base.llm_service.LlmService.__init__')
|
||||||
|
async def test_generate_content_with_model_override(self, mock_llm_init, mock_async_init, mock_requests):
|
||||||
|
"""Test generate_content with model parameter override"""
|
||||||
|
# Arrange
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.status_code = 200
|
||||||
|
mock_response.json.return_value = {
|
||||||
|
'choices': [{
|
||||||
|
'message': {
|
||||||
|
'content': 'Response with model override'
|
||||||
|
}
|
||||||
|
}],
|
||||||
|
'usage': {
|
||||||
|
'prompt_tokens': 15,
|
||||||
|
'completion_tokens': 10
|
||||||
|
}
|
||||||
|
}
|
||||||
|
mock_requests.post.return_value = mock_response
|
||||||
|
|
||||||
|
mock_async_init.return_value = None
|
||||||
|
mock_llm_init.return_value = None
|
||||||
|
|
||||||
|
config = {
|
||||||
|
'endpoint': 'https://test.inference.ai.azure.com/v1/chat/completions',
|
||||||
|
'token': 'test-token',
|
||||||
|
'temperature': 0.0,
|
||||||
|
'max_output': 4192,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override model
|
||||||
|
result = await processor.generate_content("System", "Prompt", model="custom-azure-model")
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.model == "custom-azure-model" # Should use overridden model
|
||||||
|
assert result.text == "Response with model override"
|
||||||
|
|
||||||
|
@patch('trustgraph.model.text_completion.azure.llm.requests')
|
||||||
|
@patch('trustgraph.base.async_processor.AsyncProcessor.__init__')
|
||||||
|
@patch('trustgraph.base.llm_service.LlmService.__init__')
|
||||||
|
async def test_generate_content_with_temperature_override(self, mock_llm_init, mock_async_init, mock_requests):
|
||||||
|
"""Test generate_content with temperature parameter override"""
|
||||||
|
# Arrange
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.status_code = 200
|
||||||
|
mock_response.json.return_value = {
|
||||||
|
'choices': [{
|
||||||
|
'message': {
|
||||||
|
'content': 'Response with temperature override'
|
||||||
|
}
|
||||||
|
}],
|
||||||
|
'usage': {
|
||||||
|
'prompt_tokens': 15,
|
||||||
|
'completion_tokens': 10
|
||||||
|
}
|
||||||
|
}
|
||||||
|
mock_requests.post.return_value = mock_response
|
||||||
|
|
||||||
|
mock_async_init.return_value = None
|
||||||
|
mock_llm_init.return_value = None
|
||||||
|
|
||||||
|
config = {
|
||||||
|
'endpoint': 'https://test.inference.ai.azure.com/v1/chat/completions',
|
||||||
|
'token': 'test-token',
|
||||||
|
'temperature': 0.0, # Default temperature
|
||||||
|
'max_output': 4192,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override temperature
|
||||||
|
result = await processor.generate_content("System", "Prompt", temperature=0.8)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.text == "Response with temperature override"
|
||||||
|
|
||||||
|
# Verify the request was made with the overridden temperature
|
||||||
|
mock_requests.post.assert_called_once()
|
||||||
|
call_args = mock_requests.post.call_args
|
||||||
|
|
||||||
|
import json
|
||||||
|
request_body = json.loads(call_args[1]['data'])
|
||||||
|
assert request_body['temperature'] == 0.8
|
||||||
|
|
||||||
|
@patch('trustgraph.model.text_completion.azure.llm.requests')
|
||||||
|
@patch('trustgraph.base.async_processor.AsyncProcessor.__init__')
|
||||||
|
@patch('trustgraph.base.llm_service.LlmService.__init__')
|
||||||
|
async def test_generate_content_with_both_parameters_override(self, mock_llm_init, mock_async_init, mock_requests):
|
||||||
|
"""Test generate_content with both model and temperature overrides"""
|
||||||
|
# Arrange
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.status_code = 200
|
||||||
|
mock_response.json.return_value = {
|
||||||
|
'choices': [{
|
||||||
|
'message': {
|
||||||
|
'content': 'Response with both parameters override'
|
||||||
|
}
|
||||||
|
}],
|
||||||
|
'usage': {
|
||||||
|
'prompt_tokens': 18,
|
||||||
|
'completion_tokens': 12
|
||||||
|
}
|
||||||
|
}
|
||||||
|
mock_requests.post.return_value = mock_response
|
||||||
|
|
||||||
|
mock_async_init.return_value = None
|
||||||
|
mock_llm_init.return_value = None
|
||||||
|
|
||||||
|
config = {
|
||||||
|
'endpoint': 'https://test.inference.ai.azure.com/v1/chat/completions',
|
||||||
|
'token': 'test-token',
|
||||||
|
'temperature': 0.0,
|
||||||
|
'max_output': 4192,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override both parameters
|
||||||
|
result = await processor.generate_content("System", "Prompt", model="override-model", temperature=0.9)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.model == "override-model"
|
||||||
|
assert result.text == "Response with both parameters override"
|
||||||
|
|
||||||
|
# Verify the request was made with overridden temperature
|
||||||
|
mock_requests.post.assert_called_once()
|
||||||
|
call_args = mock_requests.post.call_args
|
||||||
|
|
||||||
|
import json
|
||||||
|
request_body = json.loads(call_args[1]['data'])
|
||||||
|
assert request_body['temperature'] == 0.9
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
pytest.main([__file__])
|
pytest.main([__file__])
|
||||||
|
|
@ -458,5 +458,132 @@ class TestLlamaFileProcessorSimple(IsolatedAsyncioTestCase):
|
||||||
# No specific rate limit error handling tested since SLM presumably has no rate limits
|
# No specific rate limit error handling tested since SLM presumably has no rate limits
|
||||||
|
|
||||||
|
|
||||||
|
@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_with_model_override(self, mock_llm_init, mock_async_init, mock_openai_class):
|
||||||
|
"""Test generate_content with model parameter override"""
|
||||||
|
# Arrange
|
||||||
|
mock_openai_client = MagicMock()
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.choices = [MagicMock()]
|
||||||
|
mock_response.choices[0].message.content = "Response from overridden model"
|
||||||
|
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': 'LLaMA_CPP',
|
||||||
|
'llamafile': 'http://localhost:8080/v1',
|
||||||
|
'temperature': 0.0,
|
||||||
|
'max_output': 4096,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override model
|
||||||
|
result = await processor.generate_content("System", "Prompt", model="custom-llamafile-model")
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.model == "custom-llamafile-model" # Should use overridden model
|
||||||
|
assert result.text == "Response from overridden model"
|
||||||
|
|
||||||
|
# Verify the API call was made with overridden model
|
||||||
|
call_args = mock_openai_client.chat.completions.create.call_args
|
||||||
|
assert call_args[1]['model'] == "custom-llamafile-model"
|
||||||
|
|
||||||
|
@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_with_temperature_override(self, mock_llm_init, mock_async_init, mock_openai_class):
|
||||||
|
"""Test generate_content with temperature parameter override"""
|
||||||
|
# Arrange
|
||||||
|
mock_openai_client = MagicMock()
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.choices = [MagicMock()]
|
||||||
|
mock_response.choices[0].message.content = "Response with temperature override"
|
||||||
|
mock_response.usage.prompt_tokens = 18
|
||||||
|
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, # Default temperature
|
||||||
|
'max_output': 4096,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override temperature
|
||||||
|
result = await processor.generate_content("System", "Prompt", temperature=0.7)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.text == "Response with temperature override"
|
||||||
|
|
||||||
|
# Verify the API call was made with overridden temperature
|
||||||
|
call_args = mock_openai_client.chat.completions.create.call_args
|
||||||
|
assert call_args[1]['temperature'] == 0.7
|
||||||
|
|
||||||
|
@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_with_both_parameters_override(self, mock_llm_init, mock_async_init, mock_openai_class):
|
||||||
|
"""Test generate_content with both model and temperature overrides"""
|
||||||
|
# Arrange
|
||||||
|
mock_openai_client = MagicMock()
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.choices = [MagicMock()]
|
||||||
|
mock_response.choices[0].message.content = "Response with both parameters override"
|
||||||
|
mock_response.usage.prompt_tokens = 20
|
||||||
|
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 - Override both parameters
|
||||||
|
result = await processor.generate_content("System", "Prompt", model="override-model", temperature=0.8)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.model == "override-model"
|
||||||
|
assert result.text == "Response with both parameters override"
|
||||||
|
|
||||||
|
# Verify the API call was made with overridden parameters
|
||||||
|
call_args = mock_openai_client.chat.completions.create.call_args
|
||||||
|
assert call_args[1]['model'] == "override-model"
|
||||||
|
assert call_args[1]['temperature'] == 0.8
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
pytest.main([__file__])
|
pytest.main([__file__])
|
||||||
|
|
@ -485,5 +485,148 @@ class TestVLLMProcessorSimple(IsolatedAsyncioTestCase):
|
||||||
assert call_args[1]['json']['prompt'] == expected_prompt
|
assert call_args[1]['json']['prompt'] == expected_prompt
|
||||||
|
|
||||||
|
|
||||||
|
@patch('trustgraph.model.text_completion.vllm.llm.aiohttp.ClientSession')
|
||||||
|
@patch('trustgraph.base.async_processor.AsyncProcessor.__init__')
|
||||||
|
@patch('trustgraph.base.llm_service.LlmService.__init__')
|
||||||
|
async def test_generate_content_with_model_override(self, mock_llm_init, mock_async_init, mock_session_class):
|
||||||
|
"""Test generate_content with model parameter override"""
|
||||||
|
# Arrange
|
||||||
|
mock_session = MagicMock()
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.status = 200
|
||||||
|
mock_response.json = AsyncMock(return_value={
|
||||||
|
'choices': [{
|
||||||
|
'text': 'Response from overridden model'
|
||||||
|
}],
|
||||||
|
'usage': {
|
||||||
|
'prompt_tokens': 12,
|
||||||
|
'completion_tokens': 8
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
mock_session.post.return_value.__aenter__.return_value = mock_response
|
||||||
|
mock_session.post.return_value.__aexit__.return_value = None
|
||||||
|
mock_session_class.return_value = mock_session
|
||||||
|
|
||||||
|
mock_async_init.return_value = None
|
||||||
|
mock_llm_init.return_value = None
|
||||||
|
|
||||||
|
config = {
|
||||||
|
'model': 'TheBloke/Mistral-7B-v0.1-AWQ',
|
||||||
|
'url': 'http://vllm-service:8899/v1',
|
||||||
|
'temperature': 0.0,
|
||||||
|
'max_output': 2048,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override model
|
||||||
|
result = await processor.generate_content("System", "Prompt", model="custom-vllm-model")
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.model == "custom-vllm-model" # Should use overridden model
|
||||||
|
assert result.text == "Response from overridden model"
|
||||||
|
|
||||||
|
@patch('trustgraph.model.text_completion.vllm.llm.aiohttp.ClientSession')
|
||||||
|
@patch('trustgraph.base.async_processor.AsyncProcessor.__init__')
|
||||||
|
@patch('trustgraph.base.llm_service.LlmService.__init__')
|
||||||
|
async def test_generate_content_with_temperature_override(self, mock_llm_init, mock_async_init, mock_session_class):
|
||||||
|
"""Test generate_content with temperature parameter override"""
|
||||||
|
# Arrange
|
||||||
|
mock_session = MagicMock()
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.status = 200
|
||||||
|
mock_response.json = AsyncMock(return_value={
|
||||||
|
'choices': [{
|
||||||
|
'text': 'Response with temperature override'
|
||||||
|
}],
|
||||||
|
'usage': {
|
||||||
|
'prompt_tokens': 15,
|
||||||
|
'completion_tokens': 10
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
mock_session.post.return_value.__aenter__.return_value = mock_response
|
||||||
|
mock_session.post.return_value.__aexit__.return_value = None
|
||||||
|
mock_session_class.return_value = mock_session
|
||||||
|
|
||||||
|
mock_async_init.return_value = None
|
||||||
|
mock_llm_init.return_value = None
|
||||||
|
|
||||||
|
config = {
|
||||||
|
'model': 'TheBloke/Mistral-7B-v0.1-AWQ',
|
||||||
|
'url': 'http://vllm-service:8899/v1',
|
||||||
|
'temperature': 0.0, # Default temperature
|
||||||
|
'max_output': 2048,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override temperature
|
||||||
|
result = await processor.generate_content("System", "Prompt", temperature=0.7)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.text == "Response with temperature override"
|
||||||
|
|
||||||
|
# Verify the request was made with overridden temperature
|
||||||
|
call_args = mock_session.post.call_args
|
||||||
|
assert call_args[1]['json']['temperature'] == 0.7
|
||||||
|
|
||||||
|
@patch('trustgraph.model.text_completion.vllm.llm.aiohttp.ClientSession')
|
||||||
|
@patch('trustgraph.base.async_processor.AsyncProcessor.__init__')
|
||||||
|
@patch('trustgraph.base.llm_service.LlmService.__init__')
|
||||||
|
async def test_generate_content_with_both_parameters_override(self, mock_llm_init, mock_async_init, mock_session_class):
|
||||||
|
"""Test generate_content with both model and temperature overrides"""
|
||||||
|
# Arrange
|
||||||
|
mock_session = MagicMock()
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.status = 200
|
||||||
|
mock_response.json = AsyncMock(return_value={
|
||||||
|
'choices': [{
|
||||||
|
'text': 'Response with both parameters override'
|
||||||
|
}],
|
||||||
|
'usage': {
|
||||||
|
'prompt_tokens': 18,
|
||||||
|
'completion_tokens': 12
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
mock_session.post.return_value.__aenter__.return_value = mock_response
|
||||||
|
mock_session.post.return_value.__aexit__.return_value = None
|
||||||
|
mock_session_class.return_value = mock_session
|
||||||
|
|
||||||
|
mock_async_init.return_value = None
|
||||||
|
mock_llm_init.return_value = None
|
||||||
|
|
||||||
|
config = {
|
||||||
|
'model': 'TheBloke/Mistral-7B-v0.1-AWQ',
|
||||||
|
'url': 'http://vllm-service:8899/v1',
|
||||||
|
'temperature': 0.0,
|
||||||
|
'max_output': 2048,
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
processor = Processor(**config)
|
||||||
|
|
||||||
|
# Act - Override both parameters
|
||||||
|
result = await processor.generate_content("System", "Prompt", model="override-model", temperature=0.8)
|
||||||
|
|
||||||
|
# Assert
|
||||||
|
assert result.model == "override-model"
|
||||||
|
assert result.text == "Response with both parameters override"
|
||||||
|
|
||||||
|
# Verify the request was made with overridden temperature
|
||||||
|
call_args = mock_session.post.call_args
|
||||||
|
assert call_args[1]['json']['temperature'] == 0.8
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
pytest.main([__file__])
|
pytest.main([__file__])
|
||||||
Loading…
Add table
Add a link
Reference in a new issue