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.

This commit is contained in:
Bruno Domingues
2025-10-15 16:16:09 -03:00
parent 4b28d5f1c8
commit 305a39b44c
4 changed files with 277 additions and 9 deletions

View File

@@ -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': [],