mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-23 12:11:02 +02:00
Commonality extracted
This commit is contained in:
parent
d25c2f4d10
commit
c496c53ff7
4 changed files with 225 additions and 75 deletions
3
tests/unit/test_text_completion/common/__init__.py
Normal file
3
tests/unit/test_text_completion/common/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
"""
|
||||||
|
Common utilities for text completion tests
|
||||||
|
"""
|
||||||
69
tests/unit/test_text_completion/common/base_test_cases.py
Normal file
69
tests/unit/test_text_completion/common/base_test_cases.py
Normal file
|
|
@ -0,0 +1,69 @@
|
||||||
|
"""
|
||||||
|
Base test patterns that can be reused across different text completion models
|
||||||
|
"""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
from unittest import IsolatedAsyncioTestCase
|
||||||
|
|
||||||
|
|
||||||
|
class BaseTextCompletionTestCase(IsolatedAsyncioTestCase, ABC):
|
||||||
|
"""
|
||||||
|
Base test class for text completion processors
|
||||||
|
Provides common test patterns that can be reused
|
||||||
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_processor_class(self):
|
||||||
|
"""Return the processor class to test"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_base_config(self):
|
||||||
|
"""Return base configuration for the processor"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_mock_patches(self):
|
||||||
|
"""Return list of patch decorators for mocking dependencies"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def create_base_config(self, **overrides):
|
||||||
|
"""Create base config with optional overrides"""
|
||||||
|
config = self.get_base_config()
|
||||||
|
config.update(overrides)
|
||||||
|
return config
|
||||||
|
|
||||||
|
def create_mock_llm_result(self, text="Test response", in_token=10, out_token=5):
|
||||||
|
"""Create a mock LLM result"""
|
||||||
|
from trustgraph.base import LlmResult
|
||||||
|
return LlmResult(text=text, in_token=in_token, out_token=out_token)
|
||||||
|
|
||||||
|
|
||||||
|
class CommonTestPatterns:
|
||||||
|
"""
|
||||||
|
Common test patterns that can be used across different models
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def basic_initialization_test_pattern(test_instance):
|
||||||
|
"""
|
||||||
|
Test pattern for basic processor initialization
|
||||||
|
test_instance should be a BaseTextCompletionTestCase
|
||||||
|
"""
|
||||||
|
# This would contain the common pattern for initialization testing
|
||||||
|
pass
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def successful_generation_test_pattern(test_instance):
|
||||||
|
"""
|
||||||
|
Test pattern for successful content generation
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def error_handling_test_pattern(test_instance):
|
||||||
|
"""
|
||||||
|
Test pattern for error handling
|
||||||
|
"""
|
||||||
|
pass
|
||||||
53
tests/unit/test_text_completion/common/mock_helpers.py
Normal file
53
tests/unit/test_text_completion/common/mock_helpers.py
Normal file
|
|
@ -0,0 +1,53 @@
|
||||||
|
"""
|
||||||
|
Common mocking utilities for text completion tests
|
||||||
|
"""
|
||||||
|
|
||||||
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
|
||||||
|
class CommonMocks:
|
||||||
|
"""Common mock objects used across text completion tests"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def create_mock_async_processor_init():
|
||||||
|
"""Create mock for AsyncProcessor.__init__"""
|
||||||
|
mock = MagicMock()
|
||||||
|
mock.return_value = None
|
||||||
|
return mock
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def create_mock_llm_service_init():
|
||||||
|
"""Create mock for LlmService.__init__"""
|
||||||
|
mock = MagicMock()
|
||||||
|
mock.return_value = None
|
||||||
|
return mock
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def create_mock_response(text="Test response", prompt_tokens=10, completion_tokens=5):
|
||||||
|
"""Create a mock response object"""
|
||||||
|
response = MagicMock()
|
||||||
|
response.text = text
|
||||||
|
response.usage_metadata.prompt_token_count = prompt_tokens
|
||||||
|
response.usage_metadata.candidates_token_count = completion_tokens
|
||||||
|
return response
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def create_basic_config():
|
||||||
|
"""Create basic config with required fields"""
|
||||||
|
return {
|
||||||
|
'concurrency': 1,
|
||||||
|
'taskgroup': AsyncMock(),
|
||||||
|
'id': 'test-processor'
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class MockPatches:
|
||||||
|
"""Common patch decorators for different services"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_base_patches():
|
||||||
|
"""Get patches that are common to all processors"""
|
||||||
|
return [
|
||||||
|
'trustgraph.base.async_processor.AsyncProcessor.__init__',
|
||||||
|
'trustgraph.base.llm_service.LlmService.__init__'
|
||||||
|
]
|
||||||
|
|
@ -4,49 +4,19 @@ Pytest configuration and fixtures for text completion tests
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from unittest.mock import MagicMock, AsyncMock
|
from unittest.mock import MagicMock, AsyncMock
|
||||||
from trustgraph.schema import TextCompletionRequest, TextCompletionResponse
|
|
||||||
from trustgraph.base import LlmResult
|
from trustgraph.base import LlmResult
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def mock_vertexai_credentials():
|
|
||||||
"""Mock Google Cloud service account credentials"""
|
|
||||||
return MagicMock()
|
|
||||||
|
|
||||||
|
# === Common Fixtures for All Text Completion Models ===
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_vertexai_model():
|
def base_processor_config():
|
||||||
"""Mock VertexAI GenerativeModel"""
|
"""Base configuration required by all processors"""
|
||||||
mock_model = MagicMock()
|
return {
|
||||||
mock_response = MagicMock()
|
'concurrency': 1,
|
||||||
mock_response.text = "Test response"
|
'taskgroup': AsyncMock(),
|
||||||
mock_response.usage_metadata.prompt_token_count = 10
|
'id': 'test-processor'
|
||||||
mock_response.usage_metadata.candidates_token_count = 5
|
}
|
||||||
mock_model.generate_content.return_value = mock_response
|
|
||||||
return mock_model
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def sample_text_completion_request():
|
|
||||||
"""Sample TextCompletionRequest for testing"""
|
|
||||||
return TextCompletionRequest(
|
|
||||||
id="test-request-id",
|
|
||||||
prompt="Test prompt",
|
|
||||||
system="Test system prompt",
|
|
||||||
temperature=0.7,
|
|
||||||
max_output=1024
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def sample_text_completion_response():
|
|
||||||
"""Sample TextCompletionResponse for testing"""
|
|
||||||
return TextCompletionResponse(
|
|
||||||
id="test-response-id",
|
|
||||||
response="Test response",
|
|
||||||
in_token=10,
|
|
||||||
out_token=5,
|
|
||||||
model="gemini-2.0-flash-001"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
|
|
@ -60,16 +30,19 @@ def sample_llm_result():
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def vertexai_processor_config():
|
def mock_async_processor_init():
|
||||||
"""Default configuration for VertexAI processor"""
|
"""Mock AsyncProcessor.__init__ to avoid infrastructure requirements"""
|
||||||
return {
|
mock = MagicMock()
|
||||||
'region': 'us-central1',
|
mock.return_value = None
|
||||||
'model': 'gemini-2.0-flash-001',
|
return mock
|
||||||
'temperature': 0.0,
|
|
||||||
'max_output': 8192,
|
|
||||||
'private_key': 'private.json',
|
@pytest.fixture
|
||||||
'concurrency': 1
|
def mock_llm_service_init():
|
||||||
}
|
"""Mock LlmService.__init__ to avoid infrastructure requirements"""
|
||||||
|
mock = MagicMock()
|
||||||
|
mock.return_value = None
|
||||||
|
return mock
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
|
|
@ -92,21 +65,66 @@ def mock_pulsar_producer():
|
||||||
return AsyncMock()
|
return AsyncMock()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def mock_env_vars(monkeypatch):
|
||||||
|
"""Mock environment variables for testing"""
|
||||||
|
monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project")
|
||||||
|
monkeypatch.setenv("GOOGLE_APPLICATION_CREDENTIALS", "/path/to/test-credentials.json")
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_flow_processor_config():
|
def mock_async_context_manager():
|
||||||
"""Mock flow processor configuration"""
|
"""Mock async context manager for testing"""
|
||||||
return {
|
class MockAsyncContextManager:
|
||||||
'service_id': 'test-vertexai-service',
|
def __init__(self, return_value):
|
||||||
'flow_name': 'test-flow',
|
self.return_value = return_value
|
||||||
'consumer_name': 'test-consumer'
|
|
||||||
}
|
async def __aenter__(self):
|
||||||
|
return self.return_value
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||||
|
pass
|
||||||
|
|
||||||
|
return MockAsyncContextManager
|
||||||
|
|
||||||
|
|
||||||
|
# === VertexAI Specific Fixtures ===
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def mock_vertexai_credentials():
|
||||||
|
"""Mock Google Cloud service account credentials"""
|
||||||
|
return MagicMock()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def mock_vertexai_model():
|
||||||
|
"""Mock VertexAI GenerativeModel"""
|
||||||
|
mock_model = MagicMock()
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.text = "Test response"
|
||||||
|
mock_response.usage_metadata.prompt_token_count = 10
|
||||||
|
mock_response.usage_metadata.candidates_token_count = 5
|
||||||
|
mock_model.generate_content.return_value = mock_response
|
||||||
|
return mock_model
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def vertexai_processor_config(base_processor_config):
|
||||||
|
"""Default configuration for VertexAI processor"""
|
||||||
|
config = base_processor_config.copy()
|
||||||
|
config.update({
|
||||||
|
'region': 'us-central1',
|
||||||
|
'model': 'gemini-2.0-flash-001',
|
||||||
|
'temperature': 0.0,
|
||||||
|
'max_output': 8192,
|
||||||
|
'private_key': 'private.json'
|
||||||
|
})
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_safety_settings():
|
def mock_safety_settings():
|
||||||
"""Mock safety settings for VertexAI"""
|
"""Mock safety settings for VertexAI"""
|
||||||
from unittest.mock import MagicMock
|
|
||||||
|
|
||||||
safety_settings = []
|
safety_settings = []
|
||||||
for i in range(4): # 4 safety categories
|
for i in range(4): # 4 safety categories
|
||||||
setting = MagicMock()
|
setting = MagicMock()
|
||||||
|
|
@ -136,24 +154,31 @@ def mock_vertexai_exception():
|
||||||
return ResourceExhausted("Test resource exhausted error")
|
return ResourceExhausted("Test resource exhausted error")
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(autouse=True)
|
# === Ollama Specific Fixtures (for next implementation) ===
|
||||||
def mock_env_vars(monkeypatch):
|
|
||||||
"""Mock environment variables for testing"""
|
@pytest.fixture
|
||||||
monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project")
|
def ollama_processor_config(base_processor_config):
|
||||||
monkeypatch.setenv("GOOGLE_APPLICATION_CREDENTIALS", "/path/to/test-credentials.json")
|
"""Default configuration for Ollama processor"""
|
||||||
|
config = base_processor_config.copy()
|
||||||
|
config.update({
|
||||||
|
'model': 'llama2',
|
||||||
|
'temperature': 0.0,
|
||||||
|
'max_output': 8192,
|
||||||
|
'host': 'localhost',
|
||||||
|
'port': 11434
|
||||||
|
})
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_async_context_manager():
|
def mock_ollama_client():
|
||||||
"""Mock async context manager for testing"""
|
"""Mock Ollama client"""
|
||||||
class MockAsyncContextManager:
|
mock_client = MagicMock()
|
||||||
def __init__(self, return_value):
|
mock_response = {
|
||||||
self.return_value = return_value
|
'response': 'Test response from Ollama',
|
||||||
|
'done': True,
|
||||||
async def __aenter__(self):
|
'eval_count': 5,
|
||||||
return self.return_value
|
'prompt_eval_count': 10
|
||||||
|
}
|
||||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
mock_client.generate.return_value = mock_response
|
||||||
pass
|
return mock_client
|
||||||
|
|
||||||
return MockAsyncContextManager
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue