"""Unit tests for TrainModel workflow.""" from unittest.mock import AsyncMock, Mock, patch import pytest 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', # 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 = 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', # noqa: S108 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', # noqa: S108 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.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.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.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.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', # noqa: S108 } 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', # noqa: S108 } 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_DELETE_FILE, TIMEOUT_TRAIN_MODEL, TIMEOUT_UPDATE_DATABASE, TIMEOUT_VALIDATE_PARAMS, database_retry_policy, network_retry_policy, 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 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'] == ''