SIENTIAPDE-1250: Refactor: Remove 'laborious' directory from test structure and update README.md accordingly.

This commit is contained in:
Bruno Domingues
2025-10-07 10:02:38 -03:00
parent 851f34c80e
commit 2628dc92d2
23 changed files with 11 additions and 13 deletions

0
tests/utils/__init__.py Normal file
View File

View File

View File

@@ -0,0 +1,37 @@
from pandas import DataFrame
from model_manager.utils.filters.conditional_filters import (
filter_empty_data,
filter_specific_variables_null_values,
)
def test_filter_specific_variables_null_values():
assert (
filter_specific_variables_null_values(
DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, 2]}),
config={'variables': ['variable2']},
)
is False
)
def test_filter_specific_variables_null_values_with_null_values():
assert (
filter_specific_variables_null_values(
DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, None]}),
config={'variables': ['variable2']},
)
is True
)
def test_filter_empty_data():
assert filter_empty_data(DataFrame(), {}) is True
def test_filter_empty_data_with_data():
assert (
filter_empty_data(DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, 2]}), {})
is False
)

View File

@@ -0,0 +1,23 @@
from pandas import DataFrame
from model_manager.utils.filters.mlflow_filters import api_error_filter, nan_values_filter
def test_api_error_filter_invalid_response():
assert api_error_filter(None, {})
def test_api_error_filter_valid_response_fail():
assert api_error_filter({'success': False}, {})
def test_api_error_filter_valid_response_success():
assert not api_error_filter({'success': True}, {})
def test_nan_values_filter_all_nan_values():
assert nan_values_filter(DataFrame({'variable': [None, None]}), {})
def test_nan_values_filter_no_nan_values():
assert not nan_values_filter(DataFrame({'variable': [1, 2]}), {})

View File

View File

@@ -0,0 +1,79 @@
"""Unit tests for ExperimentStatus enum."""
from model_manager.utils.models.experiment_status import ExperimentStatus
def test_experiment_status_values():
"""Test that all expected status values exist."""
assert ExperimentStatus.MAGE_WAITING_PROC == 'MAGE_WAITING_PROC'
assert ExperimentStatus.TRAINING_SUCCESS == 'TRAINING_SUCCESS'
assert ExperimentStatus.TRAINING_ERROR == 'TRAINING_ERROR'
assert ExperimentStatus.MLFLOW_SENT == 'MLFLOW_SENT'
assert ExperimentStatus.MLFLOW_SEND_ERROR == 'MLFLOW_SEND_ERROR'
assert ExperimentStatus.FILE_DELETED == 'FILE_DELETED'
assert ExperimentStatus.FILE_DELETE_ERROR == 'FILE_DELETE_ERROR'
def test_experiment_status_count():
"""Test that enum has exactly 7 status values."""
assert len(ExperimentStatus) == 7
def test_experiment_status_is_string():
"""Test that enum values are strings."""
for status in ExperimentStatus:
assert isinstance(status.value, str)
assert isinstance(status, str)
def test_experiment_status_membership():
"""Test membership checks for status values."""
assert 'MAGE_WAITING_PROC' in [s.value for s in ExperimentStatus]
assert 'TRAINING_SUCCESS' in [s.value for s in ExperimentStatus]
assert 'TRAINING_ERROR' in [s.value for s in ExperimentStatus]
assert 'MLFLOW_SENT' in [s.value for s in ExperimentStatus]
assert 'MLFLOW_SEND_ERROR' in [s.value for s in ExperimentStatus]
assert 'FILE_DELETED' in [s.value for s in ExperimentStatus]
assert 'FILE_DELETE_ERROR' in [s.value for s in ExperimentStatus]
def test_experiment_status_iteration():
"""Test that enum can be iterated."""
statuses = list(ExperimentStatus)
assert len(statuses) == 7
assert ExperimentStatus.MAGE_WAITING_PROC in statuses
assert ExperimentStatus.TRAINING_SUCCESS in statuses
assert ExperimentStatus.TRAINING_ERROR in statuses
assert ExperimentStatus.MLFLOW_SENT in statuses
assert ExperimentStatus.MLFLOW_SEND_ERROR in statuses
assert ExperimentStatus.FILE_DELETED in statuses
assert ExperimentStatus.FILE_DELETE_ERROR in statuses
def test_experiment_status_comparison():
"""Test that enum values can be compared with strings."""
assert ExperimentStatus.MAGE_WAITING_PROC == 'MAGE_WAITING_PROC'
assert ExperimentStatus.TRAINING_SUCCESS == 'TRAINING_SUCCESS'
assert ExperimentStatus.TRAINING_ERROR != 'TRAINING_SUCCESS'
def test_experiment_status_access_by_name():
"""Test accessing enum members by name."""
assert ExperimentStatus['MAGE_WAITING_PROC'] == ExperimentStatus.MAGE_WAITING_PROC
assert ExperimentStatus['TRAINING_SUCCESS'] == ExperimentStatus.TRAINING_SUCCESS
assert ExperimentStatus['TRAINING_ERROR'] == ExperimentStatus.TRAINING_ERROR
assert ExperimentStatus['MLFLOW_SENT'] == ExperimentStatus.MLFLOW_SENT
assert ExperimentStatus['MLFLOW_SEND_ERROR'] == ExperimentStatus.MLFLOW_SEND_ERROR
assert ExperimentStatus['FILE_DELETED'] == ExperimentStatus.FILE_DELETED
assert ExperimentStatus['FILE_DELETE_ERROR'] == ExperimentStatus.FILE_DELETE_ERROR
def test_experiment_status_access_by_value():
"""Test accessing enum members by value."""
assert ExperimentStatus('MAGE_WAITING_PROC') == ExperimentStatus.MAGE_WAITING_PROC
assert ExperimentStatus('TRAINING_SUCCESS') == ExperimentStatus.TRAINING_SUCCESS
assert ExperimentStatus('TRAINING_ERROR') == ExperimentStatus.TRAINING_ERROR
assert ExperimentStatus('MLFLOW_SENT') == ExperimentStatus.MLFLOW_SENT
assert ExperimentStatus('MLFLOW_SEND_ERROR') == ExperimentStatus.MLFLOW_SEND_ERROR
assert ExperimentStatus('FILE_DELETED') == ExperimentStatus.FILE_DELETED
assert ExperimentStatus('FILE_DELETE_ERROR') == ExperimentStatus.FILE_DELETE_ERROR

