SIENTIAPDE-1579: Updated tests

This commit is contained in:
Kou Kinoshita
2026-02-18 15:05:37 -03:00
parent b6eb346a0b
commit c6d6f94e05
5 changed files with 179 additions and 33 deletions

View File

@@ -41,6 +41,8 @@ def sample_params():
scaler_name='Standard Scaler',
support_filters={},
static_threshold=None,
date_column=None,
date_format=None,
)

View File

@@ -344,9 +344,14 @@ def test_create_run_directory_success(
result = repo._create_run_directory('/tmp/reports', 'test_run') # noqa: S108
expected_path = os.path.join('/tmp/reports/temp', 'test_run_20240101_120000_123456') # noqa: S108
assert result == expected_path
mock_makedirs.assert_called_once_with(expected_path, exist_ok=True)
expected_path = os.path.normpath(
os.path.join('/tmp/reports', 'temp', 'test_run_20240101_120000_123456') # noqa: S108
)
assert os.path.normpath(result) == expected_path
mock_makedirs.assert_called_once()
call_path = mock_makedirs.call_args[0][0]
assert os.path.normpath(call_path) == expected_path
assert mock_makedirs.call_args[1] == {'exist_ok': True}
@patch('model_manager.utils.repository.model_repository.ModelServing')

View File

@@ -10,7 +10,10 @@ import pytest
from model_manager.sientia.models import LinearRegressionModel
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
from model_manager.utils.repository.training_repository import TrainingRepository
from model_manager.utils.repository.training_repository import (
TrainingRepository,
_apply_support_filters,
)
@pytest.fixture
@@ -57,6 +60,8 @@ def sample_params():
scaler_name='None',
support_filters={},
static_threshold=None,
date_column=None,
date_format=None,
)
@@ -139,6 +144,8 @@ class TestExtractModelEquation:
scaler_name='None',
support_filters={},
static_threshold=None,
date_column=None,
date_format=None,
)
# Mock model with single coefficient
@@ -585,7 +592,7 @@ class TestTrain:
"""
return BytesIO(csv_content.encode('utf-8'))
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_basic_workflow(
@@ -631,7 +638,7 @@ class TestTrain:
# Verify split was called
assert mock_split_train_test.called
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_with_scaler(
@@ -661,7 +668,7 @@ class TestTrain:
assert result is not None
assert result.scaler_dict is not None
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_with_shuffle_enabled(
@@ -692,7 +699,7 @@ class TestTrain:
call_kwargs = mock_split_train_test.call_args[1]
assert call_kwargs['shuffle'] is True
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_with_different_train_size(
@@ -723,7 +730,7 @@ class TestTrain:
call_kwargs = mock_split_train_test.call_args[1]
assert call_kwargs['train_size'] == 0.7
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_raises_on_empty_data_after_transform(
@@ -744,7 +751,7 @@ class TestTrain:
with pytest.raises(ValueError, match='Data view is empty after transformation'):
training_repo.train(sample_csv_data, sample_params)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_logs_success(
@@ -781,7 +788,7 @@ class TestTrain:
'Model trained successfully' in str(call) for call in mock_logger.info.call_args_list
)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_with_custom_separators(
@@ -812,7 +819,7 @@ class TestTrain:
# Verify load_data was called with custom separators
mock_load_data.assert_called_once_with(sample_csv_data, ';', ',')
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
@patch('model_manager.utils.repository.training_repository.split_train_test')
@patch('model_manager.utils.repository.training_repository.load_data')
def test_train_result_contains_all_fields(
@@ -862,19 +869,55 @@ class TestConfigureDatetimeIndex:
"""Create a TrainingRepository instance."""
return TrainingRepository(logger=mock_logger)
def test_configure_datetime_index_already_datetime(self, training_repo):
@pytest.fixture
def datetime_params(self):
"""Minimal params for _configure_datetime_index (date_column/date_format can be None)."""
return TrainModelParams(
experiment_run_id=1,
experiment_name='test',
target_variable='y',
variable_columns=['x'],
lag_train={'x': 0},
lag_val={'x': 0},
rem_static_win=False,
low_lim={},
upp_lim={},
window=0,
use_scaler=False,
include_ar=False,
train_size=80,
shuffle=True,
bucket_name='b',
file_name='f.csv',
line_separator=',',
decimal_separator='.',
removed_intervals=[],
model_name='Linear Regression',
degree=1,
interaction_only=False,
nan_treatment='drop',
start_date=None,
end_date=None,
scaler_name='None',
support_filters={},
static_threshold=None,
date_column=None,
date_format=None,
)
def test_configure_datetime_index_already_datetime(self, training_repo, datetime_params):
"""Test _configure_datetime_index when index is already DatetimeIndex."""
data = pd.DataFrame(
{'var1': [1, 2, 3], 'var2': [4, 5, 6]},
index=pd.to_datetime(['2023-01-01', '2023-01-02', '2023-01-03']),
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
assert isinstance(result.index, pd.DatetimeIndex)
assert len(result) == 3
def test_configure_datetime_index_with_timestamp_column(self, training_repo):
def test_configure_datetime_index_with_timestamp_column(self, training_repo, datetime_params):
"""Test _configure_datetime_index with 'timestamp' column."""
data = pd.DataFrame(
{
@@ -884,12 +927,12 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
assert isinstance(result.index, pd.DatetimeIndex)
assert 'timestamp' not in result.columns
def test_configure_datetime_index_with_date_column(self, training_repo):
def test_configure_datetime_index_with_date_column(self, training_repo, datetime_params):
"""Test _configure_datetime_index with 'date' column."""
data = pd.DataFrame(
{
@@ -898,12 +941,12 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
assert isinstance(result.index, pd.DatetimeIndex)
assert 'date' not in result.columns
def test_configure_datetime_index_with_datetime_column(self, training_repo):
def test_configure_datetime_index_with_datetime_column(self, training_repo, datetime_params):
"""Test _configure_datetime_index with 'datetime' column."""
data = pd.DataFrame(
{
@@ -912,12 +955,12 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
assert isinstance(result.index, pd.DatetimeIndex)
assert 'datetime' not in result.columns
def test_configure_datetime_index_first_column_datetime(self, training_repo):
def test_configure_datetime_index_first_column_datetime(self, training_repo, datetime_params):
"""Test _configure_datetime_index when first column looks like datetime."""
data = pd.DataFrame(
{
@@ -926,13 +969,13 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
assert isinstance(result.index, pd.DatetimeIndex)
assert 'my_date' not in result.columns
@pytest.mark.filterwarnings('ignore::UserWarning')
def test_configure_datetime_index_no_timestamp_column(self, training_repo):
def test_configure_datetime_index_no_timestamp_column(self, training_repo, datetime_params):
"""Test _configure_datetime_index when no timestamp column found."""
data = pd.DataFrame(
{
@@ -941,14 +984,14 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
# Should return original data unchanged (no valid datetime columns)
assert 'var1' in result.columns
assert 'var2' in result.columns
@pytest.mark.filterwarnings('ignore::UserWarning')
def test_configure_datetime_index_invalid_timestamp_column(self, training_repo):
def test_configure_datetime_index_invalid_timestamp_column(self, training_repo, datetime_params):
"""Test _configure_datetime_index with invalid timestamp values."""
data = pd.DataFrame(
{
@@ -957,13 +1000,13 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
# Should skip invalid column and try first column
assert 'var1' in result.columns
@pytest.mark.filterwarnings('ignore::UserWarning')
def test_configure_datetime_index_invalid_first_column(self, training_repo):
def test_configure_datetime_index_invalid_first_column(self, training_repo, datetime_params):
"""Test _configure_datetime_index when first column is not datetime."""
data = pd.DataFrame(
{
@@ -972,13 +1015,13 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
# Should return original data unchanged
assert 'var1' in result.columns
assert 'var2' in result.columns
def test_configure_datetime_index_first_column_all_nan(self, training_repo):
def test_configure_datetime_index_first_column_all_nan(self, training_repo, datetime_params):
"""Test _configure_datetime_index when first column has all NaN values."""
data = pd.DataFrame(
{
@@ -987,7 +1030,7 @@ class TestConfigureDatetimeIndex:
}
)
result = training_repo._configure_datetime_index(data)
result = training_repo._configure_datetime_index(data, datetime_params)
# Should return original data unchanged (first column has no valid values)
assert 'first_col' in result.columns
@@ -1072,3 +1115,98 @@ class TestExtractModelEquationPolynomial:
assert 'equation_string' in result
assert 'latex_equation' in result
# ============================================================================
# Tests for _apply_support_filters
# ============================================================================
class TestApplySupportFilters:
"""Tests for _apply_support_filters function."""
def test_empty_support_filters_returns_unchanged(self):
"""When support_filters is empty, data_view is returned unchanged."""
data = pd.DataFrame({'x': [1, 2, 3], 'target': [10, 20, 30]})
result = _apply_support_filters(data, 'target', {})
pd.testing.assert_frame_equal(result, data)
def test_target_not_in_columns_returns_unchanged(self):
"""When target_variable is not in data_view columns, return unchanged."""
data = pd.DataFrame({'x': [1, 2, 3], 'y': [10, 20, 30]})
result = _apply_support_filters(data, 'target', {'x': {}})
pd.testing.assert_frame_equal(result, data)
def test_variable_not_in_columns_skipped(self):
"""When a filter variable is not in data_view, that variable is skipped."""
data = pd.DataFrame({'x': [1, 2, 3], 'target': [10, 20, 30]})
support_filters = {
'missing_var': {
'upper_line': {'intercept': 100, 'angle': 10},
'lower_line': {'intercept': 0, 'angle': -10},
},
}
result = _apply_support_filters(data, 'target', support_filters)
pd.testing.assert_frame_equal(result, data)
def test_snake_case_upper_lower_line(self):
"""Support filters with upper_line/lower_line (snake_case) filter rows."""
data = pd.DataFrame({
'x': [1.0, 2.0, 3.0, 4.0],
'target': [2.0, 4.0, 6.0, 8.0],
})
support_filters = {
'x': {
'upper_line': {'intercept': 1.0, 'angle': 50},
'lower_line': {'intercept': -1.0, 'angle': -50},
},
}
result = _apply_support_filters(data, 'target', support_filters)
assert len(result) <= 4
assert list(result.columns) == ['x', 'target']
def test_camel_case_upper_lower_line(self):
"""Support filters with upperLine/lowerLine (camelCase) are accepted."""
data = pd.DataFrame({
'x': [1.0, 2.0, 3.0],
'target': [1.0, 2.0, 3.0],
})
support_filters = {
'x': {
'upperLine': {'intercept': 2, 'angle': 5},
'lowerLine': {'intercept': 0, 'angle': -5},
},
}
result = _apply_support_filters(data, 'target', support_filters)
assert len(result) <= 3
assert list(result.columns) == ['x', 'target']
def test_two_variables_ands_masks(self):
"""Two variables apply AND of both masks."""
data = pd.DataFrame({
'a': [1.0, 2.0, 3.0],
'b': [1.0, 2.0, 3.0],
'target': [2.0, 2.0, 2.0],
})
support_filters = {
'a': {
'upper_line': {'intercept': 10, 'angle': 45},
'lower_line': {'intercept': -10, 'angle': -45},
},
'b': {
'upper_line': {'intercept': 10, 'angle': 45},
'lower_line': {'intercept': -10, 'angle': -45},
},
}
result = _apply_support_filters(data, 'target', support_filters)
assert len(result) <= 3
assert list(result.columns) == ['a', 'b', 'target']
def test_missing_upper_or_lower_skips_variable(self):
"""If upper_line or lower_line is missing, that variable is skipped."""
data = pd.DataFrame({'x': [1, 2, 3], 'target': [10, 20, 30]})
support_filters = {
'x': {'upper_line': {'intercept': 100, 'angle': 0}},
}
result = _apply_support_filters(data, 'target', support_filters)
pd.testing.assert_frame_equal(result, data)