From 3c09e59f6aba900437025ff8ebd7d1bdc4ae8f29 Mon Sep 17 00:00:00 2001 From: Bruno Domingues Date: Mon, 20 Oct 2025 20:24:50 -0300 Subject: [PATCH] SIENTIAPDE-1255: Add unit tests for the ModelServing class. --- tests/sientia/test_model_serving.py | 314 ++++++++++++++++++++++++++++ 1 file changed, 314 insertions(+) create mode 100644 tests/sientia/test_model_serving.py diff --git a/tests/sientia/test_model_serving.py b/tests/sientia/test_model_serving.py new file mode 100644 index 0000000..9e6117c --- /dev/null +++ b/tests/sientia/test_model_serving.py @@ -0,0 +1,314 @@ +"""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' + logger = MagicMock() + + ModelServing(tracking_uri=tracking_uri, username=username, password=password, logger=logger) + + 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()