View File

@@ -0,0 +1,51 @@
"""Unit tests for models __init__.py module."""
from model_manager.utils.models import (
ExperimentStatus,
TrainModelParams,
TrainModelResult,
)
def test_experiment_status_import():
"""Test that ExperimentStatus can be imported from models package."""
assert ExperimentStatus is not None
assert hasattr(ExperimentStatus, 'MAGE_WAITING_PROC')
assert hasattr(ExperimentStatus, 'TRAINING_SUCCESS')
def test_train_model_params_import():
"""Test that TrainModelParams can be imported from models package."""
assert TrainModelParams is not None
assert callable(TrainModelParams)
def test_train_model_result_import():
"""Test that TrainModelResult can be imported from models package."""
assert TrainModelResult is not None
# Dataclasses have __dataclass_fields__
assert hasattr(TrainModelResult, '__dataclass_fields__')
def test_all_exports():
"""Test that __all__ contains all expected exports."""
from model_manager.utils.models import __all__
assert 'ExperimentStatus' in __all__
assert 'TrainModelParams' in __all__
assert 'TrainModelResult' in __all__
assert len(__all__) == 3
def test_no_extra_exports():
"""Test that only expected items are exported."""
import model_manager.utils.models as models_module
# Get all public attributes (not starting with _)
public_attrs = [attr for attr in dir(models_module) if not attr.startswith('_')]
# Should only have the 3 main classes
expected_public = {'ExperimentStatus', 'TrainModelParams', 'TrainModelResult'}
# Check that our expected classes are present
assert expected_public.issubset(set(public_attrs))

View File

