"""Unit tests for ModelServing class.""" from unittest.mock import MagicMock, patch import pandas as pd from pytest import raises from model_manager.sientia.exceptions import SientiaMlException from model_manager.sientia.model_serving import ModelServing @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch.dict('os.environ', {}, clear=True) def test_init_with_all_credentials(mock_set_tracking_uri): """Test initialization with tracking URI, username, and password.""" tracking_uri = 'http://mlflow.example.com' username = 'test_user' password = 'test_pass' ModelServing(tracking_uri=tracking_uri, username=username, password=password) mock_set_tracking_uri.assert_called_once_with(tracking_uri) import os assert os.environ['MLFLOW_TRACKING_USERNAME'] == username assert os.environ['MLFLOW_TRACKING_PASSWORD'] == password @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch.dict('os.environ', {}, clear=True) def test_init_without_credentials(mock_set_tracking_uri): """Test initialization without username and password.""" tracking_uri = 'http://mlflow.example.com' ModelServing(tracking_uri=tracking_uri) mock_set_tracking_uri.assert_called_once_with(tracking_uri) import os assert 'MLFLOW_TRACKING_USERNAME' not in os.environ assert 'MLFLOW_TRACKING_PASSWORD' not in os.environ @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch.dict('os.environ', {}, clear=True) def test_init_with_only_username(mock_set_tracking_uri): """Test initialization with only username (no password).""" tracking_uri = 'http://mlflow.example.com' username = 'test_user' ModelServing(tracking_uri=tracking_uri, username=username) mock_set_tracking_uri.assert_called_once_with(tracking_uri) import os assert os.environ['MLFLOW_TRACKING_USERNAME'] == username assert 'MLFLOW_TRACKING_PASSWORD' not in os.environ @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch.dict('os.environ', {}, clear=True) def test_init_with_only_password(mock_set_tracking_uri): """Test initialization with only password (no username).""" tracking_uri = 'http://mlflow.example.com' password = 'test_pass' ModelServing(tracking_uri=tracking_uri, password=password) mock_set_tracking_uri.assert_called_once_with(tracking_uri) import os assert 'MLFLOW_TRACKING_USERNAME' not in os.environ assert os.environ['MLFLOW_TRACKING_PASSWORD'] == password @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.search_runs') def test_search_runs_by_name_success(mock_search_runs, mock_set_tracking_uri): """Test successful search_runs_by_name.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') expected_df = pd.DataFrame({'run_id': ['123', '456'], 'status': ['FINISHED', 'RUNNING']}) mock_search_runs.return_value = expected_df experiment_names = ['experiment1', 'experiment2'] result = model_serving.search_runs_by_name(experiment_names) mock_search_runs.assert_called_once_with(experiment_names=experiment_names, order_by=None) pd.testing.assert_frame_equal(result, expected_df) @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.search_runs') def test_search_runs_by_name_with_order_by(mock_search_runs, mock_set_tracking_uri): """Test search_runs_by_name with order_by parameter.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') expected_df = pd.DataFrame({'run_id': ['123'], 'status': ['FINISHED']}) mock_search_runs.return_value = expected_df experiment_names = ['experiment1'] order_by = ['start_time DESC'] result = model_serving.search_runs_by_name(experiment_names, order_by=order_by) mock_search_runs.assert_called_once_with(experiment_names=experiment_names, order_by=order_by) pd.testing.assert_frame_equal(result, expected_df) @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.search_runs') @patch('model_manager.sientia.model_serving.logging.error') def test_search_runs_by_name_raises_exception( mock_logging_error, mock_search_runs, mock_set_tracking_uri ): """Test search_runs_by_name raises TypeError due to bug in line 80 of model_serving.py.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') exception = SientiaMlException(message='Search failed') mock_search_runs.side_effect = exception # The code has a bug on line 80: "raise SientiaMlException from e" # This raises TypeError because SientiaMlException requires 'message' argument with raises(TypeError, match="missing 1 required positional argument: 'message'"): model_serving.search_runs_by_name(['experiment1']) mock_logging_error.assert_called_once() @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.set_experiment') def test_set_experiment(mock_set_experiment, mock_set_tracking_uri): """Test set_experiment method.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') experiment_identifier = 'my_experiment' model_serving.set_experiment(experiment_identifier) mock_set_experiment.assert_called_once_with(experiment_identifier) @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.sklearn.log_model') def test_log_model(mock_log_model, mock_set_tracking_uri): """Test log_model method.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') sk_model = MagicMock() artifact_path = 'model' model_serving.log_model(sk_model, artifact_path) mock_log_model.assert_called_once() call_args = mock_log_model.call_args assert call_args[0][0] == sk_model assert call_args[0][1] == artifact_path assert 'extra_pip_requirements' in call_args[1] @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.sklearn.log_model') def test_log_model_with_kwargs(mock_log_model, mock_set_tracking_uri): """Test log_model method with additional kwargs.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') sk_model = MagicMock() artifact_path = 'model' registered_model_name = 'my_model' model_serving.log_model(sk_model, artifact_path, registered_model_name=registered_model_name) mock_log_model.assert_called_once() call_args = mock_log_model.call_args assert call_args[1]['registered_model_name'] == registered_model_name @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.log_param') def test_log_param(mock_log_param, mock_set_tracking_uri): """Test log_param method.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') key = 'learning_rate' value = 0.01 model_serving.log_param(key, value) mock_log_param.assert_called_once_with(key, value) @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.log_metric') def test_log_metric(mock_log_metric, mock_set_tracking_uri): """Test log_metric method.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') key = 'accuracy' value = 0.95 model_serving.log_metric(key, value) mock_log_metric.assert_called_once_with(key, value) @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.log_artifact') def test_log_artifact_with_all_params(mock_log_artifact, mock_set_tracking_uri): """Test log_artifact method with all parameters.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') local_path = '/path/to/artifact.txt' artifact_path = 'artifacts' run_id = 'run_123' model_serving.log_artifact(local_path, artifact_path, run_id) mock_log_artifact.assert_called_once_with( local_path=local_path, artifact_path=artifact_path, run_id=run_id ) @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.log_artifact') def test_log_artifact_with_minimal_params(mock_log_artifact, mock_set_tracking_uri): """Test log_artifact method with only required parameter.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') local_path = '/path/to/artifact.txt' model_serving.log_artifact(local_path) mock_log_artifact.assert_called_once_with( local_path=local_path, artifact_path=None, run_id=None ) @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.start_run') @patch('model_manager.sientia.model_serving.mlflow.end_run') def test_save_experiment_context_manager(mock_end_run, mock_start_run, mock_set_tracking_uri): """Test save_experiment context manager.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') mock_run = MagicMock() mock_start_run.return_value = mock_run with model_serving.save_experiment(run_name='test_run') as run: assert run == mock_run mock_start_run.assert_called_once_with( run_id=None, experiment_id=None, run_name='test_run', nested=False, tags=None, description=None, log_system_metrics=None, ) mock_end_run.assert_called_once() @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.start_run') @patch('model_manager.sientia.model_serving.mlflow.end_run') def test_save_experiment_with_all_params(mock_end_run, mock_start_run, mock_set_tracking_uri): """Test save_experiment with all parameters.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') mock_run = MagicMock() mock_start_run.return_value = mock_run run_id = 'run_123' experiment_id = 'exp_456' run_name = 'test_run' nested = True tags = {'key': 'value'} description = 'Test description' log_system_metrics = True with model_serving.save_experiment( run_id=run_id, experiment_id=experiment_id, run_name=run_name, nested=nested, tags=tags, description=description, log_system_metrics=log_system_metrics, ) as run: assert run == mock_run mock_start_run.assert_called_once_with( run_id=run_id, experiment_id=experiment_id, run_name=run_name, nested=nested, tags=tags, description=description, log_system_metrics=log_system_metrics, ) mock_end_run.assert_called_once() @patch('model_manager.sientia.model_serving.mlflow.set_tracking_uri') @patch('model_manager.sientia.model_serving.mlflow.start_run') @patch('model_manager.sientia.model_serving.mlflow.end_run') def test_save_experiment_ensures_end_run_on_exception( mock_end_run, mock_start_run, mock_set_tracking_uri ): """Test save_experiment ensures end_run is called even when exception occurs.""" model_serving = ModelServing(tracking_uri='http://mlflow.example.com') mock_run = MagicMock() mock_start_run.return_value = mock_run with raises(ValueError): with model_serving.save_experiment(run_name='test_run'): raise ValueError('Test exception') mock_start_run.assert_called_once() mock_end_run.assert_called_once()