SIENTIAPDE-1241: Add unit tests for ModelRepository with 100% coverage.
This commit is contained in:
201
tests/utils/repository/test_model_repository.py
Normal file
201
tests/utils/repository/test_model_repository.py
Normal file
@@ -0,0 +1,201 @@
|
|||||||
|
"""Unit tests for ModelRepository with 100% coverage."""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
@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.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
|
||||||
|
)
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
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('model_manager/reports')
|
||||||
|
assert os.path.isabs(result)
|
||||||
Reference in New Issue
Block a user