SIENTIAPDE-1255: Add unit tests for the ModelServing class.
This commit is contained in:
314
tests/sientia/test_model_serving.py
Normal file
314
tests/sientia/test_model_serving.py
Normal 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()
|
||||
Reference in New Issue
Block a user