Files
sientia-dataops-model-manager/tests/workflows/test_train_model.py
2025-10-30 10:11:07 -03:00

648 lines
22 KiB
Python

"""Unit tests for TrainModel workflow."""
import pytest
from unittest.mock import AsyncMock, MagicMock, Mock, patch
from model_manager.utils.exceptions import ModelTrainingError
from model_manager.utils.models.experiment_status import ExperimentStatus
from model_manager.utils.models.train_model_params import TrainModelParams
@pytest.fixture
def mock_train_params():
"""Create a mock TrainModelParams object."""
params = Mock(spec=TrainModelParams)
params.experiment_run_id = 123
params.bucket_name = 'test-bucket'
params.file_name = 'test-file.csv'
params.target_variable = 'target'
params.variable_columns = ['var1', 'var2']
return params
@pytest.fixture
def sample_input_data():
"""Create sample input data for workflow."""
return {
'experiment_run_id': 123,
'experiment_name': 'test_experiment',
'target_variable': 'target',
'variable_columns': ['var1', 'var2'],
'train_size': 80,
'bucket_name': 'test-bucket',
'file_name': 'test-file.csv',
'lag_train': 0,
'lag_val': 0,
'rem_static_win': False,
'low_lim': {},
'upp_lim': {},
'window': 0,
'use_scaler': False,
'include_ar': False,
'shuffle': True,
'line_separator': ',',
'decimal_separator': '.',
'removed_intervals': [],
}
@pytest.fixture
def mock_workflow():
"""Create a mock workflow module."""
workflow_mock = Mock()
workflow_mock.execute_activity_method = AsyncMock()
return workflow_mock
def test_validate_experiment_run_id_success():
"""Test successful experiment_run_id validation."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
input_data = {'experiment_run_id': 123}
result = workflow_instance._validate_experiment_run_id(input_data)
assert result == 123
def test_validate_experiment_run_id_missing():
"""Test validation fails when experiment_run_id is missing."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
input_data = {}
with pytest.raises(ValueError, match='experiment_run_id is required but was not provided'):
workflow_instance._validate_experiment_run_id(input_data)
def test_validate_experiment_run_id_not_integer():
"""Test validation fails when experiment_run_id is not an integer."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
input_data = {'experiment_run_id': 'not_an_int'}
with pytest.raises(ValueError, match='experiment_run_id must be an integer, got str'):
workflow_instance._validate_experiment_run_id(input_data)
def test_validate_experiment_run_id_none():
"""Test validation fails when experiment_run_id is None."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
input_data = {'experiment_run_id': None}
with pytest.raises(ValueError, match='experiment_run_id is required but was not provided'):
workflow_instance._validate_experiment_run_id(input_data)
def test_extract_error_message_simple():
"""Test extracting error message from simple exception."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
exc = ValueError('Test error message')
result = workflow_instance._extract_error_message(exc)
assert result == 'Test error message'
def test_extract_error_message_with_cause():
"""Test extracting error message from exception with cause."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
# Create exception chain
cause = ValueError('Root cause')
exc = RuntimeError('Outer error')
exc.cause = cause
result = workflow_instance._extract_error_message(exc)
assert 'Outer error' in result
assert 'Root cause' in result
def test_extract_error_message_empty():
"""Test extracting error message from exception with empty string."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
exc = ValueError('')
result = workflow_instance._extract_error_message(exc)
# Should return repr when no message
assert 'ValueError' in result
def test_extract_error_message_circular_reference():
"""Test extracting error message handles circular references."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
# Create circular reference
exc1 = ValueError('Error 1')
exc2 = ValueError('Error 2')
exc1.cause = exc2
exc2.cause = exc1 # Circular!
result = workflow_instance._extract_error_message(exc1)
# Should handle circular reference without infinite loop
assert 'Error 1' in result
assert 'Error 2' in result
def test_extract_error_message_duplicate_messages():
"""Test that duplicate error messages are not repeated."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
# Create chain with duplicate messages
exc1 = ValueError('Same error')
exc2 = ValueError('Same error')
exc1.cause = exc2
result = workflow_instance._extract_error_message(exc1)
# Should only appear once
assert result.count('Same error') == 1
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_validate_training_parameters_success(mock_workflow_module, sample_input_data, mock_train_params):
"""Test successful parameter validation."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks
mock_workflow_module.execute_activity_method = AsyncMock(return_value=mock_train_params)
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
result = await workflow_instance._validate_training_parameters(
sample_input_data, 123, metadata
)
assert result == mock_train_params
# Verify validate_train_params activity was called
assert mock_workflow_module.execute_activity_method.call_count == 2 # validate + update status
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_validate_training_parameters_failure(mock_workflow_module, sample_input_data):
"""Test parameter validation handles errors."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - first call fails, second succeeds (update status)
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[ValueError('Invalid params'), None]
)
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
with pytest.raises(ValueError, match='Invalid params'):
await workflow_instance._validate_training_parameters(
sample_input_data, 123, metadata
)
# Verify error status update was called
assert mock_workflow_module.execute_activity_method.call_count == 2
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_train_model_success(mock_workflow_module, mock_train_params):
"""Test successful model training."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks
train_result = {
'run_name': 'test-run-123',
'run_dir': '/tmp/test-run',
'mse_val': 0.5,
'r2_val': 0.9,
}
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[train_result, None] # train + update status
)
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
result = await workflow_instance._train_model(mock_train_params, 123, metadata)
assert result == train_result
assert mock_workflow_module.execute_activity_method.call_count == 2
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_train_model_training_error(mock_workflow_module, mock_train_params):
"""Test model training handles training errors."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - training fails
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[RuntimeError('Training failed'), None]
)
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
with pytest.raises(RuntimeError, match='Training failed'):
await workflow_instance._train_model(mock_train_params, 123, metadata)
# Verify error status update was called
assert mock_workflow_module.execute_activity_method.call_count == 2
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_train_model_mlflow_error(mock_workflow_module, mock_train_params):
"""Test model training handles MLflow save errors."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - MLflow save fails
mlflow_error = ModelTrainingError(model_trained=True, model_saved=False, message='MLflow save failed')
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[mlflow_error, None]
)
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
with pytest.raises(ModelTrainingError):
await workflow_instance._train_model(mock_train_params, 123, metadata)
# Verify MLFLOW_SEND_ERROR status was set
call_args = mock_workflow_module.execute_activity_method.call_args_list[1]
assert call_args[0][1]['status'] == ExperimentStatus.MLFLOW_SEND_ERROR
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_cleanup_resources_success(mock_workflow_module):
"""Test successful resource cleanup."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[None, None] # cleanup + update status
)
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
await workflow_instance._cleanup_resources(
experiment_run_id=123,
run_dir='/tmp/test-run',
bucket_name='test-bucket',
file_name='test-file.csv',
metadata=metadata,
)
assert mock_workflow_module.execute_activity_method.call_count == 2
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_cleanup_resources_failure(mock_workflow_module):
"""Test resource cleanup handles errors."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - cleanup fails
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[RuntimeError('Cleanup failed'), None]
)
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod', 'experiment_run_id': 123}}
with pytest.raises(RuntimeError, match='Cleanup failed'):
await workflow_instance._cleanup_resources(
experiment_run_id=123,
run_dir='/tmp/test-run',
bucket_name='test-bucket',
file_name='test-file.csv',
metadata=metadata,
)
# Verify error status update was called
assert mock_workflow_module.execute_activity_method.call_count == 2
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_update_experiment_run_status_only(mock_workflow_module):
"""Test update_experiment_run with status only."""
from model_manager.workflows.train_model import TrainModel
from model_manager.activities.experiment_tracking import UpdateType
mock_workflow_module.execute_activity_method = AsyncMock()
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod'}}
await workflow_instance._update_experiment_run(
metadata=metadata,
experiment_run_id=123,
update_type=UpdateType.STATUS,
status=ExperimentStatus.MAGE_WAITING_PROC,
)
# Verify activity was called with correct parameters
call_args = mock_workflow_module.execute_activity_method.call_args[0]
assert call_args[1]['experiment_run_id'] == 123
assert call_args[1]['update_type'] == UpdateType.STATUS
assert call_args[1]['status'] == ExperimentStatus.MAGE_WAITING_PROC
assert 'error_message' not in call_args[1] or call_args[1].get('error_message') is None
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_update_experiment_run_with_error(mock_workflow_module):
"""Test update_experiment_run with error message."""
from model_manager.workflows.train_model import TrainModel
from model_manager.activities.experiment_tracking import UpdateType
mock_workflow_module.execute_activity_method = AsyncMock()
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod'}}
await workflow_instance._update_experiment_run(
metadata=metadata,
experiment_run_id=123,
update_type=UpdateType.STATUS_WITH_ERROR,
status=ExperimentStatus.TRAINING_ERROR,
error_message='Test error',
)
# Verify activity was called with error message
call_args = mock_workflow_module.execute_activity_method.call_args[0]
assert call_args[1]['error_message'] == 'Test error'
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_update_experiment_run_with_run_name(mock_workflow_module):
"""Test update_experiment_run with run_name."""
from model_manager.workflows.train_model import TrainModel
from model_manager.activities.experiment_tracking import UpdateType
mock_workflow_module.execute_activity_method = AsyncMock()
workflow_instance = TrainModel()
metadata = {'metadata': {'pod_id': 'test-pod'}}
await workflow_instance._update_experiment_run(
metadata=metadata,
experiment_run_id=123,
update_type=UpdateType.MODEL_SAVED,
status=ExperimentStatus.MLFLOW_SENT,
run_name='test-run-123',
)
# Verify activity was called with run_name
call_args = mock_workflow_module.execute_activity_method.call_args[0]
assert call_args[1]['run_name'] == 'test-run-123'
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
@patch('model_manager.workflows.train_model.POD_ID', 'test-pod-456')
async def test_run_complete_workflow_success(mock_workflow_module, sample_input_data, mock_train_params):
"""Test complete workflow execution success path."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks for all activities
train_result = {
'run_name': 'test-run-123',
'run_dir': '/tmp/test-run',
}
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate_train_params
None, # update status (MAGE_WAITING_PROC)
train_result, # train_model
None, # update status (MLFLOW_SENT)
None, # cleanup_resources
None, # update status (FILE_DELETED)
]
)
workflow_instance = TrainModel()
# Should not raise any exceptions
await workflow_instance.run(sample_input_data)
# Verify all activities were called
assert mock_workflow_module.execute_activity_method.call_count == 6
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_run_workflow_validation_error(mock_workflow_module, sample_input_data):
"""Test workflow handles validation errors."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - validation fails
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[
ValueError('Invalid parameters'), # validate_train_params fails
None, # update status (ORCHESTRATOR_VALIDATION_ERROR)
]
)
workflow_instance = TrainModel()
with pytest.raises(ValueError, match='Invalid parameters'):
await workflow_instance.run(sample_input_data)
# Verify status update was called
assert mock_workflow_module.execute_activity_method.call_count == 2
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_run_workflow_training_error(mock_workflow_module, sample_input_data, mock_train_params):
"""Test workflow handles training errors."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - training fails
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate_train_params
None, # update status (MAGE_WAITING_PROC)
RuntimeError('Training failed'), # train_model fails
None, # update status (TRAINING_ERROR)
]
)
workflow_instance = TrainModel()
with pytest.raises(RuntimeError, match='Training failed'):
await workflow_instance.run(sample_input_data)
assert mock_workflow_module.execute_activity_method.call_count == 4
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_run_workflow_cleanup_error(mock_workflow_module, sample_input_data, mock_train_params):
"""Test workflow handles cleanup errors."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - cleanup fails
train_result = {
'run_name': 'test-run-123',
'run_dir': '/tmp/test-run',
}
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate_train_params
None, # update status (MAGE_WAITING_PROC)
train_result, # train_model
None, # update status (MLFLOW_SENT)
RuntimeError('Cleanup failed'), # cleanup_resources fails
None, # update status (FILE_DELETE_ERROR)
]
)
workflow_instance = TrainModel()
with pytest.raises(RuntimeError, match='Cleanup failed'):
await workflow_instance.run(sample_input_data)
assert mock_workflow_module.execute_activity_method.call_count == 6
@pytest.mark.asyncio
async def test_run_workflow_missing_experiment_run_id():
"""Test workflow fails early when experiment_run_id is missing."""
from model_manager.workflows.train_model import TrainModel
workflow_instance = TrainModel()
input_data = {} # Missing experiment_run_id
with pytest.raises(ValueError, match='experiment_run_id is required'):
await workflow_instance.run(input_data)
def test_module_constants():
"""Test that module-level constants are defined correctly."""
from model_manager.workflows.train_model import (
TIMEOUT_VALIDATE_PARAMS,
TIMEOUT_TRAIN_MODEL,
TIMEOUT_DELETE_FILE,
TIMEOUT_UPDATE_DATABASE,
network_retry_policy,
no_retry_policy,
database_retry_policy,
)
# Verify timeouts are integers
assert isinstance(TIMEOUT_VALIDATE_PARAMS, int)
assert isinstance(TIMEOUT_TRAIN_MODEL, int)
assert isinstance(TIMEOUT_DELETE_FILE, int)
assert isinstance(TIMEOUT_UPDATE_DATABASE, int)
# Verify default values
assert TIMEOUT_VALIDATE_PARAMS == 30
assert TIMEOUT_TRAIN_MODEL == 2700
assert TIMEOUT_DELETE_FILE == 120
assert TIMEOUT_UPDATE_DATABASE == 30
# Verify retry policies exist
assert network_retry_policy is not None
assert no_retry_policy is not None
assert database_retry_policy is not None
# Verify retry policy configurations
assert network_retry_policy.maximum_attempts == 5
assert no_retry_policy.maximum_attempts == 1
assert database_retry_policy.maximum_attempts == 5
def test_workflow_class_definition():
"""Test that TrainModel workflow class is properly defined."""
from model_manager.workflows.train_model import TrainModel
# Verify class exists and has required methods
assert hasattr(TrainModel, 'run')
assert hasattr(TrainModel, '_validate_experiment_run_id')
assert hasattr(TrainModel, '_validate_training_parameters')
assert hasattr(TrainModel, '_train_model')
assert hasattr(TrainModel, '_cleanup_resources')
assert hasattr(TrainModel, '_update_experiment_run')
assert hasattr(TrainModel, '_extract_error_message')
@pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow')
async def test_train_model_empty_run_dir(mock_workflow_module, mock_train_params):
"""Test cleanup handles empty run_dir gracefully."""
from model_manager.workflows.train_model import TrainModel
# Setup mocks - train returns empty run_dir
train_result = {
'run_name': 'test-run-123',
'run_dir': None, # Empty run_dir
}
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate
None, # update status
train_result, # train
None, # update status
None, # cleanup
None, # update status
]
)
workflow_instance = TrainModel()
sample_input = {
'experiment_run_id': 123,
'target_variable': 'target',
'variable_columns': ['var1'],
'train_size': 80,
'bucket_name': 'test-bucket',
'file_name': 'test.csv',
'lag_train': 0,
'lag_val': 0,
'rem_static_win': False,
'low_lim': {},
'upp_lim': {},
'window': 0,
'use_scaler': False,
'include_ar': False,
'shuffle': True,
'line_separator': ',',
'decimal_separator': '.',
'removed_intervals': [],
}
# Should handle None run_dir gracefully
await workflow_instance.run(sample_input)
# Verify cleanup was called with empty string
cleanup_call = mock_workflow_module.execute_activity_method.call_args_list[4]
assert cleanup_call[0][1]['run_dir'] == ''