diff --git a/tests/utils/repository/test_model_repository.py b/tests/utils/repository/test_model_repository.py new file mode 100644 index 0000000..ba773d6 --- /dev/null +++ b/tests/utils/repository/test_model_repository.py @@ -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)