SIENTIAPDE-1255: Add unit tests for the ModelServing class.

This commit is contained in:
Bruno Domingues
2025-10-20 20:24:50 -03:00
parent b1fe62f68b
commit 3c09e59f6a

View File

@@ -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()