From 305a39b44c40260424fd2bf6c388426d7f79b868 Mon Sep 17 00:00:00 2001 From: Bruno Domingues Date: Wed, 15 Oct 2025 16:16:09 -0300 Subject: [PATCH] SIENTIAPDE-1253: Implement business rule validation for training parameters. Adds a validate_business_rules method to the TrainModelParams class to enforce constraints on training parameters, improving data integrity and preventing errors. Also updates tests to include target variable in variable columns and adds tests for business rule validations. --- model_manager/activities/training.py | 6 +- .../utils/models/train_model_params.py | 68 ++++++ tests/activities/test_training.py | 14 +- tests/utils/models/test_train_model_params.py | 198 ++++++++++++++++++ 4 files changed, 277 insertions(+), 9 deletions(-) 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()