SIENTIAPDE-1253: Refactor training workflow and activities to raise exceptions on failure

This commit refactors the training workflow and associated activities to raise exceptions on failure instead of returning success/failure dictionaries. This allows the Temporal workflow to handle errors more effectively and ensures that the workflow stops when a critical error occurs.

Key changes:

- The train_model workflow is introduced to orchestrate the entire training process, including parameter validation, data download, model training, and model saving.
- The validate_train_params activity is added to validate and convert training parameters.
- The train_model and save_model activities are updated to raise exceptions on failure.
- The ExperimentStatus enum is updated to include a new status for orchestrator validation errors.
- The tests are updated to reflect the new exception-based error handling.
- The activities now return the TrainModelResult directly instead of a dictionary.
This commit is contained in:
Bruno Domingues
2025-10-15 15:19:29 -03:00
parent 61267ec49d
commit 8ea98360c3
8 changed files with 1472 additions and 129 deletions

View File

@@ -3,6 +3,7 @@
from io import BytesIO
from unittest.mock import MagicMock, patch
import pytest
from pytest import mark
from model_manager.activities.training import Training
@@ -74,10 +75,11 @@ async def test_train_model_success(mock_training_repository_class):
# Execute
result = await training.train_model(input_data)
# Assertions
assert result['success'] is True
assert result['result'] == mock_final_result
assert result['error_message'] is None
# Assertions - now returns TrainModelResult directly
assert result == mock_final_result
assert result.mse_val == 0.5
assert result.mae_val == 0.3
assert result.r2_val == 0.95
# Verify repository calls
mock_repository.train.assert_called_once()
@@ -124,11 +126,10 @@ async def test_train_model_invalid_file_type(mock_training_repository_class):
'train_params': train_params,
}
result = await training.train_model(input_data)
# Should raise ValueError
with pytest.raises(ValueError, match='uploaded_file must be BytesIO'):
await training.train_model(input_data)
assert result['success'] is False
assert result['result'] is None
assert 'uploaded_file must be BytesIO' in result['error_message']
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@@ -176,11 +177,10 @@ async def test_train_model_training_error(mock_training_repository_class):
'train_params': train_params,
}
result = await training.train_model(input_data)
# Should raise ValueError
with pytest.raises(ValueError, match='Training data is empty'):
await training.train_model(input_data)
assert result['success'] is False
assert result['result'] is None
assert 'Training data is empty' in result['error_message']
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@@ -227,16 +227,13 @@ async def test_train_model_sends_notification_on_error(mock_training_repository_
'train_params': train_params,
}
result = await training.train_model(input_data)
# Should raise Exception
with pytest.raises(Exception, match='Database connection failed'):
await training.train_model(input_data)
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
# Verify result - the important part is that error was caught and returned
assert result['success'] is False
assert result['result'] is None
assert 'Database connection failed' in result['error_message']
@mark.asyncio
@patch('model_manager.activities.training.TrainingRepository')
@@ -283,11 +280,10 @@ async def test_train_model_after_calculation_error(mock_training_repository_clas
'train_params': train_params,
}
result = await training.train_model(input_data)
# Should raise Exception
with pytest.raises(Exception, match='Metric calculation failed'):
await training.train_model(input_data)
assert result['success'] is False
assert result['result'] is None
assert 'Metric calculation failed' in result['error_message']
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@@ -316,11 +312,188 @@ async def test_train_model_invalid_train_params_type(mock_training_repository_cl
}, # This is a dict, not TrainModelParams
}
result = await training.train_model(input_data)
# Should raise ValueError
with pytest.raises(ValueError, match='train_params must be TrainModelParams.*dict'):
await training.train_model(input_data)
# ============================================================================
# Tests for validate_train_params
# ============================================================================
@mark.asyncio
async def test_validate_train_params_success():
"""Test successful validation of training parameters."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
'experiment_run_id': 456,
'target_variable': 'price',
'variable_columns': ['feature1', 'feature2'],
'train_size': 80,
'shuffle': True,
'use_scaler': True,
'include_ar': False,
'bucket_name': 'test-bucket',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'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},
'window': 10,
'experiment_name': 'test_experiment',
'removed_intervals': [],
}
result = await training.validate_train_params(input_data)
assert isinstance(result, TrainModelParams)
assert result.experiment_run_id == 456
assert result.target_variable == 'price'
assert result.variable_columns == ['feature1', 'feature2']
assert result.train_size == 80
assert result.experiment_name == 'test_experiment'
assert training.info.call_count == 2
@mark.asyncio
async def test_validate_train_params_missing_required_field():
"""Test validation fails when required field is missing."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
training.error = MagicMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
'experiment_run_id': 456,
'variable_columns': ['feature1'],
'train_size': 80,
'shuffle': True,
'use_scaler': False,
'include_ar': False,
'bucket_name': 'test',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 1,
'lag_val': 1,
'rem_static_win': False,
'low_lim': {'feature1': 0.0},
'upp_lim': {'feature1': 100.0},
'window': 10,
'experiment_name': 'test',
'removed_intervals': [],
}
with pytest.raises(ValueError, match='target_variable'):
await training.validate_train_params(input_data)
assert result['success'] is False
assert result['result'] is None
assert 'train_params must be TrainModelParams' in result['error_message']
assert 'dict' in result['error_message']
# Verify notification was sent (via BaseActivity)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
async def test_validate_train_params_invalid_type():
"""Test validation fails when field has invalid type."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
training.error = MagicMock()
input_data = {
'metadata': {'workflow_id': 'test-invalid'},
'experiment_run_id': 456,
'target_variable': 'price',
'variable_columns': ['feature1'],
'train_size': 'invalid',
'shuffle': True,
'use_scaler': False,
'include_ar': False,
'bucket_name': 'test',
'file_name': 'test.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 1,
'lag_val': 1,
'rem_static_win': False,
'low_lim': {'feature1': 0.0},
'upp_lim': {'feature1': 100.0},
'window': 10,
'experiment_name': 'test',
'removed_intervals': [],
}
with pytest.raises((ValueError, TypeError)):
await training.validate_train_params(input_data)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
async def test_validate_train_params_empty_input():
"""Test validation fails with empty input."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
training.error = MagicMock()
input_data = {'metadata': {}}
with pytest.raises(ValueError):
await training.validate_train_params(input_data)
notification_handler.send_notification.assert_called_once()
@mark.asyncio
async def test_validate_train_params_without_metadata():
"""Test validation works even without metadata key."""
logger = MagicMock()
notification_handler = MagicMock()
training = Training(logger=logger, notification_handler=notification_handler)
training.info = MagicMock()
input_data = {
'experiment_run_id': 789,
'target_variable': 'temperature',
'variable_columns': ['sensor1'],
'train_size': 75,
'shuffle': False,
'use_scaler': True,
'include_ar': True,
'bucket_name': 'sensors',
'file_name': 'data.csv',
'line_separator': '\n',
'decimal_separator': '.',
'lag_train': 2,
'lag_val': 2,
'rem_static_win': True,
'low_lim': {'sensor1': -50.0},
'upp_lim': {'sensor1': 150.0},
'window': 20,
'experiment_name': 'sensor_experiment',
'removed_intervals': [],
}
result = await training.validate_train_params(input_data)
assert isinstance(result, TrainModelParams)
assert result.experiment_run_id == 789
assert result.target_variable == 'temperature'
assert result.experiment_name == 'sensor_experiment'