@@ -0,0 +1,304 @@
"""Unit tests for TrainModelParams class."""
import pytest
from model_manager.utils.models.train_model_params import TrainModelParams
@pytest.fixture
def valid_params_dict():
"""Create valid parameters dictionary for testing."""
return {
'variable_columns': ['var1', 'var2', 'var3'],
'lag_train': 5,
'lag_val': 3,
'target_variable': 'target',
'rem_static_win': True,
'low_lim': {'var1': 0.0, 'var2': 0.0, 'var3': 0.0},
'upp_lim': {'var1': 100.0, 'var2': 100.0, 'var3': 100.0},
'window': 10,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'test-bucket',
'file_name': 'test-file.csv',
'line_separator': '\n',
'decimal_separator': '.',
'train_size': 80,
'shuffle': True,
'experiment_run_id': 123,
'experiment_name': 'test-experiment',
'experiment_description': 'Test experiment description',
'removed_intervals': [],
}
def test_train_model_params_creation_with_valid_params(valid_params_dict):
"""Test creating TrainModelParams with all valid parameters."""
params = TrainModelParams(**valid_params_dict)
assert params.variable_columns == ['var1', 'var2', 'var3']
assert params.lag_train == 5
assert params.lag_val == 3
assert params.target_variable == 'target'
assert params.rem_static_win is True
assert params.low_lim == {'var1': 0.0, 'var2': 0.0, 'var3': 0.0}
assert params.upp_lim == {'var1': 100.0, 'var2': 100.0, 'var3': 100.0}
assert params.window == 10
assert params.use_scaler is True
assert params.include_ar is False
assert params.bucket_name == 'test-bucket'
assert params.file_name == 'test-file.csv'
assert params.line_separator == '\n'
assert params.decimal_separator == '.'
assert params.train_size == 80
assert params.shuffle is True
assert params.experiment_run_id == 123
assert params.experiment_name == 'test-experiment'
assert params.experiment_description == 'Test experiment description'
assert params.removed_intervals == []
def test_train_model_params_from_dict_creation(valid_params_dict):
"""Test creating TrainModelParams using from_dict method."""
params = TrainModelParams.from_dict(valid_params_dict)
assert params.variable_columns == ['var1', 'var2', 'var3']
assert params.lag_train == 5
assert params.experiment_run_id == 123
def test_train_model_params_variable_columns_none_raises_error(valid_params_dict):
"""Test that None variable_columns raises ValueError."""
valid_params_dict['variable_columns'] = None
with pytest.raises(ValueError, match='variable_columns is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_variable_columns_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for variable_columns raises TypeError."""
valid_params_dict['variable_columns'] = 'not a list'
with pytest.raises(TypeError, match='variable_columns must be of type list'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_lag_train_none_raises_error(valid_params_dict):
"""Test that None lag_train raises ValueError."""
valid_params_dict['lag_train'] = None
with pytest.raises(ValueError, match='lag_train is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_lag_train_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for lag_train raises TypeError."""
valid_params_dict['lag_train'] = '5'
with pytest.raises(TypeError, match='lag_train must be of type int'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_target_variable_none_raises_error(valid_params_dict):
"""Test that None target_variable raises ValueError."""
valid_params_dict['target_variable'] = None
with pytest.raises(ValueError, match='target_variable is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_target_variable_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for target_variable raises TypeError."""
valid_params_dict['target_variable'] = 123
with pytest.raises(TypeError, match='target_variable must be of type str'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_boolean_fields(valid_params_dict):
"""Test boolean fields validation."""
# Test rem_static_win
valid_params_dict['rem_static_win'] = None
with pytest.raises(ValueError, match='rem_static_win is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
valid_params_dict['rem_static_win'] = 'true'
with pytest.raises(TypeError, match='rem_static_win must be of type bool'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_dict_fields(valid_params_dict):
"""Test dict fields validation."""
# Test low_lim
valid_params_dict['low_lim'] = None
with pytest.raises(ValueError, match='low_lim is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
valid_params_dict['low_lim'] = {'var1': 0.0}
valid_params_dict['upp_lim'] = 'not a dict'
with pytest.raises(TypeError, match='upp_lim must be of type dict'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_bucket_name_none_raises_error(valid_params_dict):
"""Test that None bucket_name raises ValueError."""
valid_params_dict['bucket_name'] = None
with pytest.raises(ValueError, match='bucket_name is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_file_name_none_raises_error(valid_params_dict):
"""Test that None file_name raises ValueError."""
valid_params_dict['file_name'] = None
with pytest.raises(ValueError, match='file_name is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_experiment_run_id_none_raises_error(valid_params_dict):
"""Test that None experiment_run_id raises ValueError."""
valid_params_dict['experiment_run_id'] = None
with pytest.raises(ValueError, match='experiment_run_id is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_experiment_name_none_raises_error(valid_params_dict):
"""Test that None experiment_name raises ValueError."""
valid_params_dict['experiment_name'] = None
with pytest.raises(ValueError, match='experiment_name is required and cannot be None'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_removed_intervals_can_be_none(valid_params_dict):
"""Test that removed_intervals can be None (uses _check_type not _check_none)."""
valid_params_dict['removed_intervals'] = None
params = TrainModelParams.from_dict(valid_params_dict)
assert params.removed_intervals is None
def test_train_model_params_removed_intervals_wrong_type_raises_error(valid_params_dict):
"""Test that wrong type for removed_intervals raises TypeError."""
valid_params_dict['removed_intervals'] = 'not a list'
with pytest.raises(TypeError, match='removed_intervals must be of type list'):
TrainModelParams.from_dict(valid_params_dict)
def test_train_model_params_removed_intervals_with_values(valid_params_dict):
"""Test removed_intervals with actual interval values."""
valid_params_dict['removed_intervals'] = [
('2023-01-01', '2023-01-10'),
('2023-02-01', '2023-02-05'),
]
params = TrainModelParams.from_dict(valid_params_dict)
assert len(params.removed_intervals) == 2
assert params.removed_intervals[0] == ('2023-01-01', '2023-01-10')
def test_train_model_params_all_fields_count():
"""Test that TrainModelParams has exactly 20 required fields."""
import inspect
sig = inspect.signature(TrainModelParams.__init__)
# Subtract 1 for 'self'
param_count = len(sig.parameters) - 1
assert param_count == 20
def test_train_model_params_with_minimal_valid_data():
"""Test creating params with minimal valid data."""
params = TrainModelParams(
variable_columns=['x'],
lag_train=1,
lag_val=1,
target_variable='y',
rem_static_win=False,
low_lim={},
upp_lim={},
window=1,
use_scaler=False,
include_ar=False,
bucket_name='bucket',
file_name='file.csv',
line_separator='\n',
decimal_separator='.',
train_size=50,
shuffle=False,
experiment_run_id=1,
experiment_name='exp',
experiment_description='desc',
removed_intervals=[],
)
assert params.variable_columns == ['x']
assert params.lag_train == 1
assert params.experiment_run_id == 1
def test_train_model_params_check_none_method():
"""Test _check_none method behavior."""
params_dict = {
'variable_columns': ['var1'],
'lag_train': 5,
'lag_val': 3,
'target_variable': 'target',
'rem_static_win': True,
'low_lim': {},
'upp_lim': {},
'window': 10,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'bucket',
'file_name': 'file.csv',
'line_separator': '\n',
'decimal_separator': '.',
'train_size': 80,
'shuffle': True,
'experiment_run_id': 123,
'experiment_name': 'exp',
'experiment_description': 'desc',
'removed_intervals': [],
}
params = TrainModelParams(**params_dict)
# Test that _check_none is a private method
assert hasattr(params, '_check_none')
assert callable(params._check_none)
def test_train_model_params_check_type_method():
"""Test _check_type method behavior."""
params_dict = {
'variable_columns': ['var1'],
'lag_train': 5,
'lag_val': 3,
'target_variable': 'target',
'rem_static_win': True,
'low_lim': {},
'upp_lim': {},
'window': 10,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'bucket',
'file_name': 'file.csv',
'line_separator': '\n',
'decimal_separator': '.',
'train_size': 80,
'shuffle': True,
'experiment_run_id': 123,
'experiment_name': 'exp',
'experiment_description': 'desc',
'removed_intervals': [],
}
params = TrainModelParams(**params_dict)
# Test that _check_type is a private method
assert hasattr(params, '_check_type')
assert callable(params._check_type)

View File

@@ -0,0 +1,246 @@
"""Unit tests for TrainModelResult dataclass."""
from unittest.mock import MagicMock
import pandas as pd
import pytest
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
@pytest.fixture
def sample_params():
"""Create sample TrainModelParams for testing."""
return TrainModelParams(
variable_columns=['var1', 'var2'],
lag_train=5,
lag_val=3,
target_variable='target',
rem_static_win=True,
low_lim={'var1': 0.0, 'var2': 0.0},
upp_lim={'var1': 100.0, 'var2': 100.0},
window=10,
use_scaler=True,
include_ar=False,
bucket_name='test-bucket',
file_name='test-file.csv',
line_separator='\n',
decimal_separator='.',
train_size=80,
shuffle=True,
experiment_run_id=123,
experiment_name='test-experiment',
experiment_description='Test experiment description',
removed_intervals=[],
)
@pytest.fixture
def sample_dataframes():
"""Create sample DataFrames for testing."""
X_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [4, 5, 6]})
X_test = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10]})
y_train = pd.DataFrame({'target': [10, 20, 30]})
y_test = pd.DataFrame({'target': [40, 50]})
return X_train, X_test, y_train, y_test
def test_train_model_result_creation(sample_params, sample_dataframes):
"""Test creating TrainModelResult with required fields."""
x_train, x_test, y_train, y_test = sample_dataframes
process_data = MagicMock()
regr = MagicMock()
scaler_dict = {'var1': {'min': 0, 'max': 100}}
result = TrainModelResult(
params=sample_params,
process_data=process_data,
x_train=x_train,
x_test=x_test,
y_train=y_train,
y_test=y_test,
regr=regr,
scaler_dict=scaler_dict,
)
assert result.params == sample_params
assert result.process_data == process_data
assert result.x_train.equals(x_train)
assert result.x_test.equals(x_test)
assert result.y_train.equals(y_train)
assert result.y_test.equals(y_test)
assert result.regr == regr
assert result.scaler_dict == scaler_dict
def test_train_model_result_optional_fields_default_none(sample_params, sample_dataframes):
"""Test that optional fields default to None."""
x_train, x_test, y_train, y_test = sample_dataframes
result = TrainModelResult(
params=sample_params,
process_data=MagicMock(),
x_train=x_train,
x_test=x_test,
y_train=y_train,
y_test=y_test,
regr=MagicMock(),
scaler_dict={},
)
assert result.y_pred is None
assert result.mse_val is None
assert result.mae_val is None
assert result.r2_val is None
assert result.run_name is None
assert result.report_path is None
assert result.train_data_path is None
assert result.test_data_path is None
assert result.run_dir is None
def test_train_model_result_with_metrics(sample_params, sample_dataframes):
"""Test TrainModelResult with metrics populated."""
x_train, x_test, y_train, y_test = sample_dataframes
y_pred = pd.Series([41, 49])
result = TrainModelResult(
params=sample_params,
process_data=MagicMock(),
x_train=x_train,
x_test=x_test,
y_train=y_train,
y_test=y_test,
regr=MagicMock(),
scaler_dict={},
y_pred=y_pred,
mse_val=1.5,
mae_val=1.2,
r2_val=0.95,
)
assert result.y_pred.equals(y_pred)
assert result.mse_val == 1.5
assert result.mae_val == 1.2
assert result.r2_val == 0.95
def test_train_model_result_with_artifact_paths(sample_params, sample_dataframes):
"""Test TrainModelResult with artifact paths populated."""
x_train, x_test, y_train, y_test = sample_dataframes
result = TrainModelResult(
params=sample_params,
process_data=MagicMock(),
x_train=x_train,
x_test=x_test,
y_train=y_train,
y_test=y_test,
regr=MagicMock(),
scaler_dict={},
run_name='test-experiment-1',
report_path='/path/to/report.html',
train_data_path='/path/to/train_data.csv',
test_data_path='/path/to/test_data.csv',
run_dir='/path/to/run_dir',
)
assert result.run_name == 'test-experiment-1'
assert result.report_path == '/path/to/report.html'
assert result.train_data_path == '/path/to/train_data.csv'
assert result.test_data_path == '/path/to/test_data.csv'
assert result.run_dir == '/path/to/run_dir'
def test_train_model_result_is_dataclass(sample_params, sample_dataframes):
"""Test that TrainModelResult is a dataclass."""
x_train, x_test, y_train, y_test = sample_dataframes
result = TrainModelResult(
params=sample_params,
process_data=MagicMock(),
x_train=x_train,
x_test=x_test,
y_train=y_train,
y_test=y_test,
regr=MagicMock(),
scaler_dict={},
)
# Dataclasses have __dataclass_fields__ attribute
assert hasattr(result, '__dataclass_fields__')
assert 'params' in result.__dataclass_fields__
assert 'process_data' in result.__dataclass_fields__
assert 'x_train' in result.__dataclass_fields__
def test_train_model_result_field_count():
"""Test that TrainModelResult has exactly 17 fields."""
from dataclasses import fields
result_fields = fields(TrainModelResult)
assert len(result_fields) == 17
field_names = {f.name for f in result_fields}
expected_fields = {
'params',
'process_data',
'x_train',
'x_test',
'y_train',
'y_test',
'regr',
'scaler_dict',
'y_pred',
'mse_val',
'mae_val',
'r2_val',
'run_name',
'report_path',
'train_data_path',
'test_data_path',
'run_dir',
}
assert field_names == expected_fields
def test_train_model_result_complete_workflow(sample_params, sample_dataframes):
"""Test TrainModelResult through a complete workflow simulation."""
x_train, x_test, y_train, y_test = sample_dataframes
# Step 1: Create result after training
result = TrainModelResult(
params=sample_params,
process_data=MagicMock(),
x_train=x_train,
x_test=x_test,
y_train=y_train,
y_test=y_test,
regr=MagicMock(),
scaler_dict={'var1': {'min': 0, 'max': 100}},
)
# Step 2: Add predictions and metrics
result.y_pred = pd.Series([41, 49])
result.mse_val = 1.5
result.mae_val = 1.2
result.r2_val = 0.95
# Step 3: Add artifact paths
result.run_name = 'test-experiment-1'
result.report_path = '/path/to/report.html'
result.train_data_path = '/path/to/train_data.csv'
result.test_data_path = '/path/to/test_data.csv'
result.run_dir = '/path/to/run_dir'
# Verify all fields are populated
assert result.y_pred is not None
assert result.mse_val == 1.5
assert result.mae_val == 1.2
assert result.r2_val == 0.95
assert result.run_name == 'test-experiment-1'
assert result.report_path == '/path/to/report.html'
assert result.train_data_path == '/path/to/train_data.csv'
assert result.test_data_path == '/path/to/test_data.csv'
assert result.run_dir == '/path/to/run_dir'

View File

@@ -0,0 +1,499 @@
from datetime import UTC, datetime
from unittest.mock import ANY, MagicMock, call, patch
import numpy as np
import pytest
from pandas import DataFrame, Timestamp
from model_manager.utils.repository.model_repository import MLFlowRepository
@pytest.fixture
def mlflow_repository():
with patch(
'model_manager.utils.repository.model_repository.ModelServing', autospec=True
) as mock_model_serving:
mock_instance = mock_model_serving.return_value
mock_instance.get_transformed_data = MagicMock()
repo = MLFlowRepository(
host='http://localhost:5000', username='admin', password='admin', logger=MagicMock()
)
return repo
metadata = {
'metadata': {
'model_id': 'test_model',
'model_name': 'test_model',
'workflow_name': 'test_workflow',
'schema_name': 'test_schedule',
},
}
class Any:
pass
invalid_cases = [
({'value': {'2024-01-01 12:00:00': 1, 2024: 2}}),
({'value': {'2024-01-01': 1, '2024-01-02': 2}}),
({'value': {Any(): 1, Any(): 2}}),
]
@pytest.mark.parametrize('data', invalid_cases)
def test_detect_and_parse_datetime_index_error_cases(mlflow_repository, data):
input_data = DataFrame(data)
with pytest.raises(ValueError) as e:
mlflow_repository.detect_and_parse_datetime_index(input_data, metadata['metadata'])
assert (
str(e)
== 'Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format %Y-%m-%d %H:%M:%S'
)
valid_cases = [
(
{'value': {'2024-01-01 12:00:00+0000': 1, '2024-01-02 12:00:00+0000': 2}},
['2024-01-01 12:00:00+0000', '2024-01-02 12:00:00+0000'],
),
(
{
'value': {
datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC): 1,
datetime(2025, 1, 2, 12, 0, 0, tzinfo=UTC): 2,
}
},
['2025-01-01 12:00:00+0000', '2025-01-02 12:00:00+0000'],
),
(
{
'value': {
Timestamp(2026, 1, 1, 12, 0, 0, tzinfo=UTC): 1,
Timestamp(2026, 1, 2, 12, 0, 0, tzinfo=UTC): 2,
}
},
['2026-01-01 12:00:00+0000', '2026-01-02 12:00:00+0000'],
),
]
@pytest.mark.parametrize('data,expected', valid_cases)
def test_detect_and_parse_datetime_index_valid_format(mlflow_repository, data, expected):
input_data = DataFrame(data)
response = mlflow_repository.detect_and_parse_datetime_index(input_data, metadata['metadata'])
assert response.index.tolist() == expected
def test_transform_success(mlflow_repository):
data = MagicMock()
model_name = 'model'
mlflow_repository.detect_and_parse_datetime_index = MagicMock()
output = mlflow_repository.transform(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
model_name, data, 0, 'sklearn', False, 'model', 'predict'
)
mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with(
mlflow_repository.model_serving.get_cached_transform.return_value, metadata['metadata']
)
assert output == {
'success': True,
'content': mlflow_repository.detect_and_parse_datetime_index.return_value.to_dict.return_value,
}
def test_transform_error(mlflow_repository):
data = MagicMock()
model_name = 'model'
mlflow_repository.model_serving.get_cached_transform.side_effect = Exception('error')
output = mlflow_repository.transform(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
model_name, data, 0, 'sklearn', False, 'model', 'predict'
)
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
def test_predict_success(mlflow_repository):
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_name = 'model'
mlflow_repository.model_serving.get_cached_predict.return_value = np.array([2, 3])
output = mlflow_repository.predict(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_predict.assert_called_once_with(
model_name, data, 0, 'pyfunc', False, 'model'
)
assert output['success'] is True
assert output['content'] == {
'prediction': {'index_1': 2, 'index_2': 3},
'response_time': {'index_1': ANY, 'index_2': ANY},
}
def test_predict_error(mlflow_repository):
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_name = 'model'
mlflow_repository.model_serving.get_cached_predict = MagicMock(side_effect=Exception('error'))
output = mlflow_repository.predict(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_predict.assert_called_once_with(
model_name, data, 0, 'pyfunc', False, 'model'
)
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_by_run_id(mlflow, mlflow_repository):
mlflow.get_run.return_value = MagicMock(
info=MagicMock(
experiment_id='0',
)
)
mlflow.get_experiment.return_value = MagicMock()
mlflow.get_experiment.return_value.name = 'test'
output = mlflow_repository.get_experiment_by_run_id('0')
assert output == 'test'
mlflow.get_run.assert_called_once_with('0')
mlflow.get_experiment.assert_called_once_with('0')
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_next_run_name(mlflow, mlflow_repository):
mlflow.search_runs.return_value = [1, 2, 3]
output = mlflow_repository.get_next_run_name('run')
assert output == 'run-4'
mlflow.search_runs.assert_called_once_with(
experiment_names=['run'],
order_by=['start_time desc'],
)
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_success(mlflow, mlflow_repository):
mlflow.get_experiment_by_name.return_value = MagicMock(experiment_id='0')
output = mlflow_repository.get_experiment('test')
assert output == '0'
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_error(mlflow, mlflow_repository):
mlflow.get_experiment_by_name.return_value = None
try:
mlflow_repository.get_experiment('test')
except ValueError as e:
assert str(e) == 'Experiment test not found'
else:
raise AssertionError('Expected exception')
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_last_run(mlflow, mlflow_repository):
mlflow.search_runs.return_value = DataFrame(
{
'params.retrain': ['True', 'False', 'True', 'False'],
'end_time': ['2021-01-01', '2021-01-02', '2021-01-03', '2021-01-04'],
'run_id': ['0', '1', '2', '3'],
}
)
output = mlflow_repository.get_experiment_last_run(0)
mlflow.search_runs.assert_called_once_with(
experiment_ids=[0],
filter_string='',
output_format='pandas',
)
assert output == '2'
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_last_run_error(mlflow, mlflow_repository):
mlflow.search_runs.return_value = []
try:
mlflow_repository.get_experiment_last_run(0)
except ValueError as e:
assert str(e) == 'Runs is not a pandas DataFrame'
else:
raise AssertionError('Expected exception')
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn')
@patch('model_manager.utils.repository.model_repository.mlflow.set_experiment')
def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
mlflow_repository.model_serving.get_model_info = MagicMock(return_value='0')
mlflow_repository.model_serving.get_model_uri = MagicMock(return_value='test')
mlflow_repository.get_experiment_by_run_id = MagicMock()
data_model_mock = MagicMock()
prediction_model_mock = MagicMock()
sklearn.load_model.side_effect = [data_model_mock, prediction_model_mock]
data_model_mock.fit.return_value = data_model_mock
data_model_mock.predict.return_value = DataFrame(
{
'x': [10, 20, 30],
}
)
data_model_mock.target_variable = 'y'
prediction_model_mock.fit.return_value = prediction_model_mock
data = DataFrame({'x': [1, 2, 3], 'y': [4, 5, 6]})
output = mlflow_repository.create_model_experiment('test', data)
mlflow_repository.model_serving.get_model_info.assert_called_once_with('test')
mlflow_repository.model_serving.get_model_uri.assert_called_once_with('0', prediction=False)
sklearn.load_model.assert_has_calls(
[
call(mlflow_repository.model_serving.get_model_uri.return_value),
call('models:/test/production'),
]
)
assert sklearn.load_model.call_count == 2
data_model_mock.fit.assert_called_once_with(data)
data_model_mock.predict.assert_called_once_with(data)
fit_args = prediction_model_mock.fit.call_args[0][0]
assert fit_args.equals(
DataFrame(
{
'x': [10, 20, 30],
'y': [4, 5, 6],
}
)
)
mlflow_repository.get_experiment_by_run_id.assert_called_once_with('0')
set_experiment.assert_called_once_with(mlflow_repository.get_experiment_by_run_id.return_value)
assert output == (
prediction_model_mock,
data_model_mock,
mlflow_repository.get_experiment_by_run_id.return_value,
)
@patch('model_manager.utils.repository.model_repository.path.exists')
@patch('model_manager.utils.repository.model_repository.remove')
@patch('model_manager.utils.repository.model_repository.mlflow.start_run')
@patch('model_manager.utils.repository.model_repository.mlflow.log_param')
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn.log_model')
@patch('model_manager.utils.repository.model_repository.mlflow.log_artifact')
def test_perform_model_retrain(
log_artifact, log_model, log_param, start_run, mock_remove, mock_path_exists, mlflow_repository
):
# Create mock models with attributes to test the for loops (lines 268-274)
prediction_model_mock = MagicMock()
prediction_model_mock.__dict__ = {'model': 'pred_model', 'param1': 'value1', 'param2': 'value2'}
data_model_mock = MagicMock()
data_model_mock.__dict__ = {'model': 'data_model', 'param3': 'value3', 'param4': 'value4'}
experiment = 'test'
model_name = 'test'
data = MagicMock()
mlflow_repository.get_next_run_name = MagicMock(return_value='test-1')
run = MagicMock()
start_run.__enter__.return_value = run
mock_path_exists.return_value = True
output = mlflow_repository.perform_model_retrain(
prediction_model_mock, data_model_mock, experiment, model_name, data
)
mlflow_repository.get_next_run_name.assert_called_once_with(experiment)
start_run.assert_called_once_with(
run_name='test-1', description='Retrain model test with new data'
)
log_model.assert_has_calls(
[
call(data_model_mock, 'data_model'),
call(prediction_model_mock, 'prediction_model'),
]
)
data.to_csv.assert_called_once_with('temp/raw_data_test.csv', index=True)
log_artifact.assert_called_once_with('temp/raw_data_test.csv')
# Verify that model attributes were logged (excluding 'model' key)
log_param.assert_has_calls(
[
call('param1', 'value1'), # from prediction_model
call('param2', 'value2'), # from prediction_model
call('param3', 'value3'), # from data_model
call('param4', 'value4'), # from data_model
call('retrain', True),
],
any_order=True,
)
# Verify temp file cleanup
mock_path_exists.assert_called_once_with('temp/raw_data_test.csv')
mock_remove.assert_called_once_with('temp/raw_data_test.csv')
assert output == ('Model retrained successfully', experiment)
@patch('model_manager.utils.repository.model_repository.path.exists')
@patch('model_manager.utils.repository.model_repository.remove')
@patch('model_manager.utils.repository.model_repository.mlflow.start_run')
@patch('model_manager.utils.repository.model_repository.mlflow.log_param')
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn.log_model')
@patch('model_manager.utils.repository.model_repository.mlflow.log_artifact')
def test_perform_model_retrain_file_not_exists(
log_artifact, log_model, log_param, start_run, mock_remove, mock_path_exists, mlflow_repository
):
"""Test perform_model_retrain when temp file doesn't exist (line 291->294 branch)."""
prediction_model_mock = MagicMock()
prediction_model_mock.__dict__ = {'model': 'pred_model'}
data_model_mock = MagicMock()
data_model_mock.__dict__ = {'model': 'data_model'}
experiment = 'test'
model_name = 'test'
data = MagicMock()
mlflow_repository.get_next_run_name = MagicMock(return_value='test-1')
run = MagicMock()
start_run.__enter__.return_value = run
mock_path_exists.return_value = False # File doesn't exist
output = mlflow_repository.perform_model_retrain(
prediction_model_mock, data_model_mock, experiment, model_name, data
)
# Verify temp file cleanup was checked but not executed
mock_path_exists.assert_called_once_with('temp/raw_data_test.csv')
mock_remove.assert_not_called() # Should not be called when file doesn't exist
assert output == ('Model retrained successfully', experiment)
def test_retrain_model(mlflow_repository):
data = MagicMock()
model_name = 'test'
mlflow_repository.create_model_experiment = MagicMock(
return_value=('data_model', 'prediction_model', '0')
)
mlflow_repository.perform_model_retrain = MagicMock(return_value='Model retrained successfully')
output = mlflow_repository.retrain_model(data, model_name)
mlflow_repository.create_model_experiment.assert_called_once_with(model_name, data)
mlflow_repository.perform_model_retrain.assert_called_once_with(
'data_model', 'prediction_model', '0', model_name, data
)
assert output == 'Model retrained successfully'
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_update_production_model_by_run_id(mlflow, mlflow_repository):
client_mock = MagicMock()
mlflow.tracking.MlflowClient.return_value = client_mock
client_mock.get_registered_model.return_value = MagicMock(
latest_versions=[
MagicMock(version='1'),
MagicMock(version='2'),
MagicMock(version='3'),
]
)
output = mlflow_repository.update_production_model_by_run_id('0', 'test')
mlflow.register_model.assert_called_once_with(
'runs:/0/prediction_model',
'test',
)
mlflow.tracking.MlflowClient.assert_called_once()
client_mock.get_registered_model.assert_called_once_with('test')
client_mock.transition_model_version_stage.assert_called_once_with(
name='test',
version='3',
stage='Production',
archive_existing_versions=True,
)
assert output == {
'model_name': 'test',
'version': '3',
'mlflow_run_id': '0',
}
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
mlflow.tracking.MlflowClient.return_value = MagicMock(
get_registered_model=MagicMock(return_value=MagicMock(latest_versions={}))
)
try:
mlflow_repository.update_production_model_by_run_id('0', 'test')
except Exception as e: # noqa: BLE001
assert str(e) == 'Model versions is not a list'
else:
raise AssertionError('Expected exception')
def test_update_production_model(mlflow_repository):
connector = mlflow_repository
with patch.object(connector, 'get_experiment', return_value='0') as get_experiment:
with patch.object(
connector, 'get_experiment_last_run', return_value='2'
) as get_experiment_last_run:
with patch.object(
connector,
'update_production_model_by_run_id',
return_value={'model_name': 'test', 'version': '3', 'mlflow_run_id': '0'},
) as update_production_model_by_run_id:
output = connector.update_production_model('0', 'test')
get_experiment.assert_called_once_with('0')
get_experiment_last_run.assert_called_once_with('0')
update_production_model_by_run_id.assert_called_once_with('2', 'test')
assert output == {
'model_name': 'test',
'version': '3',
'mlflow_run_id': '0',
'mlflow_experiment_id': '0',
}

View File

@@ -0,0 +1,171 @@
from os import environ
from model_manager.utils.connectors_config import (
build_minio_config,
build_mlflow_config,
build_mongodb_config,
build_postgres_config,
)
def test_build_mlflow_config_with_env_vars():
# Arrange
environ['MLFLOW_HOST'] = 'http://test-host'
environ['MLFLOW_PORT'] = '8080'
environ['MLFLOW_USERNAME'] = 'test-user'
environ['MLFLOW_PASSWORD'] = 'test-pass'
# Act
config = build_mlflow_config()
# Assert
assert config['host'] == 'http://test-host'
assert config['port'] == 8080
assert config['username'] == 'test-user'
assert config['password'] == 'test-pass'
def test_build_mlflow_config_with_defaults():
# Arrange
# Clear any existing env vars
environ.pop('MLFLOW_HOST', None)
environ.pop('MLFLOW_PORT', None)
environ.pop('MLFLOW_USERNAME', None)
environ.pop('MLFLOW_PASSWORD', None)
# Act
config = build_mlflow_config()
# Assert
assert config['host'] == 'http://localhost'
assert config['port'] == 5080
assert config['username'] == 'aignosi'
assert config['password'] == 'aignosi'
def test_build_postgres_config_with_env_vars():
# Arrange
environ['POSTGRES_HOST'] = 'test-host'
environ['POSTGRES_PORT'] = '5433'
environ['POSTGRES_USER'] = 'test-user'
environ['POSTGRES_PASSWORD'] = 'test-pass'
environ['POSTGRES_DBNAME'] = 'test-db'
environ['POSTGRES_MIN_CONNECTIONS'] = '10'
environ['POSTGRES_MAX_CONNECTIONS'] = '30'
# Act
config = build_postgres_config()
# Assert
assert config['host'] == 'test-host'
assert config['port'] == 5433
assert config['user'] == 'test-user'
assert config['password'] == 'test-pass'
assert config['dbname'] == 'test-db'
assert config['min_connections'] == 10
assert config['max_connections'] == 30
def test_build_postgres_config_with_defaults():
# Arrange
environ.pop('POSTGRES_HOST', None)
environ.pop('POSTGRES_PORT', None)
environ.pop('POSTGRES_USER', None)
environ.pop('POSTGRES_PASSWORD', None)
environ.pop('POSTGRES_DBNAME', None)
environ.pop('POSTGRES_MIN_CONNECTIONS', None)
environ.pop('POSTGRES_MAX_CONNECTIONS', None)
# Act
config = build_postgres_config()
# Assert
assert config['host'] == 'localhost'
assert config['port'] == 5432
assert config['user'] == 'sientia'
assert config['password'] == 'sientia'
assert config['dbname'] == 'sientia'
assert config['min_connections'] == 5
assert config['max_connections'] == 20
def test_build_mongo_db_config_with_env_vars():
environ['MONGODB_USERNAME'] = 'sientia1'
environ['MONGODB_PASSWORD'] = 'sientia1'
environ['MONGODB_URL'] = 'localhost:27018'
environ['MONGODB_DATABASE_NAME'] = 'test_db'
environ['MONGODB_TTL_INDEX_HOURS'] = '1'
assert build_mongodb_config() == {
'connection_string': 'mongodb://sientia1:sientia1@localhost:27018',
'database_name': 'test_db',
'ttl_index_seconds': 3600,
}
def test_build_mongo_db_config_with_defaults():
environ.pop('MONGODB_USERNAME', None)
environ.pop('MONGODB_PASSWORD', None)
environ.pop('MONGODB_DATABASE_NAME', None)
environ.pop('MONGODB_URL', None)
environ.pop('MONGODB_TTL_INDEX_HOURS', None)
assert build_mongodb_config() == {
'connection_string': 'mongodb://root:wKZDbMNU1c@localhost:27018',
'database_name': 'sientia',
'ttl_index_seconds': 3600,
}
def test_build_minio_config_with_env_vars():
# Arrange
environ['MINIO_ENDPOINT_URL'] = 'http://test-minio:9000'
environ['MINIO_ACCESS_KEY'] = 'test-access-key'
environ['MINIO_SECRET_KEY'] = 'test-secret-key'
environ['MINIO_REGION'] = 'eu-west-1'
environ['MINIO_USE_SSL'] = 'true'
environ['MINIO_MAX_RETRY_ATTEMPTS'] = '5'
environ['MINIO_RETRY_MODE'] = 'standard'
environ['MINIO_CONNECT_TIMEOUT'] = '20'
environ['MINIO_READ_TIMEOUT'] = '120'
# Act
config = build_minio_config()
# Assert
assert config['endpoint_url'] == 'http://test-minio:9000'
assert config['access_key'] == 'test-access-key'
assert config['secret_key'] == 'test-secret-key'
assert config['region'] == 'eu-west-1'
assert config['use_ssl'] is True
assert config['max_retry_attempts'] == 5
assert config['retry_mode'] == 'standard'
assert config['connect_timeout'] == 20
assert config['read_timeout'] == 120
def test_build_minio_config_with_defaults():
# Arrange
# Clear any existing env vars
environ.pop('MINIO_ENDPOINT_URL', None)
environ.pop('MINIO_ACCESS_KEY', None)
environ.pop('MINIO_SECRET_KEY', None)
environ.pop('MINIO_REGION', None)
environ.pop('MINIO_USE_SSL', None)
environ.pop('MINIO_MAX_RETRY_ATTEMPTS', None)
environ.pop('MINIO_RETRY_MODE', None)
environ.pop('MINIO_CONNECT_TIMEOUT', None)
environ.pop('MINIO_READ_TIMEOUT', None)
# Act
config = build_minio_config()
# Assert
assert config['endpoint_url'] == 'http://localhost:9000'
assert config['access_key'] == 'minioadmin'
assert config['secret_key'] == 'minioadmin'
assert config['region'] == 'us-east-1'
assert config['use_ssl'] is False
assert config['max_retry_attempts'] == 3
assert config['retry_mode'] == 'adaptive'
assert config['connect_timeout'] == 10
assert config['read_timeout'] == 60