SIENTIAPDE-1241: Train_model tests
This commit is contained in:
647
tests/workflows/test_train_model.py
Normal file
647
tests/workflows/test_train_model.py
Normal file
@@ -0,0 +1,647 @@
|
||||
"""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'] == ''
|
||||
|
||||
Reference in New Issue
Block a user