SIENTIAPDE-1579: Updated tests and refactored apply filters method
This commit is contained in:
@@ -9,9 +9,34 @@ from pytest import raises
|
||||
from model_manager.sientia.models import (
|
||||
DataPreprocessor,
|
||||
LinearRegressionModel,
|
||||
_frontend_date_format_to_strftime,
|
||||
)
|
||||
|
||||
|
||||
class TestFrontendDateFormatToStrftime:
|
||||
"""Tests for _frontend_date_format_to_strftime (models module)."""
|
||||
|
||||
def test_none_or_empty_returns_none_or_empty(self):
|
||||
"""None or empty string returns None or falsy."""
|
||||
assert _frontend_date_format_to_strftime(None) is None
|
||||
assert _frontend_date_format_to_strftime('') is None
|
||||
|
||||
def test_dd_mm_yyyy_hh_mm_ss(self):
|
||||
"""Converts dd/MM/yyyy HH:mm:ss to strftime."""
|
||||
result = _frontend_date_format_to_strftime('dd/MM/yyyy HH:mm:ss')
|
||||
assert result == '%d/%m/%Y %H:%M:%S'
|
||||
|
||||
def test_yyyy_mm_dd(self):
|
||||
"""Converts yyyy-MM-dd to strftime."""
|
||||
result = _frontend_date_format_to_strftime('yyyy-MM-dd')
|
||||
assert result == '%Y-%m-%d'
|
||||
|
||||
def test_iso_datetime(self):
|
||||
"""Converts yyyy-MM-dd HH:mm:ss to strftime."""
|
||||
result = _frontend_date_format_to_strftime('yyyy-MM-dd HH:mm:ss')
|
||||
assert result == '%Y-%m-%d %H:%M:%S'
|
||||
|
||||
|
||||
class _IterableWithContains:
|
||||
def __init__(self, iterable, contains_values):
|
||||
self._iterable = iterable
|
||||
|
||||
@@ -13,6 +13,8 @@ from model_manager.utils.models.train_model_result import TrainModelResult
|
||||
from model_manager.utils.repository.training_repository import (
|
||||
TrainingRepository,
|
||||
_apply_support_filters,
|
||||
_ensure_date_column_parsed,
|
||||
_frontend_format_to_strftime,
|
||||
)
|
||||
|
||||
|
||||
@@ -991,9 +993,7 @@ class TestConfigureDatetimeIndex:
|
||||
assert 'var2' in result.columns
|
||||
|
||||
@pytest.mark.filterwarnings('ignore::UserWarning')
|
||||
def test_configure_datetime_index_invalid_timestamp_column(
|
||||
self, training_repo, datetime_params
|
||||
):
|
||||
def test_configure_datetime_index_invalid_timestamp_column(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index with invalid timestamp values."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -1118,6 +1118,109 @@ class TestExtractModelEquationPolynomial:
|
||||
assert 'equation_string' in result
|
||||
assert 'latex_equation' in result
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for _frontend_format_to_strftime and _ensure_date_column_parsed
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestFrontendFormatToStrftime:
|
||||
"""Tests for _frontend_format_to_strftime."""
|
||||
|
||||
def test_empty_returns_unchanged(self):
|
||||
"""Empty string is returned as-is."""
|
||||
assert _frontend_format_to_strftime('') == ''
|
||||
|
||||
def test_dd_mm_yyyy_hh_mm_ss(self):
|
||||
"""Converts dd/MM/yyyy HH:mm:ss to strftime."""
|
||||
result = _frontend_format_to_strftime('dd/MM/yyyy HH:mm:ss')
|
||||
assert result == '%d/%m/%Y %H:%M:%S'
|
||||
|
||||
def test_yyyy_mm_dd(self):
|
||||
"""Converts yyyy-MM-dd to strftime."""
|
||||
result = _frontend_format_to_strftime('yyyy-MM-dd')
|
||||
assert result == '%Y-%m-%d'
|
||||
|
||||
def test_iso_like_datetime(self):
|
||||
"""Converts yyyy-MM-dd HH:mm:ss to strftime."""
|
||||
result = _frontend_format_to_strftime('yyyy-MM-dd HH:mm:ss')
|
||||
assert result == '%Y-%m-%d %H:%M:%S'
|
||||
|
||||
|
||||
class TestEnsureDateColumnParsed:
|
||||
"""Tests for _ensure_date_column_parsed."""
|
||||
|
||||
@pytest.fixture
|
||||
def date_params(self):
|
||||
"""Params with date_column and date_format set."""
|
||||
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='ts',
|
||||
date_format='yyyy-MM-dd HH:mm:ss',
|
||||
)
|
||||
|
||||
def test_returns_unchanged_when_no_date_column(self, date_params):
|
||||
"""When params.date_column is None, data is returned unchanged."""
|
||||
date_params.date_column = None
|
||||
date_params.date_format = None
|
||||
data = pd.DataFrame({'ts': ['2023-01-01'], 'x': [1]})
|
||||
result = _ensure_date_column_parsed(data, date_params)
|
||||
pd.testing.assert_frame_equal(result, data)
|
||||
|
||||
def test_returns_unchanged_when_column_missing(self, date_params):
|
||||
"""When date_column not in data columns, data is returned unchanged."""
|
||||
data = pd.DataFrame({'other': [1], 'x': [2]})
|
||||
result = _ensure_date_column_parsed(data, date_params)
|
||||
pd.testing.assert_frame_equal(result, data)
|
||||
|
||||
def test_parses_column_with_format(self, date_params):
|
||||
"""When date_column and date_format set, column is parsed as datetime."""
|
||||
data = pd.DataFrame({
|
||||
'ts': ['2023-01-01 10:00:00', '2023-06-15 14:30:00'],
|
||||
'x': [1, 2],
|
||||
})
|
||||
result = _ensure_date_column_parsed(data, date_params)
|
||||
assert result['ts'].dtype == 'datetime64[ns]'
|
||||
assert result['ts'].iloc[0].year == 2023
|
||||
assert result['ts'].iloc[0].month == 1
|
||||
assert result['ts'].iloc[1].month == 6
|
||||
|
||||
def test_invalid_values_coerced_to_nat(self, date_params):
|
||||
"""Invalid date strings are coerced to NaT when format is set."""
|
||||
date_params.date_format = 'yyyy-MM-dd'
|
||||
data = pd.DataFrame({'ts': ['2023-01-01', 'not-a-date', '2023-12-31'], 'x': [1, 2, 3]})
|
||||
result = _ensure_date_column_parsed(data, date_params)
|
||||
assert pd.isna(result['ts'].iloc[1])
|
||||
assert result['ts'].iloc[0].year == 2023
|
||||
assert result['ts'].iloc[2].month == 12
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for _apply_support_filters
|
||||
@@ -1153,12 +1256,10 @@ class TestApplySupportFilters:
|
||||
|
||||
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],
|
||||
}
|
||||
)
|
||||
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},
|
||||
@@ -1171,12 +1272,10 @@ class TestApplySupportFilters:
|
||||
|
||||
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],
|
||||
}
|
||||
)
|
||||
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},
|
||||
@@ -1189,13 +1288,11 @@ class TestApplySupportFilters:
|
||||
|
||||
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],
|
||||
}
|
||||
)
|
||||
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},
|
||||
|
||||
Reference in New Issue
Block a user