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

@@ -1,6 +1,7 @@
from unittest.mock import ANY, MagicMock, patch
import numpy as np
import pytest
from pytest import fixture, mark
from sientia_do.notifications.models import NotificationLevel
from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ
@@ -333,10 +334,9 @@ async def test_save_model_success(mlflow):
mlflow.model_monitoring_repository.generate_artifacts.assert_called_once_with(train_result)
mlflow.model_monitoring_repository.save_run.assert_called_once_with(train_result)
# Verify response
assert response['success'] is True
assert response['result'] == train_result
assert response['error_message'] is None
# Verify response - now returns TrainModelResult directly
assert response == train_result
assert response.run_name == 'test_experiment-1'
assert train_result.run_name == 'test_experiment-1'
@@ -362,13 +362,9 @@ async def test_save_model_get_next_run_name_error(mlflow):
'train_result': train_result,
}
# Call the method
response = await mlflow.save_model(input_data)
# Verify error handling
assert response['success'] is False
assert response['result'] is None
assert 'MLflow connection error' in response['error_message']
# Call the method - should raise exception
with pytest.raises(Exception, match='MLflow connection error'):
await mlflow.save_model(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once_with(
@@ -404,13 +400,9 @@ async def test_save_model_generate_artifacts_error(mlflow):
'train_result': train_result,
}
# Call the method
response = await mlflow.save_model(input_data)
# Verify error handling
assert response['success'] is False
assert response['result'] is None
assert 'Reports directory does not exist' in response['error_message']
# Call the method - should raise exception
with pytest.raises(FileNotFoundError, match='Reports directory does not exist'):
await mlflow.save_model(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once_with(
@@ -447,13 +439,9 @@ async def test_save_model_save_run_error(mlflow):
'train_result': train_result,
}
# Call the method
response = await mlflow.save_model(input_data)
# Verify error handling
assert response['success'] is False
assert response['result'] is None
assert 'One or more metrics (MSE, R2, MAE) are None' in response['error_message']
# Call the method - should raise exception
with pytest.raises(ValueError, match=r'One or more metrics \(MSE, R2, MAE\) are None'):
await mlflow.save_model(input_data)
# Verify notification was sent
mlflow.send_notification.assert_called_once_with(
@@ -491,10 +479,9 @@ async def test_save_model_missing_metadata(mlflow):
# Call the method
response = await mlflow.save_model(input_data)
# Verify it still works (metadata defaults to {})
assert response['success'] is True
assert response['result'] == train_result
assert response['error_message'] is None
# Verify it still works (metadata defaults to {}) - returns TrainModelResult directly
assert response == train_result
assert response.run_name == 'test_experiment-1'
@mark.asyncio
@@ -540,10 +527,8 @@ async def test_save_model_complete_flow(mlflow):
mlflow.model_monitoring_repository.generate_artifacts.assert_called_once()
mlflow.model_monitoring_repository.save_run.assert_called_once_with(updated_result)
# Verify response
assert response['success'] is True
assert response['result'] == updated_result
assert response['error_message'] is None
assert updated_result.run_name == 'production_model-5'
assert updated_result.run_dir is not None
assert updated_result.report_path is not None
# Verify response - returns TrainModelResult directly
assert response == updated_result
assert response.run_name == 'production_model-5'
assert response.run_dir is not None
assert response.report_path is not None

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'