Files
sientia-dataops-model-manager/tests/utils/models/test_train_model_params.py
vitor-aignosi ba9eb3d7c7 feat: require date_column in training parameters and update documentation
- Made `date_column` a required field in `TrainModelParams`, ensuring it must be present in the input data.
- Updated related documentation in `input-sample.md`, `README.md`, and various test scenarios to reflect the change in requirement.
- Adjusted the handling of `date_format` to default to `yyyy-MM-dd HH:mm:ss` if omitted, enhancing usability.
- Refined test scenarios to include new examples and ensure compliance with the updated parameter structure.

These changes improve the robustness of the model training workflow and clarify the expectations for input data.
2026-05-05 08:35:12 -03:00

322 lines
11 KiB
Python

"""Unit tests for TrainModelParams (current schema)."""
import copy
from unittest.mock import patch
import pytest
from model_manager.utils.models.train_model_params import (
DEFAULT_TRAIN_DATE_FORMAT,
TrainModelParams,
validate_frontend_date_format,
)
@pytest.fixture
def minimal_model_metadata() -> dict:
"""Minimal truthy metadata so validate_business_rules passes schema lookup."""
return {'schemas': {'components': {'schemas': {}}}}
@pytest.fixture
def valid_train_params_dict(minimal_model_metadata) -> dict:
"""Valid dictionary for TrainModelParams.from_dict."""
return {
'variable_columns': ['var1', 'var2'],
'target_variable': 'target',
'bucket_name': 'test-bucket',
'file_name': 'test-file.csv',
'line_separator': ',',
'decimal_separator': '.',
'date_column': 'timestamp',
'date_format': 'yyyy-MM-dd HH:mm:ss',
'train_size': 80,
'shuffle': True,
'random_state': 42,
'experiment_run_id': 1,
'model_name': 'Linear Regression',
'val_file_name': None,
'data_model_kwargs': {},
'model_kwargs': {},
'opt_params': {},
'model_type': 'linear_regression',
'model_id': None,
'model_metadata': minimal_model_metadata,
}
def test_from_dict_success(valid_train_params_dict):
"""from_dict builds params and experiment_name from model_name."""
params = TrainModelParams.from_dict(valid_train_params_dict)
assert params.variable_columns == ['var1', 'var2']
assert params.target_variable == 'target'
assert params.bucket_name == 'test-bucket'
assert params.experiment_run_id == 1
assert params.experiment_name == 'Linear Regression'
assert params.model_metadata is valid_train_params_dict['model_metadata']
def test_from_dict_date_format_omitted_uses_default(valid_train_params_dict):
"""Missing date_format defaults to DEFAULT_TRAIN_DATE_FORMAT."""
d = copy.deepcopy(valid_train_params_dict)
del d['date_format']
params = TrainModelParams.from_dict(d)
assert params.date_format == DEFAULT_TRAIN_DATE_FORMAT
def test_from_dict_date_format_blank_uses_default(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['date_format'] = ' '
params = TrainModelParams.from_dict(d)
assert params.date_format == DEFAULT_TRAIN_DATE_FORMAT
def test_from_dict_superfluous_date_column_camel_key_is_ignored(valid_train_params_dict):
"""Only snake_case keys are read; dateColumn does not populate date_column."""
d = copy.deepcopy(valid_train_params_dict)
d['dateColumn'] = 'wrong_name'
params = TrainModelParams.from_dict(d)
assert params.date_column == 'timestamp'
def test_from_dict_missing_date_column_raises(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
del d['date_column']
with pytest.raises(ValueError, match='date_column is required'):
TrainModelParams.from_dict(d)
def test_from_dict_date_format_non_string_raises(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['date_format'] = 12345
with pytest.raises(TypeError, match='date_format must be a string'):
TrainModelParams.from_dict(d)
def test_from_dict_coerces_experiment_run_id_string(valid_train_params_dict):
"""Numeric string experiment_run_id is coerced to int."""
d = copy.deepcopy(valid_train_params_dict)
d['experiment_run_id'] = '42'
params = TrainModelParams.from_dict(d)
assert params.experiment_run_id == 42
def test_from_dict_model_metadata_none(valid_train_params_dict):
"""model_metadata may be None before load_model_metadata activity."""
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = None
params = TrainModelParams.from_dict(d)
assert params.model_metadata is None
def test_coerce_experiment_run_id_rejects_bool():
"""Boolean must not be accepted as experiment_run_id."""
with pytest.raises(TypeError, match='experiment_run_id must be an integer'):
TrainModelParams._coerce_experiment_run_id(True)
def test_parse_optional_model_metadata_rejects_list():
"""model_metadata must be dict or None."""
with pytest.raises(TypeError, match='model_metadata must be a dict or None'):
TrainModelParams._parse_optional_model_metadata([])
def test_check_none_raises_value_error():
with pytest.raises(ValueError, match='test_field is required'):
TrainModelParams._check_none(None, str, 'test_field')
def test_check_none_raises_type_error():
with pytest.raises(TypeError, match='test_field must be of type str'):
TrainModelParams._check_none(123, str, 'test_field')
def test_validate_business_rules_success(valid_train_params_dict):
params = TrainModelParams.from_dict(valid_train_params_dict)
params.validate_business_rules()
def test_validate_business_rules_missing_model_metadata(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = None
params = TrainModelParams.from_dict(d)
with pytest.raises(ValueError, match='model_metadata is required'):
params.validate_business_rules()
def test_validate_business_rules_train_size_out_of_range(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['train_size'] = 5
params = TrainModelParams.from_dict(d)
with pytest.raises(ValueError, match='train_size must be between'):
params.validate_business_rules()
def test_validate_business_rules_empty_variable_columns(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['variable_columns'] = []
params = TrainModelParams.from_dict(d)
with pytest.raises(ValueError, match='variable_columns cannot be empty'):
params.validate_business_rules()
def test_validate_business_rules_empty_target(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['target_variable'] = ' '
params = TrainModelParams.from_dict(d)
with pytest.raises(ValueError, match='target_variable cannot be empty'):
params.validate_business_rules()
def test_validate_business_rules_whitespace_date_column(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['date_column'] = ' '
params = TrainModelParams.from_dict(d)
with pytest.raises(ValueError, match='date_column cannot be empty'):
params.validate_business_rules()
def test_from_dict_missing_required_key(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
del d['bucket_name']
with pytest.raises(ValueError, match='bucket_name is required'):
TrainModelParams.from_dict(d)
def test_to_dict_roundtrip_keys(valid_train_params_dict):
params = TrainModelParams.from_dict(valid_train_params_dict)
d = params.to_dict()
assert 'variable_columns' in d
assert d['experiment_run_id'] == 1
def test_coerce_experiment_run_id_float():
assert TrainModelParams._coerce_experiment_run_id(2.0) == 2
def test_coerce_experiment_run_id_none_raises():
with pytest.raises(ValueError, match='experiment_run_id is required'):
TrainModelParams._coerce_experiment_run_id(None)
def test_coerce_experiment_run_id_invalid_type():
with pytest.raises(TypeError, match='integer or numeric string'):
TrainModelParams._coerce_experiment_run_id([1])
def test_validate_model_param_schema_validation_error(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = {
'schemas': {
'components': {
'schemas': {
'data_model': {
'type': 'object',
'properties': {'x': {'type': 'integer'}},
'required': ['x'],
},
}
}
}
}
p = TrainModelParams.from_dict(d)
p.data_model_kwargs = {}
with pytest.raises(ValueError, match='Model parameters validation failed'):
p.validate_business_rules()
def test_validate_model_param_unexpected_validator_error(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = {
'schemas': {
'components': {
'schemas': {
'data_model': {'type': 'object'},
}
}
}
}
p = TrainModelParams.from_dict(d)
with patch('model_manager.utils.models.train_model_params.Draft202012Validator') as m:
m.return_value.validate.side_effect = RuntimeError('boom')
with pytest.raises(RuntimeError, match='boom'):
p.validate_business_rules()
def test_validate_business_rules_date_format_invalid(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['date_format'] = 'not-an-allowed-format'
p = TrainModelParams.from_dict(d)
with pytest.raises(ValueError, match='Invalid date_format'):
p.validate_business_rules()
def test_validate_required_strings_whitespace_bucket_file_model(valid_train_params_dict):
for field, msg in [
('bucket_name', 'bucket_name cannot be empty'),
('file_name', 'file_name cannot be empty'),
('model_name', 'model_name cannot be empty'),
]:
d = copy.deepcopy(valid_train_params_dict)
d[field] = ' '
p = TrainModelParams.from_dict(d)
with pytest.raises(ValueError, match=msg):
p.validate_business_rules()
def test_validate_model_param_only_data_model_schema(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = {
'schemas': {'components': {'schemas': {'data_model': {'type': 'object'}}}}
}
p = TrainModelParams.from_dict(d)
p.data_model_kwargs = {}
p.validate_business_rules()
def test_validate_model_param_only_model_schema(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = {'schemas': {'components': {'schemas': {'model': {'type': 'object'}}}}}
p = TrainModelParams.from_dict(d)
p.model_kwargs = {}
p.validate_business_rules()
def test_validate_model_param_only_opt_params_schema(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = {
'schemas': {'components': {'schemas': {'opt_params': {'type': 'object'}}}}
}
p = TrainModelParams.from_dict(d)
p.opt_params = {}
p.validate_business_rules()
def test_validate_model_param_all_schema_branches(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = {
'schemas': {
'components': {
'schemas': {
'data_model': {'type': 'object'},
'model': {'type': 'object'},
'opt_params': {'type': 'object'},
}
}
}
}
p = TrainModelParams.from_dict(d)
p.data_model_kwargs = {}
p.model_kwargs = {}
p.opt_params = {}
p.validate_business_rules()
def test_validate_frontend_date_format_whitespace_returns():
validate_frontend_date_format(' ')
def test_validate_frontend_date_format_valid_returns():
validate_frontend_date_format('dd/MM/yyyy HH:mm:ss')