Code import - branch release/SIENTIAPDE-1645
This commit is contained in:
318
tests/workflows/test_train_model.py
Normal file
318
tests/workflows/test_train_model.py
Normal file
@@ -0,0 +1,318 @@
|
||||
"""Unit tests for TrainModel workflow."""
|
||||
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from temporalio.exceptions import ApplicationError
|
||||
|
||||
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():
|
||||
"""Minimal mock TrainModelParams."""
|
||||
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():
|
||||
"""Sample workflow input (IDs normalized in run())."""
|
||||
return {
|
||||
'experiment_run_id': 123,
|
||||
'target_variable': 'target',
|
||||
'variable_columns': ['var1', 'var2'],
|
||||
'train_size': 80,
|
||||
'bucket_name': 'test-bucket',
|
||||
'file_name': 'test-file.csv',
|
||||
'line_separator': ',',
|
||||
'decimal_separator': '.',
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'model_name': 'Linear Regression',
|
||||
'model_type': 'linear_regression',
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'val_file_name': None,
|
||||
'model_id': None,
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_success():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
wf = TrainModel()
|
||||
assert wf._validate_experiment_run_id({'experiment_run_id': 123}) == 123
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_string_numeric():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
wf = TrainModel()
|
||||
assert wf._validate_experiment_run_id({'experiment_run_id': '123'}) == 123
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_missing():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
with pytest.raises(ValueError, match='experiment_run_id is required'):
|
||||
TrainModel()._validate_experiment_run_id({})
|
||||
|
||||
|
||||
def test_validate_experiment_run_id_invalid_type():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
with pytest.raises(ValueError, match='must be an integer or numeric string'):
|
||||
TrainModel()._validate_experiment_run_id({'experiment_run_id': 'not_int'})
|
||||
|
||||
|
||||
def test_extract_error_message_simple():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
assert TrainModel()._extract_error_message(ValueError('x')) == 'x'
|
||||
|
||||
|
||||
def test_extract_error_message_with_cause():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
cause = ValueError('Root')
|
||||
exc = RuntimeError('Outer')
|
||||
exc.__cause__ = cause
|
||||
msg = TrainModel()._extract_error_message(exc)
|
||||
assert 'Outer' in msg and 'Root' in msg
|
||||
|
||||
|
||||
def test_extract_error_message_empty_message():
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
out = TrainModel()._extract_error_message(ValueError(''))
|
||||
assert 'ValueError' in out
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_validate_training_parameters_success(mock_wf, sample_input_data, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'experiment_run_id': 123, 'model_metadata': {}},
|
||||
mock_train_params,
|
||||
None,
|
||||
]
|
||||
)
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
out = await TrainModel()._validate_training_parameters(sample_input_data, 123, meta)
|
||||
assert out is mock_train_params
|
||||
assert mock_wf.execute_activity_method.call_count == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_validate_training_parameters_load_fails(mock_wf, sample_input_data):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[ValueError('load'), None])
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(ValueError, match='load'):
|
||||
await TrainModel()._validate_training_parameters(sample_input_data, 123, meta)
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_train_model_success(mock_wf, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
tr = {'run_name': 'rn', 'run_id': 'rid', 'run_dir': '/tmp/r'}
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[tr, None])
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
out = await TrainModel()._train_model(mock_train_params, 123, meta)
|
||||
assert out == tr
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_validate_training_parameters_logs_when_db_update_fails(mock_wf, sample_input_data):
|
||||
"""If persisting ORCHESTRATOR_VALIDATION_ERROR fails, workflow logs a warning."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[ValueError('validation'), RuntimeError('db')],
|
||||
)
|
||||
mock_wf.logger = Mock()
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(ValueError, match='validation'):
|
||||
await TrainModel()._validate_training_parameters(sample_input_data, 123, meta)
|
||||
mock_wf.logger.warning.assert_called_once()
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_train_model_logs_when_error_status_persist_fails(mock_wf, mock_train_params):
|
||||
"""If persisting TRAINING_ERROR fails, workflow logs a warning."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[RuntimeError('train'), RuntimeError('db')],
|
||||
)
|
||||
mock_wf.logger = Mock()
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(RuntimeError, match='train'):
|
||||
await TrainModel()._train_model(mock_train_params, 123, meta)
|
||||
mock_wf.logger.warning.assert_called_once()
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_train_model_failure_updates_db(mock_wf, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[RuntimeError('fail'), None])
|
||||
meta = {'metadata': {'pod_id': 'p', 'experiment_run_id': 123}}
|
||||
with pytest.raises(RuntimeError, match='fail'):
|
||||
await TrainModel()._train_model(mock_train_params, 123, meta)
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
err_call = mock_wf.execute_activity_method.call_args_list[1]
|
||||
assert err_call[0][1]['status'] == ExperimentStatus.TRAINING_ERROR
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_cleanup_resources(mock_wf):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(return_value=None)
|
||||
meta = {'metadata': {'pod_id': 'p'}}
|
||||
await TrainModel()._cleanup_resources('/tmp/x', meta)
|
||||
assert mock_wf.execute_activity_method.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_cleanup_resources_none_skips(mock_wf):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
await TrainModel()._cleanup_resources(None, {'metadata': {}})
|
||||
mock_wf.execute_activity_method.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_success_six_activities(mock_wf, sample_input_data, mock_train_params):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
tr = {'run_name': 'rn', 'run_id': 'i', 'run_dir': '/tmp/t'}
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'x': 1},
|
||||
mock_train_params,
|
||||
None,
|
||||
tr,
|
||||
None,
|
||||
None,
|
||||
]
|
||||
)
|
||||
result = await TrainModel().run(sample_input_data)
|
||||
assert result == tr
|
||||
assert mock_wf.execute_activity_method.call_count == 6
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_validation_error(mock_wf, sample_input_data):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(side_effect=[ValueError('bad'), None])
|
||||
with pytest.raises(ValueError, match='bad'):
|
||||
await TrainModel().run(sample_input_data)
|
||||
assert mock_wf.execute_activity_method.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_cleanup_failure_does_not_fail_workflow(
|
||||
mock_wf, sample_input_data, mock_train_params
|
||||
):
|
||||
"""After successful training, cleanup failure is logged, workflow still returns result."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
tr = {'run_name': 'rn', 'run_id': 'i', 'run_dir': '/tmp/t'}
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'x': 1},
|
||||
mock_train_params,
|
||||
None,
|
||||
tr,
|
||||
None,
|
||||
RuntimeError('cleanup'),
|
||||
]
|
||||
)
|
||||
mock_wf.logger = Mock()
|
||||
out = await TrainModel().run(sample_input_data)
|
||||
assert out == tr
|
||||
mock_wf.logger.warning.assert_called_once()
|
||||
assert mock_wf.execute_activity_method.call_count == 6
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_training_failure_skips_cleanup_activity(
|
||||
mock_wf, sample_input_data, mock_train_params
|
||||
):
|
||||
"""When train_model raises, train_result stays None and cleanup activity is not scheduled."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.execute_activity_method = AsyncMock(
|
||||
side_effect=[
|
||||
{'x': 1},
|
||||
mock_train_params,
|
||||
None,
|
||||
RuntimeError('train failed'),
|
||||
]
|
||||
)
|
||||
with pytest.raises(RuntimeError, match='train failed'):
|
||||
await TrainModel().run(sample_input_data)
|
||||
# validate (3) + train activity (1) + TRAINING_ERROR DB update (1); no cleanup (6th) when train_result is unset
|
||||
assert mock_wf.execute_activity_method.call_count == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('model_manager.workflows.train_model.workflow')
|
||||
async def test_run_missing_experiment_run_id(mock_wf):
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
mock_wf.logger = Mock()
|
||||
with pytest.raises(ApplicationError, match='experiment_run_id is required'):
|
||||
await TrainModel().run({})
|
||||
|
||||
|
||||
def test_module_constants():
|
||||
from model_manager.workflows.train_model import (
|
||||
TIMEOUT_DELETE_FILE,
|
||||
TIMEOUT_TRAIN_MODEL,
|
||||
TIMEOUT_VALIDATE_PARAMS,
|
||||
database_retry_policy,
|
||||
network_retry_policy,
|
||||
no_retry_policy,
|
||||
)
|
||||
|
||||
assert isinstance(TIMEOUT_VALIDATE_PARAMS, int)
|
||||
assert no_retry_policy.maximum_attempts == 1
|
||||
assert network_retry_policy.maximum_attempts == 5
|
||||
assert database_retry_policy.maximum_attempts == 5
|
||||
assert TIMEOUT_TRAIN_MODEL == 2700
|
||||
assert TIMEOUT_DELETE_FILE == 120
|
||||
Reference in New Issue
Block a user