From c63389620a40eec83da292c6420327fb3c5b313f Mon Sep 17 00:00:00 2001 From: Kou-Kinoshita Date: Thu, 30 Oct 2025 10:11:07 -0300 Subject: [PATCH] SIENTIAPDE-1241: Train_model tests --- .../utils/repository/test_model_repository.py | 63 -- tests/workflows/test_train_model.py | 647 ++++++++++++++++++ 2 files changed, 647 insertions(+), 63 deletions(-) create mode 100644 tests/workflows/test_train_model.py diff --git a/tests/utils/repository/test_model_repository.py b/tests/utils/repository/test_model_repository.py index f909b4b..72b18d9 100644 --- a/tests/utils/repository/test_model_repository.py +++ b/tests/utils/repository/test_model_repository.py @@ -802,66 +802,3 @@ def test_generate_artifacts_header_not_found( with pytest.raises(FileNotFoundError, match='Header file does not exist'): repo._generate_artifacts(mock_train_result) - - -@patch('model_manager.utils.repository.model_repository.ModelServing') -@patch('model_manager.utils.repository.model_repository.json.dump') -@patch('model_manager.utils.repository.model_repository.Reports') -@patch('model_manager.utils.repository.model_repository.path.exists') -@patch('model_manager.utils.repository.model_repository.path.join') -@patch('model_manager.utils.repository.model_repository.os.makedirs') -@patch('model_manager.utils.repository.model_repository.shutil.copy') -@patch('builtins.open', create=True) -def test_generate_artifacts_full_success_path( - mock_open, - mock_copy, - mock_makedirs, - mock_path_join, - mock_exists, - mock_reports_class, - mock_json_dump, - mock_model_serving_class, - mock_logger, - mock_train_result, -): - """Test _generate_artifacts complete success path covering lines 147-148.""" - from model_manager.utils.repository.model_repository import ModelRepository - - repo = ModelRepository( - url='http://mlflow.test', username='user', password='pass', logger=mock_logger - ) - - # Mock all path operations - def join_side_effect(*args): - return '/'.join(str(arg) for arg in args) - - mock_path_join.side_effect = join_side_effect - - # Mock path.exists to return True for header file and other files - def exists_side_effect(path_arg): - if 'header.html' in str(path_arg): - return True # Header exists - if 'report.html' in str(path_arg): - return False # Report doesn't exist yet (will be created) - return False - - mock_exists.side_effect = exists_side_effect - - # Mock Reports class - mock_reports_instance = MagicMock() - mock_reports_instance.generate_report = MagicMock(return_value='report') - mock_reports_class.return_value = mock_reports_instance - - # Mock DataFrame.to_csv to avoid actual file I/O - with patch.object(pd.DataFrame, 'to_csv'): - result = repo._generate_artifacts(mock_train_result) - - # Verify the result is returned correctly - assert result == mock_train_result - - # Verify _setup_run_directory was called (line 147) - mock_makedirs.assert_called_once() - mock_copy.assert_called_once() - - # Verify _generate_report was called (line 148) - mock_reports_instance.generate_report.assert_called_once() diff --git a/tests/workflows/test_train_model.py b/tests/workflows/test_train_model.py new file mode 100644 index 0000000..5b74714 --- /dev/null +++ b/tests/workflows/test_train_model.py @@ -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'] == '' +