"""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 name', '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 name' 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 19 required fields.""" import inspect sig = inspect.signature(TrainModelParams.__init__) # Subtract 1 for 'self' param_count = len(sig.parameters) - 1 assert param_count == 19 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='name', 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', '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', 'removed_intervals': [], } params = TrainModelParams(**params_dict) # Test that _check_type is a private method assert hasattr(params, '_check_type') assert callable(params._check_type)