diff --git a/model_manager/activities/training.py b/model_manager/activities/training.py index 0040c37..a7698c0 100644 --- a/model_manager/activities/training.py +++ b/model_manager/activities/training.py @@ -88,10 +88,12 @@ class Training(BaseActivity): try: self.info('Validating training parameters', metadata) - # Convert input_data directly to TrainModelParams (this validates all fields) - # The from_dict method will extract only the fields it needs + # Step 1: Convert input_data to TrainModelParams (validates types and required fields) train_params = TrainModelParams.from_dict(input_data) + # Step 2: Validate business rules (ranges, consistency, etc.) + train_params.validate_business_rules() + self.info( f'Training parameters validated successfully - ' f'Target: {train_params.target_variable}, ' diff --git a/model_manager/utils/models/train_model_params.py b/model_manager/utils/models/train_model_params.py index af1b1c7..e83e527 100644 --- a/model_manager/utils/models/train_model_params.py +++ b/model_manager/utils/models/train_model_params.py @@ -167,3 +167,71 @@ class TrainModelParams: raise TypeError(error) return value + + def validate_business_rules(self) -> None: + """ + Validate business rules and constraints for training parameters. + + This method performs additional validation beyond type checking to ensure + that parameter values are within acceptable ranges and logically consistent. + It implements defense-in-depth validation to catch configuration errors + early in the workflow. + + Raises: + ValueError: If any business rule is violated + + Example: + >>> params = TrainModelParams.from_dict(data) + >>> params.validate_business_rules() # Raises ValueError if invalid + """ + # Validate train_size range (1-99%) + if not 1 <= self.train_size <= 99: + raise ValueError(f'train_size must be between 1 and 99, got {self.train_size}') + + # Validate variable_columns is not empty + if not self.variable_columns: + raise ValueError('variable_columns cannot be empty') + + # Validate positive integers + if self.lag_train <= 0: + raise ValueError(f'lag_train must be positive, got {self.lag_train}') + + if self.lag_val <= 0: + raise ValueError(f'lag_val must be positive, got {self.lag_val}') + + if self.window <= 0: + raise ValueError(f'window must be positive, got {self.window}') + + # Validate low_lim and upp_lim consistency + if set(self.low_lim.keys()) != set(self.upp_lim.keys()): + raise ValueError( + f'low_lim and upp_lim must have the same keys. ' + f'low_lim keys: {set(self.low_lim.keys())}, ' + f'upp_lim keys: {set(self.upp_lim.keys())}' + ) + + # Validate that low_lim < upp_lim for each variable + for var in self.low_lim: + if self.low_lim[var] >= self.upp_lim[var]: + raise ValueError( + f'low_lim must be less than upp_lim for variable "{var}". ' + f'Got low_lim={self.low_lim[var]}, upp_lim={self.upp_lim[var]}' + ) + + # Validate target_variable is in variable_columns + if self.target_variable not in self.variable_columns: + raise ValueError( + f'target_variable "{self.target_variable}" must be in variable_columns: ' + f'{self.variable_columns}' + ) + + # Validate bucket_name and file_name are not empty + if not self.bucket_name.strip(): + raise ValueError('bucket_name cannot be empty or whitespace') + + if not self.file_name.strip(): + raise ValueError('file_name cannot be empty or whitespace') + + # Validate experiment_name is not empty + if not self.experiment_name.strip(): + raise ValueError('experiment_name cannot be empty or whitespace') diff --git a/tests/activities/test_training.py b/tests/activities/test_training.py index 3a2b582..d86bb59 100644 --- a/tests/activities/test_training.py +++ b/tests/activities/test_training.py @@ -335,7 +335,7 @@ async def test_validate_train_params_success(): 'metadata': {'workflow_id': 'test-123'}, 'experiment_run_id': 456, 'target_variable': 'price', - 'variable_columns': ['feature1', 'feature2'], + 'variable_columns': ['feature1', 'feature2', 'price'], # target_variable must be in list 'train_size': 80, 'shuffle': True, 'use_scaler': True, @@ -347,8 +347,8 @@ async def test_validate_train_params_success(): 'lag_train': 1, 'lag_val': 1, 'rem_static_win': False, - 'low_lim': {'feature1': 0.0, 'feature2': 0.0}, - 'upp_lim': {'feature1': 100.0, 'feature2': 100.0}, + 'low_lim': {'feature1': 0.0, 'feature2': 0.0, 'price': 0.0}, + 'upp_lim': {'feature1': 100.0, 'feature2': 100.0, 'price': 1000.0}, 'window': 10, 'experiment_name': 'test_experiment', 'removed_intervals': [], @@ -359,7 +359,7 @@ async def test_validate_train_params_success(): assert isinstance(result, TrainModelParams) assert result.experiment_run_id == 456 assert result.target_variable == 'price' - assert result.variable_columns == ['feature1', 'feature2'] + assert result.variable_columns == ['feature1', 'feature2', 'price'] assert result.train_size == 80 assert result.experiment_name == 'test_experiment' assert training.info.call_count == 2 @@ -472,7 +472,7 @@ async def test_validate_train_params_without_metadata(): input_data = { 'experiment_run_id': 789, 'target_variable': 'temperature', - 'variable_columns': ['sensor1'], + 'variable_columns': ['sensor1', 'temperature'], # target_variable must be in list 'train_size': 75, 'shuffle': False, 'use_scaler': True, @@ -484,8 +484,8 @@ async def test_validate_train_params_without_metadata(): 'lag_train': 2, 'lag_val': 2, 'rem_static_win': True, - 'low_lim': {'sensor1': -50.0}, - 'upp_lim': {'sensor1': 150.0}, + 'low_lim': {'sensor1': -50.0, 'temperature': -50.0}, + 'upp_lim': {'sensor1': 150.0, 'temperature': 150.0}, 'window': 20, 'experiment_name': 'sensor_experiment', 'removed_intervals': [], diff --git a/tests/utils/models/test_train_model_params.py b/tests/utils/models/test_train_model_params.py index f71a4af..53e0f9c 100644 --- a/tests/utils/models/test_train_model_params.py +++ b/tests/utils/models/test_train_model_params.py @@ -297,3 +297,201 @@ def test_train_model_params_check_type_method(): # Test that _check_type is a private method assert hasattr(params, '_check_type') assert callable(params._check_type) + + +# ============================================================================ +# Tests for validate_business_rules method +# ============================================================================ + + +def test_validate_business_rules_success(valid_params_dict): + """Test that valid params pass business rules validation.""" + # Ensure target_variable is in variable_columns + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + # Should not raise any exception + params.validate_business_rules() + + +def test_validate_business_rules_train_size_too_low(valid_params_dict): + """Test that train_size < 1 raises ValueError.""" + valid_params_dict['train_size'] = 0 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='train_size must be between 1 and 99'): + params.validate_business_rules() + + +def test_validate_business_rules_train_size_too_high(valid_params_dict): + """Test that train_size > 99 raises ValueError.""" + valid_params_dict['train_size'] = 100 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='train_size must be between 1 and 99'): + params.validate_business_rules() + + +def test_validate_business_rules_empty_variable_columns(valid_params_dict): + """Test that empty variable_columns raises ValueError.""" + valid_params_dict['variable_columns'] = [] + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='variable_columns cannot be empty'): + params.validate_business_rules() + + +def test_validate_business_rules_lag_train_zero(valid_params_dict): + """Test that lag_train = 0 raises ValueError.""" + valid_params_dict['lag_train'] = 0 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='lag_train must be positive'): + params.validate_business_rules() + + +def test_validate_business_rules_lag_train_negative(valid_params_dict): + """Test that lag_train < 0 raises ValueError.""" + valid_params_dict['lag_train'] = -1 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='lag_train must be positive'): + params.validate_business_rules() + + +def test_validate_business_rules_lag_val_zero(valid_params_dict): + """Test that lag_val = 0 raises ValueError.""" + valid_params_dict['lag_val'] = 0 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='lag_val must be positive'): + params.validate_business_rules() + + +def test_validate_business_rules_window_zero(valid_params_dict): + """Test that window = 0 raises ValueError.""" + valid_params_dict['window'] = 0 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='window must be positive'): + params.validate_business_rules() + + +def test_validate_business_rules_low_lim_upp_lim_keys_mismatch(valid_params_dict): + """Test that mismatched keys in low_lim and upp_lim raises ValueError.""" + valid_params_dict['low_lim'] = {'var1': 0.0, 'var2': 0.0} + valid_params_dict['upp_lim'] = {'var1': 100.0, 'var3': 100.0} # var3 instead of var2 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='low_lim and upp_lim must have the same keys'): + params.validate_business_rules() + + +def test_validate_business_rules_low_lim_greater_than_upp_lim(valid_params_dict): + """Test that low_lim >= upp_lim raises ValueError.""" + valid_params_dict['low_lim'] = {'var1': 100.0, 'var2': 0.0, 'var3': 0.0} + valid_params_dict['upp_lim'] = {'var1': 50.0, 'var2': 100.0, 'var3': 100.0} + valid_params_dict['target_variable'] = 'var2' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='low_lim must be less than upp_lim for variable "var1"'): + params.validate_business_rules() + + +def test_validate_business_rules_low_lim_equal_to_upp_lim(valid_params_dict): + """Test that low_lim == upp_lim raises ValueError.""" + valid_params_dict['low_lim'] = {'var1': 50.0, 'var2': 0.0, 'var3': 0.0} + valid_params_dict['upp_lim'] = {'var1': 50.0, 'var2': 100.0, 'var3': 100.0} + valid_params_dict['target_variable'] = 'var2' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='low_lim must be less than upp_lim for variable "var1"'): + params.validate_business_rules() + + +def test_validate_business_rules_target_not_in_variable_columns(valid_params_dict): + """Test that target_variable not in variable_columns raises ValueError.""" + valid_params_dict['target_variable'] = 'nonexistent_var' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises( + ValueError, match='target_variable "nonexistent_var" must be in variable_columns' + ): + params.validate_business_rules() + + +def test_validate_business_rules_empty_bucket_name(valid_params_dict): + """Test that empty bucket_name raises ValueError.""" + valid_params_dict['bucket_name'] = ' ' + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='bucket_name cannot be empty or whitespace'): + params.validate_business_rules() + + +def test_validate_business_rules_empty_file_name(valid_params_dict): + """Test that empty file_name raises ValueError.""" + valid_params_dict['file_name'] = '' + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='file_name cannot be empty or whitespace'): + params.validate_business_rules() + + +def test_validate_business_rules_empty_experiment_name(valid_params_dict): + """Test that empty experiment_name raises ValueError.""" + valid_params_dict['experiment_name'] = ' ' + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + with pytest.raises(ValueError, match='experiment_name cannot be empty or whitespace'): + params.validate_business_rules() + + +def test_validate_business_rules_all_valid_edge_cases(valid_params_dict): + """Test that edge case valid values pass validation.""" + valid_params_dict['train_size'] = 1 # Minimum valid + valid_params_dict['lag_train'] = 1 # Minimum valid + valid_params_dict['lag_val'] = 1 # Minimum valid + valid_params_dict['window'] = 1 # Minimum valid + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + # Should not raise any exception + params.validate_business_rules() + + +def test_validate_business_rules_train_size_99(valid_params_dict): + """Test that train_size = 99 (maximum valid) passes validation.""" + valid_params_dict['train_size'] = 99 + valid_params_dict['target_variable'] = 'var1' + + params = TrainModelParams.from_dict(valid_params_dict) + + # Should not raise any exception + params.validate_business_rules()