312 lines
11 KiB
Python
312 lines
11 KiB
Python
"""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 properly propagates SientiaMlException."""
|
|
model_serving = ModelServing(tracking_uri='http://mlflow.example.com')
|
|
|
|
exception = SientiaMlException(message='Search failed')
|
|
mock_search_runs.side_effect = exception
|
|
|
|
with raises(SientiaMlException, match='Search failed'):
|
|
model_serving.search_runs_by_name(['experiment1'])
|
|
|
|
mock_logging_error.assert_called_once_with(exception)
|
|
|
|
|
|
@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()
|