feat: integrate PluginStore and MinIO repository into model manager activities
- Added PluginStore integration for model management. - Replaced StorageRepository with MinIORepository in Activities, Cleanup, and Training classes. - Updated training logic to handle validation files and improved data management. - Enhanced configuration for MinIO and PluginStore in connectors. - Removed deprecated model repository and storage repository files. - Updated environment variable handling for new configurations.
This commit is contained in:
@@ -1,316 +0,0 @@
|
||||
"""Unit tests for Activities class with 100% coverage."""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_logger():
|
||||
"""Create a mock logger."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_notification_handler():
|
||||
"""Create a mock notification handler."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def postgres_config():
|
||||
"""Create a valid PostgreSQL configuration."""
|
||||
return {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
'user': 'testuser',
|
||||
'password': 'testpass',
|
||||
'dbname': 'testdb',
|
||||
'min_connections': 1,
|
||||
'max_connections': 10,
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mlflow_config():
|
||||
"""Create a valid MLFlow configuration."""
|
||||
return {
|
||||
'url': 'http://mlflow:5080',
|
||||
'username': 'aignosi',
|
||||
'password': 'aignosi',
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def minio_config():
|
||||
"""Create a valid MinIO configuration."""
|
||||
return {
|
||||
'endpoint_url': 'http://minio:9000',
|
||||
'access_key': 'minioadmin',
|
||||
'secret_key': 'minioadmin',
|
||||
'region': 'us-east-1',
|
||||
'use_ssl': False,
|
||||
'max_retry_attempts': 3,
|
||||
'retry_mode': 'standard',
|
||||
'connect_timeout': 5,
|
||||
'read_timeout': 5,
|
||||
}
|
||||
|
||||
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.Training.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.ModelRepository')
|
||||
@patch('model_manager.activities.activities.StorageRepository')
|
||||
def test_activities_init_success(
|
||||
mock_storage_repo,
|
||||
mock_model_repo,
|
||||
mock_training_init,
|
||||
mock_et_init,
|
||||
postgres_config,
|
||||
mlflow_config,
|
||||
minio_config,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
):
|
||||
"""Test successful initialization of Activities."""
|
||||
from model_manager.activities.activities import Activities
|
||||
|
||||
activities = Activities(
|
||||
postgres_config=postgres_config,
|
||||
mlflow_config=mlflow_config,
|
||||
minio_config=minio_config,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
)
|
||||
|
||||
mock_et_init.assert_called_once()
|
||||
assert mock_et_init.call_args[1]['host'] == postgres_config['host']
|
||||
assert mock_et_init.call_args[1]['port'] == postgres_config['port']
|
||||
assert mock_et_init.call_args[1]['user'] == postgres_config['user']
|
||||
assert mock_et_init.call_args[1]['password'] == postgres_config['password']
|
||||
assert mock_et_init.call_args[1]['dbname'] == postgres_config['dbname']
|
||||
assert mock_et_init.call_args[1]['min_connections'] == postgres_config['min_connections']
|
||||
assert mock_et_init.call_args[1]['max_connections'] == postgres_config['max_connections']
|
||||
assert mock_et_init.call_args[1]['logger'] is mock_logger
|
||||
assert mock_et_init.call_args[1]['notification_handler'] is mock_notification_handler
|
||||
|
||||
mock_model_repo.assert_called_once_with(
|
||||
url=mlflow_config['url'],
|
||||
username=mlflow_config['username'],
|
||||
password=mlflow_config['password'],
|
||||
logger=mock_logger,
|
||||
)
|
||||
|
||||
mock_storage_repo.assert_called_once_with(
|
||||
endpoint_url=minio_config['endpoint_url'],
|
||||
access_key=minio_config['access_key'],
|
||||
secret_key=minio_config['secret_key'],
|
||||
region=minio_config['region'],
|
||||
use_ssl=minio_config['use_ssl'],
|
||||
max_retry_attempts=minio_config['max_retry_attempts'],
|
||||
retry_mode=minio_config['retry_mode'],
|
||||
connect_timeout=minio_config['connect_timeout'],
|
||||
read_timeout=minio_config['read_timeout'],
|
||||
logger=mock_logger,
|
||||
)
|
||||
|
||||
mock_training_init.assert_called_once()
|
||||
assert mock_training_init.call_args[1]['model_repository'] is mock_model_repo.return_value
|
||||
assert mock_training_init.call_args[1]['storage_repository'] is mock_storage_repo.return_value
|
||||
assert mock_training_init.call_args[1]['logger'] is mock_logger
|
||||
assert mock_training_init.call_args[1]['notification_handler'] is mock_notification_handler
|
||||
|
||||
assert hasattr(activities, 'model_repository')
|
||||
assert hasattr(activities, 'storage_repository')
|
||||
|
||||
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.close')
|
||||
@patch('model_manager.activities.activities.Training.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.ModelRepository')
|
||||
@patch('model_manager.activities.activities.StorageRepository')
|
||||
def test_activities_shutdown(
|
||||
mock_storage_repo,
|
||||
mock_model_repo,
|
||||
mock_training_init,
|
||||
mock_et_close,
|
||||
mock_et_init,
|
||||
postgres_config,
|
||||
mlflow_config,
|
||||
minio_config,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
):
|
||||
"""Test Activities.shutdown() calls ExperimentTracking.close()."""
|
||||
from model_manager.activities.activities import Activities
|
||||
|
||||
activities = Activities(
|
||||
postgres_config=postgres_config,
|
||||
mlflow_config=mlflow_config,
|
||||
minio_config=minio_config,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
)
|
||||
|
||||
asyncio.run(activities.shutdown())
|
||||
|
||||
mock_et_close.assert_called_once_with(activities)
|
||||
|
||||
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.Training.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.ModelRepository')
|
||||
@patch('model_manager.activities.activities.StorageRepository')
|
||||
def test_activities_del_without_engine(
|
||||
mock_storage_repo,
|
||||
mock_model_repo,
|
||||
mock_training_init,
|
||||
mock_et_init,
|
||||
postgres_config,
|
||||
mlflow_config,
|
||||
minio_config,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
):
|
||||
"""Test __del__ when engine attribute does not exist."""
|
||||
from model_manager.activities.activities import Activities
|
||||
|
||||
activities = Activities(
|
||||
postgres_config=postgres_config,
|
||||
mlflow_config=mlflow_config,
|
||||
minio_config=minio_config,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
)
|
||||
|
||||
if hasattr(activities, 'engine'):
|
||||
delattr(activities, 'engine')
|
||||
|
||||
activities.__del__()
|
||||
|
||||
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.Training.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.ModelRepository')
|
||||
@patch('model_manager.activities.activities.StorageRepository')
|
||||
def test_activities_del_with_engine_no_super_del(
|
||||
mock_storage_repo,
|
||||
mock_model_repo,
|
||||
mock_training_init,
|
||||
mock_et_init,
|
||||
postgres_config,
|
||||
mlflow_config,
|
||||
minio_config,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
):
|
||||
"""Test __del__ when engine exists but super has no __del__."""
|
||||
from model_manager.activities.activities import Activities
|
||||
|
||||
activities = Activities(
|
||||
postgres_config=postgres_config,
|
||||
mlflow_config=mlflow_config,
|
||||
minio_config=minio_config,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
)
|
||||
|
||||
activities.engine = MagicMock()
|
||||
|
||||
with patch('builtins.super') as mock_super:
|
||||
mock_super_instance = MagicMock()
|
||||
del mock_super_instance.__del__
|
||||
mock_super.return_value = mock_super_instance
|
||||
|
||||
activities.__del__()
|
||||
|
||||
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.Training.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.ModelRepository')
|
||||
@patch('model_manager.activities.activities.StorageRepository')
|
||||
def test_activities_del_with_engine_and_super_del(
|
||||
mock_storage_repo,
|
||||
mock_model_repo,
|
||||
mock_training_init,
|
||||
mock_et_init,
|
||||
postgres_config,
|
||||
mlflow_config,
|
||||
minio_config,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
):
|
||||
"""Test __del__ when engine exists and super has __del__."""
|
||||
from model_manager.activities.activities import Activities
|
||||
|
||||
activities = Activities(
|
||||
postgres_config=postgres_config,
|
||||
mlflow_config=mlflow_config,
|
||||
minio_config=minio_config,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
)
|
||||
|
||||
activities.engine = MagicMock()
|
||||
|
||||
mock_super_del = MagicMock()
|
||||
|
||||
class MockSuper:
|
||||
def __del__(self):
|
||||
mock_super_del()
|
||||
|
||||
with patch('builtins.super', return_value=MockSuper()):
|
||||
activities.__del__()
|
||||
|
||||
mock_super_del.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.Training.__init__', return_value=None)
|
||||
@patch('model_manager.activities.activities.ModelRepository')
|
||||
@patch('model_manager.activities.activities.StorageRepository')
|
||||
def test_activities_del_with_engine_exception_caught(
|
||||
mock_storage_repo,
|
||||
mock_model_repo,
|
||||
mock_training_init,
|
||||
mock_et_init,
|
||||
postgres_config,
|
||||
mlflow_config,
|
||||
minio_config,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
):
|
||||
"""Test __del__ catches exceptions when super().__del__() raises."""
|
||||
from model_manager.activities.activities import Activities
|
||||
|
||||
activities = Activities(
|
||||
postgres_config=postgres_config,
|
||||
mlflow_config=mlflow_config,
|
||||
minio_config=minio_config,
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
)
|
||||
|
||||
activities.engine = MagicMock()
|
||||
|
||||
class MockSuperWithError:
|
||||
def __del__(self):
|
||||
# Only raise error if not being cleaned up by garbage collector
|
||||
# This prevents the PytestUnraisableExceptionWarning
|
||||
if hasattr(self, '_should_raise') and self._should_raise:
|
||||
raise RuntimeError('Test error')
|
||||
|
||||
# Suppress the PytestUnraisableExceptionWarning for this specific test
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
||||
|
||||
mock_super = MockSuperWithError()
|
||||
mock_super._should_raise = True
|
||||
try:
|
||||
with patch('builtins.super', return_value=mock_super):
|
||||
activities.__del__()
|
||||
finally:
|
||||
# Prevent the exception from being raised during garbage collection
|
||||
mock_super._should_raise = False
|
||||
@@ -61,7 +61,6 @@ def test_train_model_result_creation(sample_params, sample_dataframes):
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
process_data = MagicMock()
|
||||
regr = MagicMock()
|
||||
scaler_dict = {'var1': {'min': 0, 'max': 100}}
|
||||
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
@@ -71,7 +70,6 @@ def test_train_model_result_creation(sample_params, sample_dataframes):
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=regr,
|
||||
scaler_dict=scaler_dict,
|
||||
)
|
||||
|
||||
assert result.params == sample_params
|
||||
@@ -96,7 +94,6 @@ def test_train_model_result_optional_fields_default_none(sample_params, sample_d
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
scaler_dict={},
|
||||
)
|
||||
|
||||
assert result.y_pred is None
|
||||
@@ -123,7 +120,6 @@ def test_train_model_result_with_metrics(sample_params, sample_dataframes):
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
scaler_dict={},
|
||||
y_pred=y_pred,
|
||||
mse_val=1.5,
|
||||
mae_val=1.2,
|
||||
@@ -148,7 +144,6 @@ def test_train_model_result_with_artifact_paths(sample_params, sample_dataframes
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
scaler_dict={},
|
||||
run_name='test-experiment-1',
|
||||
report_path='/path/to/report.html',
|
||||
train_data_path='/path/to/train_data.csv',
|
||||
@@ -186,11 +181,11 @@ def test_train_model_result_is_dataclass(sample_params, sample_dataframes):
|
||||
|
||||
|
||||
def test_train_model_result_field_count():
|
||||
"""Test that TrainModelResult has exactly 20 fields."""
|
||||
"""Test that TrainModelResult has exactly 19 fields."""
|
||||
from dataclasses import fields
|
||||
|
||||
result_fields = fields(TrainModelResult)
|
||||
assert len(result_fields) == 20
|
||||
assert len(result_fields) == 19
|
||||
|
||||
field_names = {f.name for f in result_fields}
|
||||
expected_fields = {
|
||||
@@ -201,7 +196,6 @@ def test_train_model_result_field_count():
|
||||
'y_train',
|
||||
'y_test',
|
||||
'regr',
|
||||
'scaler_dict',
|
||||
'y_pred',
|
||||
'y_train_pred',
|
||||
'mse_val',
|
||||
@@ -231,7 +225,6 @@ def test_train_model_result_complete_workflow(sample_params, sample_dataframes):
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
scaler_dict={'var1': {'min': 0, 'max': 100}},
|
||||
)
|
||||
|
||||
# Step 2: Add predictions and metrics
|
||||
|
||||
@@ -1,982 +0,0 @@
|
||||
"""Unit tests for ModelRepository with 100% coverage."""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def cleanup_temp_directories():
|
||||
"""Clean up temporary directories after each test."""
|
||||
# Get the temp directory path
|
||||
current_file_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
model_manager_dir = os.path.dirname(os.path.dirname(os.path.dirname(current_file_dir)))
|
||||
temp_dir = os.path.join(model_manager_dir, 'reports', 'temp')
|
||||
|
||||
# Run the test
|
||||
yield
|
||||
|
||||
# Clean up after test
|
||||
if os.path.exists(temp_dir):
|
||||
for item in os.listdir(temp_dir):
|
||||
item_path = os.path.join(temp_dir, item)
|
||||
if os.path.isdir(item_path) and item.startswith('test_run_'):
|
||||
try:
|
||||
shutil.rmtree(item_path)
|
||||
except (OSError, PermissionError):
|
||||
# Ignore cleanup errors
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_logger():
|
||||
"""Create a mock logger."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_train_result():
|
||||
"""Create a mock TrainModelResult."""
|
||||
result = MagicMock()
|
||||
result.params = MagicMock()
|
||||
result.params.experiment_name = 'test_experiment'
|
||||
result.params.experiment_run_id = 1
|
||||
result.params.target_variable = 'target'
|
||||
result.params.variable_columns = ['var1', 'var2']
|
||||
result.params.lag_train = 5
|
||||
result.params.lag_val = 3
|
||||
result.params.window = 10
|
||||
result.params.low_lim = {'var1': 0.0}
|
||||
result.params.upp_lim = {'var1': 10.0}
|
||||
result.params.include_ar = False
|
||||
result.params.train_size = 80
|
||||
result.params.removed_intervals = []
|
||||
result.params.rem_static_win = True
|
||||
result.params.static_threshold = None
|
||||
result.run_name = 'test_run'
|
||||
result.run_dir = '/tmp/test_run' # noqa: S108
|
||||
result.report_path = '/tmp/test_run/report.html' # noqa: S108
|
||||
result.train_data_path = '/tmp/test_run/train_data.csv' # noqa: S108
|
||||
result.test_data_path = '/tmp/test_run/test_data.csv' # noqa: S108
|
||||
result.mse_val = 0.5
|
||||
result.r2_val = 0.9
|
||||
result.mae_val = 0.3
|
||||
result.scaler_dict = {'scaler': 'standard'}
|
||||
result.process_data = MagicMock()
|
||||
result.regr = MagicMock()
|
||||
result.regr.predict = MagicMock(return_value=np.array([1.0, 2.0, 3.0]))
|
||||
result.x_train = pd.DataFrame({'var1': [1, 2, 3]})
|
||||
result.y_train = pd.Series([1.0, 2.0, 3.0], name='target')
|
||||
result.x_test = pd.DataFrame({'var1': [4, 5, 6]})
|
||||
result.y_test = pd.Series([4.0, 5.0, 6.0], name='target')
|
||||
result.y_pred = np.array([4.1, 5.1, 6.1])
|
||||
return result
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_model_repository_init(mock_model_serving_class, mock_logger):
|
||||
"""Test ModelRepository initialization."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
mock_model_serving_instance = MagicMock()
|
||||
mock_model_serving_class.return_value = mock_model_serving_instance
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_model_serving_class.assert_called_once_with(
|
||||
tracking_uri='http://mlflow.test', username='user', password='pass'
|
||||
)
|
||||
assert repo.model_serving is mock_model_serving_instance
|
||||
assert repo.logger is mock_logger
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_save_model_success(mock_model_serving_class, mock_logger, mock_train_result):
|
||||
"""Test save_model successfully saves model."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Reset mock after initialization to focus on method-specific calls
|
||||
mock_logger.reset_mock()
|
||||
|
||||
repo._get_next_run_name = MagicMock(return_value='test_experiment-1')
|
||||
repo._generate_artifacts = MagicMock(return_value=mock_train_result)
|
||||
repo._save_run = MagicMock()
|
||||
|
||||
result = repo.save_model(mock_train_result)
|
||||
|
||||
repo._get_next_run_name.assert_called_once_with('test_experiment')
|
||||
repo._generate_artifacts.assert_called_once()
|
||||
repo._save_run.assert_called_once_with(mock_train_result)
|
||||
mock_logger.info.assert_called_once()
|
||||
assert result is mock_train_result
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.os.path.exists')
|
||||
@patch('model_manager.utils.repository.model_repository.shutil.rmtree')
|
||||
def test_cleanup_run_directory_exists(
|
||||
mock_rmtree, mock_exists, mock_model_serving_class, mock_logger
|
||||
):
|
||||
"""Test cleanup_run_directory when directory exists."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Reset mock after initialization to focus on method-specific calls
|
||||
mock_logger.reset_mock()
|
||||
|
||||
mock_exists.return_value = True
|
||||
|
||||
repo.cleanup_run_directory('/tmp/test_run') # noqa: S108
|
||||
|
||||
mock_exists.assert_called_once_with('/tmp/test_run') # noqa: S108
|
||||
mock_rmtree.assert_called_once_with('/tmp/test_run') # noqa: S108
|
||||
mock_logger.info.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.os.path.exists')
|
||||
def test_cleanup_run_directory_not_exists(mock_exists, mock_model_serving_class, mock_logger):
|
||||
"""Test cleanup_run_directory when directory doesn't exist."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Reset mock after initialization to focus on method-specific calls
|
||||
mock_logger.reset_mock()
|
||||
|
||||
mock_exists.return_value = False
|
||||
|
||||
repo.cleanup_run_directory('/tmp/test_run') # noqa: S108
|
||||
|
||||
mock_exists.assert_called_once_with('/tmp/test_run') # noqa: S108
|
||||
mock_logger.info.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_cleanup_run_directory_empty_path(mock_model_serving_class, mock_logger):
|
||||
"""Test cleanup_run_directory with empty path."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Reset mock after initialization to focus on method-specific calls
|
||||
mock_logger.reset_mock()
|
||||
|
||||
repo.cleanup_run_directory('')
|
||||
|
||||
mock_logger.info.assert_called_once_with('No run directory specified, skipping cleanup')
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_get_next_run_name_no_existing_runs(mock_model_serving_class, mock_logger):
|
||||
"""Test _get_next_run_name when no runs exist."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
mock_model_serving_instance = MagicMock()
|
||||
mock_model_serving_instance.search_runs_by_name.return_value = []
|
||||
mock_model_serving_class.return_value = mock_model_serving_instance
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
result = repo._get_next_run_name('test_experiment')
|
||||
|
||||
assert result == 'test_experiment-1'
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_get_next_run_name_with_existing_runs(mock_model_serving_class, mock_logger):
|
||||
"""Test _get_next_run_name when runs exist."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
mock_model_serving_instance = MagicMock()
|
||||
mock_model_serving_instance.search_runs_by_name.return_value = [
|
||||
MagicMock(),
|
||||
MagicMock(),
|
||||
MagicMock(),
|
||||
]
|
||||
mock_model_serving_class.return_value = mock_model_serving_instance
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
result = repo._get_next_run_name('test_experiment')
|
||||
|
||||
assert result == 'test_experiment-4'
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_get_reports_directory(mock_model_serving_class, mock_logger):
|
||||
"""Test _get_reports_directory returns correct path."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
result = repo._get_reports_directory()
|
||||
|
||||
assert result.endswith(os.path.join('model_manager', 'reports'))
|
||||
assert os.path.isabs(result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_init_artifacts_data_success(mock_model_serving_class, mock_logger, mock_train_result):
|
||||
"""Test _init_artifacts_data successfully prepares data."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
reference_data, current_data = repo._init_artifacts_data(mock_train_result)
|
||||
|
||||
# Check reference data
|
||||
assert 'target' in reference_data.columns
|
||||
assert 'prediction' in reference_data.columns
|
||||
assert len(reference_data) == 3
|
||||
|
||||
# Check current data
|
||||
assert 'target' in current_data.columns
|
||||
assert 'prediction' in current_data.columns
|
||||
assert len(current_data) == 3
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_init_artifacts_data_empty_x_train(
|
||||
mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _init_artifacts_data raises ValueError when x_train is empty."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.x_train = pd.DataFrame()
|
||||
|
||||
with pytest.raises(ValueError, match='Training features .* are empty'):
|
||||
repo._init_artifacts_data(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_init_artifacts_data_empty_y_train(
|
||||
mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _init_artifacts_data raises ValueError when y_train is empty."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.y_train = pd.Series(dtype=float)
|
||||
|
||||
with pytest.raises(ValueError, match='Training target .* is empty'):
|
||||
repo._init_artifacts_data(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_init_artifacts_data_none_y_pred(mock_model_serving_class, mock_logger, mock_train_result):
|
||||
"""Test _init_artifacts_data raises ValueError when y_pred is None."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.y_pred = None
|
||||
|
||||
with pytest.raises(ValueError, match='Test predictions .* are None'):
|
||||
repo._init_artifacts_data(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_init_artifacts_data_none_y_train_pred(
|
||||
mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _init_artifacts_data raises ValueError when y_train_pred is None."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.y_train_pred = None
|
||||
|
||||
with pytest.raises(ValueError, match='Training predictions .* are None'):
|
||||
repo._init_artifacts_data(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.datetime')
|
||||
@patch('model_manager.utils.repository.model_repository.makedirs')
|
||||
def test_create_run_directory_success(
|
||||
mock_makedirs, mock_datetime, mock_model_serving_class, mock_logger
|
||||
):
|
||||
"""Test _create_run_directory creates directory successfully."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_datetime.now.return_value.strftime.return_value = '20240101_120000_123456'
|
||||
|
||||
result = repo._create_run_directory('/tmp/reports', 'test_run') # noqa: S108
|
||||
|
||||
expected_path = os.path.normpath(
|
||||
os.path.join('/tmp/reports', 'temp', 'test_run_20240101_120000_123456') # noqa: S108
|
||||
)
|
||||
assert os.path.normpath(result) == expected_path
|
||||
mock_makedirs.assert_called_once()
|
||||
call_path = mock_makedirs.call_args[0][0]
|
||||
assert os.path.normpath(call_path) == expected_path
|
||||
assert mock_makedirs.call_args[1] == {'exist_ok': True}
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.makedirs')
|
||||
def test_create_run_directory_permission_error(
|
||||
mock_makedirs, mock_model_serving_class, mock_logger
|
||||
):
|
||||
"""Test _create_run_directory raises PermissionError."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_makedirs.side_effect = PermissionError('Permission denied')
|
||||
|
||||
with pytest.raises(PermissionError, match='Permission denied when creating directory'):
|
||||
repo._create_run_directory('/tmp/reports', 'test_run') # noqa: S108
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.shutil.copy')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_setup_run_directory_success(mock_open, mock_copy, mock_model_serving_class, mock_logger):
|
||||
"""Test _setup_run_directory creates files successfully."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
repo._setup_run_directory('/tmp/test_run', '/tmp/header.html') # noqa: S108
|
||||
|
||||
# Check that empty files were created
|
||||
assert mock_open.call_count == 3
|
||||
mock_copy.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.json.dump')
|
||||
@patch('model_manager.utils.repository.model_repository.Reports')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_generate_report_with_equation(
|
||||
mock_open,
|
||||
mock_reports_class,
|
||||
mock_json_dump,
|
||||
mock_model_serving_class,
|
||||
mock_logger,
|
||||
mock_train_result,
|
||||
):
|
||||
"""Test _generate_report creates equation JSON artifact."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Add equation to train result
|
||||
mock_train_result.equation = {
|
||||
'target_variable': 'target',
|
||||
'coefficients': {'var1': 1.5, 'var2': -0.75},
|
||||
'intercept': 10.5,
|
||||
'equation_string': 'target = 10.5 + 1.5 * var1 + -0.75 * var2',
|
||||
'latex_equation': 'target = 10.5 + 1.5 \\cdot var1 + -0.75 \\cdot var2',
|
||||
'model_type': 'Linear Regression',
|
||||
}
|
||||
|
||||
reference_data = pd.DataFrame({'var1': [1, 2], 'var2': [3, 4], 'target': [5, 6]})
|
||||
current_data = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10], 'target': [11, 12]})
|
||||
|
||||
# Mock DataFrame.to_csv to avoid file I/O
|
||||
with patch.object(pd.DataFrame, 'to_csv'):
|
||||
result = repo._generate_report(reference_data, current_data, mock_train_result)
|
||||
|
||||
# Verify equation path was set
|
||||
assert result.equation_path == os.path.join(mock_train_result.run_dir, 'model_equation.json')
|
||||
|
||||
# Verify JSON was written
|
||||
mock_json_dump.assert_called()
|
||||
call_args = mock_json_dump.call_args
|
||||
assert call_args[0][0] == mock_train_result.equation
|
||||
assert call_args[1]['indent'] == 2
|
||||
assert call_args[1]['ensure_ascii'] is False
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.Reports')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_generate_report_without_equation(
|
||||
mock_open, mock_reports_class, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_report works without equation."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# No equation
|
||||
mock_train_result.equation = None
|
||||
# Remove equation_path if it exists from fixture
|
||||
if hasattr(mock_train_result, 'equation_path'):
|
||||
del mock_train_result.equation_path
|
||||
|
||||
reference_data = pd.DataFrame({'var1': [1, 2], 'var2': [3, 4], 'target': [5, 6]})
|
||||
current_data = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10], 'target': [11, 12]})
|
||||
|
||||
# Mock DataFrame.to_csv to avoid file I/O
|
||||
with patch.object(pd.DataFrame, 'to_csv'):
|
||||
result = repo._generate_report(reference_data, current_data, mock_train_result)
|
||||
|
||||
# Verify equation section was not executed (equation_path not set)
|
||||
# Since equation is None, the equation block should not run
|
||||
assert result == mock_train_result
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_with_equation(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run logs equation artifact."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
mock_model_serving_instance = MagicMock()
|
||||
mock_model_serving_class.return_value = mock_model_serving_instance
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Set equation path
|
||||
mock_train_result.equation_path = '/tmp/test_run/model_equation.json' # noqa: S108
|
||||
|
||||
# Mock all path.exists calls to return True
|
||||
mock_exists.return_value = True
|
||||
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
# Verify equation artifact was logged
|
||||
logged_artifacts = [
|
||||
call[0][0] for call in mock_model_serving_instance.log_artifact.call_args_list
|
||||
]
|
||||
assert '/tmp/test_run/model_equation.json' in logged_artifacts # noqa: S108
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_with_static_threshold_value(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run logs static_threshold when rem_static_win is True and value is set."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
mock_model_serving_instance = MagicMock()
|
||||
mock_model_serving_class.return_value = mock_model_serving_instance
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Set static_threshold to a specific value
|
||||
mock_train_result.params.rem_static_win = True
|
||||
mock_train_result.params.static_threshold = 500
|
||||
mock_train_result.equation_path = None
|
||||
|
||||
# Mock all path.exists calls to return True
|
||||
mock_exists.return_value = True
|
||||
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
# Verify static_threshold was logged with the correct value
|
||||
log_param_calls = {
|
||||
call[0][0]: call[0][1] for call in mock_model_serving_instance.log_param.call_args_list
|
||||
}
|
||||
assert log_param_calls['static_threshold'] == 500
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_with_rem_static_win_false(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run logs static_threshold as None when rem_static_win is False."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
mock_model_serving_instance = MagicMock()
|
||||
mock_model_serving_class.return_value = mock_model_serving_instance
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Set rem_static_win to False
|
||||
mock_train_result.params.rem_static_win = False
|
||||
mock_train_result.params.static_threshold = 500 # Should be ignored
|
||||
mock_train_result.equation_path = None
|
||||
|
||||
# Mock all path.exists calls to return True
|
||||
mock_exists.return_value = True
|
||||
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
# Verify static_threshold was logged as None
|
||||
log_param_calls = {
|
||||
call[0][0]: call[0][1] for call in mock_model_serving_instance.log_param.call_args_list
|
||||
}
|
||||
assert log_param_calls['static_threshold'] is None
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_without_equation(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run works without equation."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
mock_model_serving_instance = MagicMock()
|
||||
mock_model_serving_class.return_value = mock_model_serving_instance
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# No equation
|
||||
mock_train_result.equation_path = None
|
||||
|
||||
# Mock path.exists to return True for required artifacts
|
||||
mock_exists.return_value = True
|
||||
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
# Verify only 3 artifacts were logged (report, train_data, test_data)
|
||||
assert mock_model_serving_instance.log_artifact.call_count == 3
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_missing_report(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run raises ValueError when report is missing."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Mock report doesn't exist
|
||||
def exists_side_effect(path):
|
||||
return not path.endswith('report.html')
|
||||
|
||||
mock_exists.side_effect = exists_side_effect
|
||||
|
||||
with pytest.raises(ValueError, match='Report file does not exist'):
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_none_metrics(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run raises ValueError when metrics are None."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.mse_val = None
|
||||
# Mock all paths exist so we reach the metrics check
|
||||
mock_exists.return_value = True
|
||||
|
||||
with pytest.raises(ValueError, match='One or more metrics .* are None'):
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_missing_train_data(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run raises ValueError when train data is missing."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Mock train_data doesn't exist
|
||||
def exists_side_effect(path):
|
||||
return not path.endswith('train_data.csv')
|
||||
|
||||
mock_exists.side_effect = exists_side_effect
|
||||
|
||||
with pytest.raises(ValueError, match='Training data file does not exist'):
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_save_run_missing_test_data(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _save_run raises ValueError when test data is missing."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Mock test_data doesn't exist
|
||||
def exists_side_effect(path):
|
||||
if path.endswith('test_data.csv'):
|
||||
return False
|
||||
return True
|
||||
|
||||
mock_exists.side_effect = exists_side_effect
|
||||
|
||||
with pytest.raises(ValueError, match='Test data file does not exist'):
|
||||
repo._save_run(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_init_artifacts_data_empty_x_test(mock_model_serving_class, mock_logger, mock_train_result):
|
||||
"""Test _init_artifacts_data raises ValueError when x_test is empty."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.x_test = pd.DataFrame()
|
||||
|
||||
with pytest.raises(ValueError, match='Test features .* are empty'):
|
||||
repo._init_artifacts_data(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
def test_init_artifacts_data_empty_y_test(mock_model_serving_class, mock_logger, mock_train_result):
|
||||
"""Test _init_artifacts_data raises ValueError when y_test is empty."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.y_test = pd.Series(dtype=float)
|
||||
|
||||
with pytest.raises(ValueError, match='Test target .* is empty'):
|
||||
repo._init_artifacts_data(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.makedirs')
|
||||
def test_create_run_directory_os_error(mock_makedirs, mock_model_serving_class, mock_logger):
|
||||
"""Test _create_run_directory raises OSError."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_makedirs.side_effect = OSError('Disk full')
|
||||
|
||||
with pytest.raises(OSError, match='Failed to create directory'):
|
||||
repo._create_run_directory('/tmp/reports', 'test_run') # noqa: S108
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.shutil.copy')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_setup_run_directory_file_not_found(
|
||||
mock_open, mock_copy, mock_model_serving_class, mock_logger
|
||||
):
|
||||
"""Test _setup_run_directory raises FileNotFoundError when header missing."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_copy.side_effect = FileNotFoundError('Header not found')
|
||||
|
||||
with pytest.raises(FileNotFoundError, match='Header file not found'):
|
||||
repo._setup_run_directory('/tmp/test_run', '/tmp/header.html') # noqa: S108
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.shutil.copy')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_setup_run_directory_permission_error(
|
||||
mock_open, mock_copy, mock_model_serving_class, mock_logger
|
||||
):
|
||||
"""Test _setup_run_directory raises PermissionError."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_open.side_effect = PermissionError('Permission denied')
|
||||
|
||||
with pytest.raises(PermissionError, match='Permission denied when setting up directory'):
|
||||
repo._setup_run_directory('/tmp/test_run', '/tmp/header.html') # noqa: S108
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.shutil.copy')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_setup_run_directory_os_error(mock_open, mock_copy, mock_model_serving_class, mock_logger):
|
||||
"""Test _setup_run_directory raises OSError."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_open.side_effect = OSError('Disk error')
|
||||
|
||||
with pytest.raises(OSError, match='Failed to setup run directory'):
|
||||
repo._setup_run_directory('/tmp/test_run', '/tmp/header.html') # noqa: S108
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.Reports')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_generate_report_value_error(
|
||||
mock_open, mock_reports_class, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_report raises ValueError on invalid data."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Create data that can't be converted to float64
|
||||
reference_data = pd.DataFrame({'var1': ['invalid', 'data']})
|
||||
current_data = pd.DataFrame({'var1': [1, 2]})
|
||||
|
||||
with pytest.raises(ValueError, match='Failed to convert data to float64'):
|
||||
repo._generate_report(reference_data, current_data, mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.Reports')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_generate_report_permission_error(
|
||||
mock_open, mock_reports_class, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_report raises PermissionError on write failure."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
reference_data = pd.DataFrame({'var1': [1, 2], 'var2': [3, 4], 'target': [5, 6]})
|
||||
current_data = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10], 'target': [11, 12]})
|
||||
|
||||
# Mock Reports to raise PermissionError
|
||||
mock_reports_class.side_effect = PermissionError('Permission denied')
|
||||
|
||||
with pytest.raises(PermissionError, match='Permission denied when writing report files'):
|
||||
repo._generate_report(reference_data, current_data, mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.Reports')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_generate_report_os_error(
|
||||
mock_open, mock_reports_class, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_report raises OSError on write failure."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
reference_data = pd.DataFrame({'var1': [1, 2], 'var2': [3, 4], 'target': [5, 6]})
|
||||
current_data = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10], 'target': [11, 12]})
|
||||
|
||||
# Mock Reports to raise OSError
|
||||
mock_reports_class.side_effect = OSError('Disk error')
|
||||
|
||||
with pytest.raises(OSError, match='Failed to generate report'):
|
||||
repo._generate_report(reference_data, current_data, mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.Reports')
|
||||
@patch('builtins.open', create=True)
|
||||
def test_generate_report_none_run_dir(
|
||||
mock_open, mock_reports_class, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_report raises ValueError when run_dir is None."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.run_dir = None
|
||||
|
||||
reference_data = pd.DataFrame({'var1': [1, 2], 'var2': [3, 4], 'target': [5, 6]})
|
||||
current_data = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10], 'target': [11, 12]})
|
||||
|
||||
with pytest.raises(ValueError, match='run_dir is not set'):
|
||||
repo._generate_report(reference_data, current_data, mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_generate_artifacts_no_run_name(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_artifacts raises ValueError when run_name is not set."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
mock_train_result.run_name = None
|
||||
mock_exists.return_value = True # Mock reports directory exists
|
||||
|
||||
# Mock _create_run_directory to avoid creating real directories
|
||||
with patch.object(repo, '_create_run_directory') as mock_create_dir:
|
||||
mock_create_dir.return_value = '/mock/run/dir'
|
||||
|
||||
with pytest.raises(ValueError, match='run_name must be set'):
|
||||
repo._generate_artifacts(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_generate_artifacts_reports_dir_not_found(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_artifacts raises FileNotFoundError when reports dir missing."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Mock reports directory doesn't exist
|
||||
mock_exists.return_value = False
|
||||
|
||||
with pytest.raises(FileNotFoundError, match='Reports directory does not exist'):
|
||||
repo._generate_artifacts(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
def test_generate_artifacts_header_not_found(
|
||||
mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_artifacts raises FileNotFoundError when header missing."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Mock: reports dir exists, but header doesn't
|
||||
def exists_side_effect(path):
|
||||
if path.endswith('header.html'):
|
||||
return False
|
||||
return True
|
||||
|
||||
mock_exists.side_effect = exists_side_effect
|
||||
|
||||
# Mock _create_run_directory to avoid creating real directories
|
||||
with patch.object(repo, '_create_run_directory') as mock_create_dir:
|
||||
mock_create_dir.return_value = '/mock/run/dir'
|
||||
|
||||
with pytest.raises(FileNotFoundError, match='Header file does not exist'):
|
||||
repo._generate_artifacts(mock_train_result)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
@patch('model_manager.utils.repository.model_repository.path.exists')
|
||||
@patch('model_manager.utils.repository.model_repository.path.join')
|
||||
def test_generate_artifacts_success(
|
||||
mock_join, mock_exists, mock_model_serving_class, mock_logger, mock_train_result
|
||||
):
|
||||
"""Test _generate_artifacts success case covering lines 147-148."""
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
|
||||
repo = ModelRepository(
|
||||
url='http://mlflow.test', username='user', password='pass', logger=mock_logger
|
||||
)
|
||||
|
||||
# Mock path.join to return predictable paths
|
||||
def join_side_effect(*args):
|
||||
return '/'.join(args)
|
||||
|
||||
mock_join.side_effect = join_side_effect
|
||||
mock_exists.return_value = True # Both reports dir and header.html exist
|
||||
|
||||
# Mock the internal methods to avoid actual file operations
|
||||
with (
|
||||
patch.object(repo, '_setup_run_directory') as mock_setup,
|
||||
patch.object(repo, '_generate_report') as mock_generate_report,
|
||||
patch.object(repo, '_init_artifacts_data') as mock_init_data,
|
||||
patch.object(repo, '_get_reports_directory') as mock_get_reports_dir,
|
||||
patch.object(repo, '_create_run_directory') as mock_create_run_dir,
|
||||
):
|
||||
# Setup mocks
|
||||
mock_init_data.return_value = (pd.DataFrame(), pd.DataFrame())
|
||||
mock_get_reports_dir.return_value = '/reports'
|
||||
mock_create_run_dir.return_value = '/reports/run_1'
|
||||
mock_generate_report.return_value = mock_train_result
|
||||
|
||||
# Call the method
|
||||
result = repo._generate_artifacts(mock_train_result)
|
||||
|
||||
# Verify the methods on lines 147-148 were called
|
||||
mock_setup.assert_called_once_with('/reports/run_1', '/reports/header.html')
|
||||
mock_generate_report.assert_called_once()
|
||||
|
||||
# Verify result
|
||||
assert result == mock_train_result
|
||||
@@ -1,513 +0,0 @@
|
||||
"""Unit tests for StorageRepository class."""
|
||||
|
||||
from io import BytesIO
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_logger():
|
||||
"""Create a mock logger for testing."""
|
||||
logger = Mock()
|
||||
logger.info = Mock()
|
||||
logger.error = Mock()
|
||||
logger.warning = Mock()
|
||||
return logger
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage_config():
|
||||
"""Create storage repository configuration."""
|
||||
return {
|
||||
'endpoint_url': 'http://localhost:9000',
|
||||
'access_key': 'test_access_key',
|
||||
'secret_key': 'test_secret_key',
|
||||
'region': 'us-east-1',
|
||||
'use_ssl': False,
|
||||
'max_retry_attempts': 3,
|
||||
'retry_mode': 'standard',
|
||||
'connect_timeout': 30,
|
||||
'read_timeout': 60,
|
||||
}
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_storage_repository_initialization(mock_boto3, mock_logger, storage_config):
|
||||
"""Test StorageRepository initialization with correct parameters."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Verify attributes are set correctly
|
||||
assert repo.endpoint_url == storage_config['endpoint_url']
|
||||
assert repo.access_key == storage_config['access_key']
|
||||
assert repo.secret_key == storage_config['secret_key']
|
||||
assert repo.region == storage_config['region']
|
||||
assert repo.use_ssl == storage_config['use_ssl']
|
||||
assert repo.max_retry_attempts == storage_config['max_retry_attempts']
|
||||
assert repo.retry_mode == storage_config['retry_mode']
|
||||
assert repo.connect_timeout == storage_config['connect_timeout']
|
||||
assert repo.read_timeout == storage_config['read_timeout']
|
||||
assert repo.logger == mock_logger
|
||||
|
||||
# Verify boto3 client was created
|
||||
mock_boto3.client.assert_called_once()
|
||||
call_args = mock_boto3.client.call_args
|
||||
|
||||
assert call_args[0][0] == 's3'
|
||||
assert call_args[1]['endpoint_url'] == storage_config['endpoint_url']
|
||||
assert call_args[1]['aws_access_key_id'] == storage_config['access_key']
|
||||
assert call_args[1]['aws_secret_access_key'] == storage_config['secret_key']
|
||||
assert call_args[1]['use_ssl'] == storage_config['use_ssl']
|
||||
|
||||
# Verify logger was called
|
||||
mock_logger.info.assert_called()
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_storage_repository_boto_config(mock_boto3, mock_logger, storage_config):
|
||||
"""Test that boto3 Config is created with correct retry settings."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Verify Config was passed with correct settings
|
||||
call_args = mock_boto3.client.call_args
|
||||
boto_config = call_args[1]['config']
|
||||
|
||||
assert boto_config.region_name == storage_config['region']
|
||||
assert boto_config.connect_timeout == storage_config['connect_timeout']
|
||||
assert boto_config.read_timeout == storage_config['read_timeout']
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_success(mock_boto3, mock_logger, storage_config):
|
||||
"""Test successful file fetch from MinIO."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock response
|
||||
file_content = b'test file content'
|
||||
mock_body = Mock()
|
||||
mock_body.read.return_value = file_content
|
||||
mock_body.__enter__ = Mock(return_value=mock_body)
|
||||
mock_body.__exit__ = Mock(return_value=False)
|
||||
|
||||
mock_response = {'Body': mock_body}
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.return_value = mock_response
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Fetch file
|
||||
result = repo.fetch_file('test-bucket', 'test-file.csv')
|
||||
|
||||
# Verify result
|
||||
assert isinstance(result, BytesIO)
|
||||
assert result.getvalue() == file_content
|
||||
|
||||
# Verify get_object was called correctly
|
||||
mock_s3_client.get_object.assert_called_once_with(Bucket='test-bucket', Key='test-file.csv')
|
||||
|
||||
# Verify logging
|
||||
assert mock_logger.info.call_count >= 2 # Init + fetch
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_with_large_content(mock_boto3, mock_logger, storage_config):
|
||||
"""Test fetch file with large content."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock response with large content
|
||||
large_content = b'x' * 1024 * 1024 # 1MB
|
||||
mock_body = Mock()
|
||||
mock_body.read.return_value = large_content
|
||||
mock_body.__enter__ = Mock(return_value=mock_body)
|
||||
mock_body.__exit__ = Mock(return_value=False)
|
||||
|
||||
mock_response = {'Body': mock_body}
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.return_value = mock_response
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Fetch file
|
||||
result = repo.fetch_file('test-bucket', 'large-file.bin')
|
||||
|
||||
# Verify result
|
||||
assert isinstance(result, BytesIO)
|
||||
assert len(result.getvalue()) == 1024 * 1024
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_empty_content(mock_boto3, mock_logger, storage_config):
|
||||
"""Test fetch file with empty content."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock response with empty content
|
||||
mock_body = Mock()
|
||||
mock_body.read.return_value = b''
|
||||
mock_body.__enter__ = Mock(return_value=mock_body)
|
||||
mock_body.__exit__ = Mock(return_value=False)
|
||||
|
||||
mock_response = {'Body': mock_body}
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.return_value = mock_response
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Fetch file
|
||||
result = repo.fetch_file('test-bucket', 'empty-file.txt')
|
||||
|
||||
# Verify result
|
||||
assert isinstance(result, BytesIO)
|
||||
assert result.getvalue() == b''
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_not_found(mock_boto3, mock_logger, storage_config):
|
||||
"""Test fetch file when object doesn't exist."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock to raise NoSuchKey error
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.side_effect = ClientError(
|
||||
{'Error': {'Code': 'NoSuchKey', 'Message': 'The specified key does not exist.'}},
|
||||
'GetObject',
|
||||
)
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Attempt to fetch non-existent file
|
||||
with pytest.raises(ClientError) as exc_info:
|
||||
repo.fetch_file('test-bucket', 'non-existent.csv')
|
||||
|
||||
assert exc_info.value.response['Error']['Code'] == 'NoSuchKey'
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_access_denied(mock_boto3, mock_logger, storage_config):
|
||||
"""Test fetch file when access is denied."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock to raise AccessDenied error
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.side_effect = ClientError(
|
||||
{'Error': {'Code': 'AccessDenied', 'Message': 'Access Denied'}}, 'GetObject'
|
||||
)
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Attempt to fetch file without permissions
|
||||
with pytest.raises(ClientError) as exc_info:
|
||||
repo.fetch_file('test-bucket', 'protected-file.csv')
|
||||
|
||||
assert exc_info.value.response['Error']['Code'] == 'AccessDenied'
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_network_error(mock_boto3, mock_logger, storage_config):
|
||||
"""Test fetch file when network error occurs."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock to raise network error
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.side_effect = ConnectionError('Network unreachable')
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Attempt to fetch file with network error
|
||||
with pytest.raises(ConnectionError):
|
||||
repo.fetch_file('test-bucket', 'test-file.csv')
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_delete_file_success(mock_boto3, mock_logger, storage_config):
|
||||
"""Test successful file deletion from MinIO."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.delete_object.return_value = {}
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Delete file
|
||||
repo.delete_file('test-bucket', 'test-file.csv')
|
||||
|
||||
# Verify delete_object was called correctly
|
||||
mock_s3_client.delete_object.assert_called_once_with(Bucket='test-bucket', Key='test-file.csv')
|
||||
|
||||
# Verify logging
|
||||
assert mock_logger.info.call_count >= 2 # Init + delete
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_delete_file_non_existent(mock_boto3, mock_logger, storage_config):
|
||||
"""Test delete file that doesn't exist (should succeed silently in S3)."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# S3/MinIO delete is idempotent - deleting non-existent file succeeds
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.delete_object.return_value = {}
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Delete non-existent file (should succeed)
|
||||
repo.delete_file('test-bucket', 'non-existent.csv')
|
||||
|
||||
mock_s3_client.delete_object.assert_called_once()
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_delete_file_access_denied(mock_boto3, mock_logger, storage_config):
|
||||
"""Test delete file when access is denied."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock to raise AccessDenied error
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.delete_object.side_effect = ClientError(
|
||||
{'Error': {'Code': 'AccessDenied', 'Message': 'Access Denied'}}, 'DeleteObject'
|
||||
)
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Attempt to delete file without permissions
|
||||
with pytest.raises(ClientError) as exc_info:
|
||||
repo.delete_file('test-bucket', 'protected-file.csv')
|
||||
|
||||
assert exc_info.value.response['Error']['Code'] == 'AccessDenied'
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_delete_file_network_error(mock_boto3, mock_logger, storage_config):
|
||||
"""Test delete file when network error occurs."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
# Setup mock to raise network error
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.delete_object.side_effect = ConnectionError('Network unreachable')
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Attempt to delete file with network error
|
||||
with pytest.raises(ConnectionError):
|
||||
repo.delete_file('test-bucket', 'test-file.csv')
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_storage_repository_with_ssl(mock_boto3, mock_logger, storage_config):
|
||||
"""Test StorageRepository initialization with SSL enabled."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
storage_config['use_ssl'] = True
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
assert repo.use_ssl is True
|
||||
|
||||
# Verify boto3 client was created with use_ssl=True
|
||||
call_args = mock_boto3.client.call_args
|
||||
assert call_args[1]['use_ssl'] is True
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_storage_repository_custom_timeouts(mock_boto3, mock_logger, storage_config):
|
||||
"""Test StorageRepository with custom timeout values."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
storage_config['connect_timeout'] = 10
|
||||
storage_config['read_timeout'] = 120
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
assert repo.connect_timeout == 10
|
||||
assert repo.read_timeout == 120
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_storage_repository_custom_retry_mode(mock_boto3, mock_logger, storage_config):
|
||||
"""Test StorageRepository with different retry modes."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
for retry_mode in ['standard', 'legacy', 'adaptive']:
|
||||
storage_config['retry_mode'] = retry_mode
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
assert repo.retry_mode == retry_mode
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_storage_repository_custom_max_retries(mock_boto3, mock_logger, storage_config):
|
||||
"""Test StorageRepository with different max retry attempts."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
storage_config['max_retry_attempts'] = 5
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
assert repo.max_retry_attempts == 5
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_with_special_characters(mock_boto3, mock_logger, storage_config):
|
||||
"""Test fetch file with special characters in name."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
file_content = b'test content'
|
||||
mock_body = Mock()
|
||||
mock_body.read.return_value = file_content
|
||||
mock_body.__enter__ = Mock(return_value=mock_body)
|
||||
mock_body.__exit__ = Mock(return_value=False)
|
||||
|
||||
mock_response = {'Body': mock_body}
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.return_value = mock_response
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Fetch file with special characters
|
||||
special_filename = 'test file (2023-01-01) #1.csv'
|
||||
result = repo.fetch_file('test-bucket', special_filename)
|
||||
|
||||
assert isinstance(result, BytesIO)
|
||||
mock_s3_client.get_object.assert_called_once_with(Bucket='test-bucket', Key=special_filename)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_delete_file_with_path_separators(mock_boto3, mock_logger, storage_config):
|
||||
"""Test delete file with path separators in object key."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.delete_object.return_value = {}
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
# Delete file with path separators
|
||||
file_path = 'data/2023/01/test-file.csv'
|
||||
repo.delete_file('test-bucket', file_path)
|
||||
|
||||
mock_s3_client.delete_object.assert_called_once_with(Bucket='test-bucket', Key=file_path)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_storage_repository_different_regions(mock_boto3, mock_logger, storage_config):
|
||||
"""Test StorageRepository with different AWS regions."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
regions = ['us-west-1', 'eu-central-1', 'ap-southeast-1']
|
||||
|
||||
for region in regions:
|
||||
storage_config['region'] = region
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
assert repo.region == region
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_fetch_file_logs_file_size(mock_boto3, mock_logger, storage_config):
|
||||
"""Test that fetch_file logs the file size."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
file_content = b'x' * 12345
|
||||
mock_body = Mock()
|
||||
mock_body.read.return_value = file_content
|
||||
mock_body.__enter__ = Mock(return_value=mock_body)
|
||||
mock_body.__exit__ = Mock(return_value=False)
|
||||
|
||||
mock_response = {'Body': mock_body}
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_s3_client.get_object.return_value = mock_response
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
|
||||
repo.fetch_file('test-bucket', 'test-file.csv')
|
||||
|
||||
# Verify logging includes file size
|
||||
log_calls = [str(call) for call in mock_logger.info.call_args_list]
|
||||
assert any('12345 bytes' in str(call) for call in log_calls)
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_close_method(mock_boto3, mock_logger, storage_config):
|
||||
"""Test that the close method calls the underlying client's close method."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
repo.close()
|
||||
|
||||
mock_s3_client.close.assert_called_once()
|
||||
mock_logger.info.assert_called_with('MinIO client closed')
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.storage_repository.boto3')
|
||||
def test_list_bucket_objects_with_pagination(mock_boto3, mock_logger, storage_config):
|
||||
"""Test list_bucket_objects with a paginated response."""
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
mock_s3_client = Mock()
|
||||
mock_paginator = Mock()
|
||||
page1 = {
|
||||
'Contents': [
|
||||
{'Key': 'file1.txt'},
|
||||
{'Key': 'file2.txt'},
|
||||
]
|
||||
}
|
||||
page2 = {
|
||||
'Contents': [
|
||||
{'Key': 'file3.txt'},
|
||||
]
|
||||
}
|
||||
page3 = {}
|
||||
|
||||
mock_paginator.paginate.return_value = [page1, page2, page3]
|
||||
mock_s3_client.get_paginator.return_value = mock_paginator
|
||||
mock_boto3.client.return_value = mock_s3_client
|
||||
|
||||
repo = StorageRepository(logger=mock_logger, **storage_config)
|
||||
objects = repo.list_bucket_objects('test-bucket', max_keys=2)
|
||||
|
||||
assert objects == ['file1.txt', 'file2.txt', 'file3.txt']
|
||||
assert len(objects) == 3
|
||||
mock_s3_client.get_paginator.assert_called_once_with('list_objects_v2')
|
||||
mock_paginator.paginate.assert_called_once_with(Bucket='test-bucket', MaxKeys=2)
|
||||
mock_logger.info.assert_any_call('Listed 3 objects from bucket test-bucket')
|
||||
File diff suppressed because it is too large
Load Diff
@@ -20,6 +20,10 @@ def mock_env_vars():
|
||||
'PROJECT_NAME': 'test-project',
|
||||
'TRAIN_TASK_QUEUE': 'train_model-local_queue',
|
||||
'CLEANUP_TASK_QUEUE': 'cleanup-local_queue',
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
'STORE_BASE_URL': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'STORE_OWNER': 'sientia',
|
||||
'STORE_REPO': 'model-library-store',
|
||||
}
|
||||
|
||||
with patch.dict(os.environ, env_vars, clear=False):
|
||||
@@ -176,6 +180,8 @@ def test_start_prometheus_server_failure(
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@@ -191,6 +197,8 @@ async def test_main_successful_startup(
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
@@ -220,6 +228,23 @@ async def test_main_successful_startup(
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
}
|
||||
|
||||
mock_runtime = Mock()
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
@@ -262,6 +287,8 @@ async def test_main_successful_startup(
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@@ -277,6 +304,8 @@ async def test_main_handles_exception(
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
@@ -320,6 +349,23 @@ async def test_main_handles_exception(
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
}
|
||||
|
||||
# Run main() and expect SystemExit
|
||||
with pytest.raises(SystemExit) as exc_info:
|
||||
await main()
|
||||
@@ -343,6 +389,8 @@ async def test_main_handles_exception(
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@@ -358,6 +406,8 @@ async def test_main_temporal_client_configuration(
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
@@ -377,6 +427,10 @@ async def test_main_temporal_client_configuration(
|
||||
'TEMPORAL_HOST': 'temporal.example.com:7233',
|
||||
'TEMPORAL_NAMESPACE': 'production',
|
||||
'TEMPORAL_USE_TLS': 'true',
|
||||
'RUNTIME': 'model-manager-worker',
|
||||
'STORE_BASE_URL': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'STORE_OWNER': 'sientia',
|
||||
'STORE_REPO': 'model-library-store',
|
||||
},
|
||||
):
|
||||
# Setup mocks
|
||||
@@ -390,6 +444,23 @@ async def test_main_temporal_client_configuration(
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
}
|
||||
|
||||
mock_notification_handler = Mock()
|
||||
mock_notification_handler.shutdown = Mock()
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
@@ -431,6 +502,8 @@ async def test_main_temporal_client_configuration(
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@@ -446,6 +519,8 @@ async def test_main_worker_configuration(
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
@@ -500,6 +575,23 @@ async def test_main_worker_configuration(
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
}
|
||||
|
||||
# Run main()
|
||||
with pytest.raises(SystemExit):
|
||||
await main()
|
||||
@@ -539,6 +631,8 @@ async def test_main_worker_configuration(
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@@ -554,6 +648,8 @@ async def test_main_schedule_creation_failure_does_not_stop_worker(
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
@@ -603,6 +699,106 @@ async def test_main_schedule_creation_failure_does_not_stop_worker(
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.worker.worker.Worker')
|
||||
@patch('model_manager.worker.worker.client.Client')
|
||||
@patch('model_manager.worker.worker.Runtime')
|
||||
@patch('model_manager.worker.worker.Activities')
|
||||
@patch('model_manager.worker.worker.NotificationHandler')
|
||||
@patch('model_manager.worker.worker.PluginStore')
|
||||
@patch('model_manager.worker.worker.build_plugin_store_config')
|
||||
@patch('model_manager.worker.worker.build_mongodb_config')
|
||||
@patch('model_manager.worker.worker.build_postgres_config')
|
||||
@patch('model_manager.worker.worker.build_mlflow_config')
|
||||
@patch('model_manager.worker.worker.build_minio_config')
|
||||
@patch('model_manager.worker.worker.get_logger')
|
||||
@patch('model_manager.worker.worker.start_prometheus_server')
|
||||
@patch('model_manager.worker.worker.metrics')
|
||||
async def test_main_missing_runtime_fails_fast(
|
||||
mock_metrics,
|
||||
mock_start_prometheus,
|
||||
mock_get_logger,
|
||||
mock_build_minio,
|
||||
mock_build_mlflow,
|
||||
mock_build_postgres,
|
||||
mock_build_mongodb,
|
||||
mock_build_plugin_store_config,
|
||||
mock_plugin_store_class,
|
||||
mock_notification_handler_class,
|
||||
mock_activities_class,
|
||||
mock_runtime_class,
|
||||
mock_client_class,
|
||||
mock_worker_class,
|
||||
mock_logger,
|
||||
):
|
||||
"""Test that main() fails fast when RUNTIME is missing."""
|
||||
from model_manager.worker.worker import main
|
||||
|
||||
mock_get_logger.return_value = mock_logger
|
||||
mock_build_mongodb.return_value = {
|
||||
'connection_string': 'mongodb://test',
|
||||
'database_name': 'test_db',
|
||||
'uri': 'localhost:27018',
|
||||
}
|
||||
mock_build_postgres.return_value = {}
|
||||
mock_build_mlflow.return_value = {}
|
||||
mock_build_minio.return_value = {}
|
||||
|
||||
mock_notification_handler = Mock()
|
||||
mock_notification_handler.shutdown = Mock()
|
||||
mock_notification_handler_class.return_value = mock_notification_handler
|
||||
|
||||
mock_activities = AsyncMock()
|
||||
mock_activities.shutdown = AsyncMock()
|
||||
mock_activities_class.return_value = mock_activities
|
||||
|
||||
mock_app_up = Mock()
|
||||
mock_metrics.APP_UP.labels.return_value = mock_app_up
|
||||
|
||||
mock_plugin_store_instance = AsyncMock()
|
||||
mock_plugin_store_instance.install_runtime = AsyncMock(
|
||||
return_value={'runtime': 'model-manager-worker', 'installed': []},
|
||||
)
|
||||
mock_plugin_store_class.return_value = mock_plugin_store_instance
|
||||
mock_build_plugin_store_config.return_value = {
|
||||
'base_url': 'http://sientia-plugin-store.svc.cluster.local',
|
||||
'owner': 'sientia',
|
||||
'repo': 'model-library-store',
|
||||
'branch': 'main',
|
||||
'username': 'gitea-user',
|
||||
'password': 'gitea-password',
|
||||
'pypi_index_url': 'http://library-distribution-server.library.svc.cluster.local:5000',
|
||||
'pypi_username': None,
|
||||
'pypi_password': None,
|
||||
}
|
||||
|
||||
# Ensure RUNTIME is not defined
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
with pytest.raises(SystemExit) as exc_info:
|
||||
await main()
|
||||
|
||||
assert exc_info.value.code == 1
|
||||
mock_logger.custom_critical.assert_called_once()
|
||||
mock_app_up.set.assert_called_with(0)
|
||||
|
||||
# Run main() - should not fail despite schedule creation error
|
||||
with pytest.raises(SystemExit):
|
||||
await main()
|
||||
|
||||
Reference in New Issue
Block a user