feat: enhance training and experiment tracking functionality

- Updated `Activities` class to improve garbage collection handling.
- Enhanced error messaging in `ExperimentTracking` for better clarity on update failures.
- Refactored `Training` class to streamline exception handling and improve type hints.
- Introduced new methods in `TrainModelParams` for better handling of experiment run IDs and model metadata.
- Added functionality to extract model equations in `DataManagerRepository` for linear regression models.
This commit is contained in:
vitor-aignosi
2026-04-06 15:05:57 -03:00
parent 1352d1ac8f
commit 6b1df7c3a7
22 changed files with 1751 additions and 2085 deletions

View File

@@ -1,5 +1,6 @@
"""Unit tests for the CleanupFiles workflow."""
import os
from unittest.mock import AsyncMock, patch
import pytest
@@ -7,7 +8,6 @@ import pytest
@pytest.mark.asyncio
@patch('model_manager.workflows.cleanup_files.workflow')
@patch('model_manager.workflows.cleanup_files.POD_ID', 'temporal-pod')
async def test_cleanup_files_workflow(mock_workflow_module):
"""Test the CleanupFiles workflow."""
from model_manager.workflows.cleanup_files import CleanupFiles
@@ -17,7 +17,8 @@ async def test_cleanup_files_workflow(mock_workflow_module):
# Instantiate and run the workflow
workflow_instance = CleanupFiles()
await workflow_instance.run({})
with patch.dict(os.environ, {'POD_ID': 'temporal-pod'}):
await workflow_instance.run({})
# Verify that the activities were called with the correct parameters
calls = mock_workflow_module.execute_activity_method.call_args_list

View File

@@ -10,7 +10,7 @@ from model_manager.utils.models.train_model_params import TrainModelParams
@pytest.fixture
def mock_train_params():
"""Create a mock TrainModelParams object."""
"""Minimal mock TrainModelParams."""
params = Mock(spec=TrainModelParams)
params.experiment_run_id = 123
params.bucket_name = 'test-bucket'
@@ -22,524 +22,278 @@ def mock_train_params():
@pytest.fixture
def sample_input_data():
"""Create sample input data for workflow."""
"""Sample workflow input (IDs normalized in run())."""
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': [],
'date_column': None,
'date_format': None,
'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': {}}}},
}
@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}
wf = TrainModel()
assert wf._validate_experiment_run_id({'experiment_run_id': 123}) == 123
result = workflow_instance._validate_experiment_run_id(input_data)
assert result == 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():
"""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)
with pytest.raises(ValueError, match='experiment_run_id is required'):
TrainModel()._validate_experiment_run_id({})
def test_validate_experiment_run_id_not_integer():
"""Test validation fails when experiment_run_id is not an integer."""
def test_validate_experiment_run_id_invalid_type():
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)
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():
"""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'
assert TrainModel()._extract_error_message(ValueError('x')) == 'x'
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
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():
"""Test extracting error message from exception with empty string."""
def test_extract_error_message_empty_message():
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
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_workflow_module, sample_input_data, mock_train_params
):
"""Test successful parameter validation."""
async def test_validate_training_parameters_success(mock_wf, sample_input_data, mock_train_params):
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', # noqa: S108
'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 = RuntimeError('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(RuntimeError):
await workflow_instance._train_model(mock_train_params, 123, metadata)
# Verify TRAINING_ERROR status was set
call_args = mock_workflow_module.execute_activity_method.call_args_list[1]
assert call_args[0][1]['status'] == ExperimentStatus.TRAINING_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(
run_dir='/tmp/test-run', # noqa: S108
metadata=metadata,
)
assert mock_workflow_module.execute_activity_method.call_count == 1
@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(
run_dir='/tmp/test-run', # noqa: S108
metadata=metadata,
)
# Verify error status update was called
assert mock_workflow_module.execute_activity_method.call_count == 1
@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.activities.experiment_tracking import UpdateType
from model_manager.workflows.train_model import TrainModel
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.ORCHESTRATOR_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.ORCHESTRATOR_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.activities.experiment_tracking import UpdateType
from model_manager.workflows.train_model import TrainModel
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.activities.experiment_tracking import UpdateType
from model_manager.workflows.train_model import TrainModel
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.TRAINING_SUCCESS,
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', # noqa: S108
}
mock_workflow_module.execute_activity_method = AsyncMock(
mock_wf.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate_train_params
None, # update status (ORCHESTRATOR_WAITING_PROC)
train_result, # train_model
None, # update status (TRAINING_SUCCESS)
None, # cleanup_resources
{'experiment_run_id': 123, 'model_metadata': {}},
mock_train_params,
None,
]
)
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 == 5
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_run_workflow_validation_error(mock_workflow_module, sample_input_data):
"""Test workflow handles validation errors."""
async def test_validate_training_parameters_load_fails(mock_wf, sample_input_data):
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
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_run_workflow_training_error(
mock_workflow_module, sample_input_data, mock_train_params
):
"""Test workflow handles training errors."""
async def test_train_model_success(mock_wf, mock_train_params):
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 (ORCHESTRATOR_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
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_run_workflow_cleanup_error(
mock_workflow_module, sample_input_data, mock_train_params
):
"""Test workflow handles cleanup errors."""
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
# Setup mocks - cleanup fails
train_result = {
'run_name': 'test-run-123',
'run_dir': '/tmp/test-run', # noqa: S108
}
mock_workflow_module.execute_activity_method = AsyncMock(
side_effect=[
mock_train_params, # validate_train_params
None, # update status (ORCHESTRATOR_WAITING_PROC)
train_result, # train_model
None, # update status (TRAINING_SUCCESS)
RuntimeError('Cleanup failed'), # cleanup_resources fails
]
mock_wf.execute_activity_method = AsyncMock(
side_effect=[ValueError('validation'), RuntimeError('db')],
)
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 == 5
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
async def test_run_workflow_missing_experiment_run_id():
"""Test workflow fails early when experiment_run_id is missing."""
@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
workflow_instance = TrainModel()
input_data = {} # Missing experiment_run_id
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
async def test_run_missing_experiment_run_id():
from model_manager.workflows.train_model import TrainModel
with pytest.raises(ValueError, match='experiment_run_id is required'):
await workflow_instance.run(input_data)
await TrainModel().run({})
def test_module_constants():
"""Test that module-level constants are defined correctly."""
from model_manager.workflows.train_model import (
TIMEOUT_DELETE_FILE,
TIMEOUT_TRAIN_MODEL,
@@ -550,91 +304,9 @@ def test_module_constants():
no_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 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
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'] == ''