"""Unit tests for TrainModel workflow.""" from unittest.mock import AsyncMock, MagicMock, patch import pytest from pytest import fixture, mark from model_manager.utils.models.experiment_status import ExperimentStatus from model_manager.utils.models.train_model_params import TrainModelParams from model_manager.utils.models.train_model_result import TrainModelResult from model_manager.workflows.train_model import TrainModel @fixture def train_model_workflow() -> TrainModel: """Fixture for TrainModel workflow instance.""" return TrainModel() @fixture def mock_train_params(): """Fixture for mock TrainModelParams.""" return TrainModelParams( experiment_run_id=123, target_variable='price', variable_columns=['feature1', 'feature2'], train_size=80, shuffle=True, use_scaler=True, include_ar=False, bucket_name='test-bucket', file_name='test.csv', line_separator='\n', decimal_separator='.', lag_train=1, lag_val=1, rem_static_win=False, low_lim={'feature1': 0.0, 'feature2': 0.0}, upp_lim={'feature1': 100.0, 'feature2': 100.0}, window=10, experiment_name='test_experiment', removed_intervals=[], ) @fixture def mock_train_result(mock_train_params): """Fixture for mock TrainModelResult.""" result = MagicMock(spec=TrainModelResult) result.params = mock_train_params result.run_name = 'test_experiment-1' result.run_dir = 'test_run_dir' # Relative path instead of /tmp return result # ============================================================================ # Tests for run() - Complete workflow # ============================================================================ @mark.asyncio @patch('model_manager.workflows.train_model.shutil') @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_run_success_complete_flow( workflow_mock: AsyncMock, mock_shutil: MagicMock, train_model_workflow: TrainModel, mock_train_params, mock_train_result, ): """Test successful complete workflow execution.""" input_data = { 'experiment_run_id': 123, 'target_variable': 'price', 'variable_columns': ['feature1', 'feature2'], 'train_size': 80, 'shuffle': True, 'use_scaler': True, 'include_ar': False, 'bucket_name': 'test-bucket', 'file_name': 'test.csv', 'line_separator': '\n', 'decimal_separator': '.', 'lag_train': 1, 'lag_val': 1, 'rem_static_win': False, 'low_lim': {'feature1': 0.0, 'feature2': 0.0}, 'upp_lim': {'feature1': 100.0, 'feature2': 100.0}, 'window': 10, 'experiment_name': 'test_experiment', 'removed_intervals': [], } # Mock activity responses workflow_mock.execute_activity_method = AsyncMock( side_effect=[ mock_train_params, # validate_train_params None, # update_experiment_run (MAGE_WAITING_PROC) b'file_content', # fetch_file_from_minio mock_train_result, # train_model None, # update_experiment_run (TRAINING_SUCCESS) mock_train_result, # save_model None, # update_experiment_run (MLFLOW_SENT with run_name) None, # delete_file_from_minio None, # update_experiment_run (FILE_DELETED) ] ) # Execute workflow await train_model_workflow.run(input_data) # Verify all activity calls assert workflow_mock.execute_activity_method.call_count == 9 # Verify shutil.rmtree was called mock_shutil.rmtree.assert_called_once_with('test_run_dir') @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_run_missing_experiment_run_id( workflow_mock: AsyncMock, train_model_workflow: TrainModel ): """Test workflow fails when experiment_run_id is missing.""" input_data = { 'target_variable': 'price', 'variable_columns': ['feature1'], } # Mock workflow.logger to avoid NotInWorkflowEventLoopError workflow_mock.logger = MagicMock() with pytest.raises(ValueError, match='experiment_run_id is required'): await train_model_workflow.run(input_data) @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_run_invalid_experiment_run_id_type( workflow_mock: AsyncMock, train_model_workflow: TrainModel ): """Test workflow fails when experiment_run_id has invalid type.""" input_data = { 'experiment_run_id': 'invalid', # Should be int 'target_variable': 'price', } # Mock workflow.logger to avoid NotInWorkflowEventLoopError workflow_mock.logger = MagicMock() with pytest.raises(ValueError, match='experiment_run_id must be an integer'): await train_model_workflow.run(input_data) # ============================================================================ # Tests for _validate_experiment_run_id() # ============================================================================ @patch('model_manager.workflows.train_model.workflow') def test_validate_experiment_run_id_success( workflow_mock: MagicMock, train_model_workflow: TrainModel ): """Test successful experiment_run_id validation.""" workflow_mock.logger = MagicMock() input_data = {'experiment_run_id': 456} result = train_model_workflow._validate_experiment_run_id(input_data) assert result == 456 @patch('model_manager.workflows.train_model.workflow') def test_validate_experiment_run_id_missing( workflow_mock: MagicMock, train_model_workflow: TrainModel ): """Test validation fails when experiment_run_id is missing.""" workflow_mock.logger = MagicMock() input_data = {} with pytest.raises(ValueError, match='experiment_run_id is required'): train_model_workflow._validate_experiment_run_id(input_data) @patch('model_manager.workflows.train_model.workflow') def test_validate_experiment_run_id_none( workflow_mock: MagicMock, train_model_workflow: TrainModel ): """Test validation fails when experiment_run_id is None.""" workflow_mock.logger = MagicMock() input_data = {'experiment_run_id': None} with pytest.raises(ValueError, match='experiment_run_id is required'): train_model_workflow._validate_experiment_run_id(input_data) @patch('model_manager.workflows.train_model.workflow') def test_validate_experiment_run_id_invalid_type( workflow_mock: MagicMock, train_model_workflow: TrainModel ): """Test validation fails when experiment_run_id is not an integer.""" workflow_mock.logger = MagicMock() input_data = {'experiment_run_id': 'not_an_int'} with pytest.raises(ValueError, match='experiment_run_id must be an integer'): train_model_workflow._validate_experiment_run_id(input_data) # ============================================================================ # Tests for _validate_training_parameters() # ============================================================================ @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_validate_training_parameters_success( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params ): """Test successful parameter validation.""" input_data = { 'experiment_run_id': 123, 'target_variable': 'price', 'variable_columns': ['feature1'], 'train_size': 80, 'shuffle': True, 'use_scaler': False, 'include_ar': False, 'bucket_name': 'test', 'file_name': 'test.csv', 'line_separator': '\n', 'decimal_separator': '.', 'lag_train': 1, 'lag_val': 1, 'rem_static_win': False, 'low_lim': {'feature1': 0.0}, 'upp_lim': {'feature1': 100.0}, 'window': 10, 'experiment_name': 'test', 'removed_intervals': [], } metadata = { 'metadata': { 'experiment_run_id': 123, 'workflow_name': 'train_model', } } workflow_mock.execute_activity_method = AsyncMock( side_effect=[ mock_train_params, # validate_train_params None, # update_experiment_run ] ) result = await train_model_workflow._validate_training_parameters(input_data, 123, metadata) assert result == mock_train_params assert workflow_mock.execute_activity_method.call_count == 2 @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_validate_training_parameters_validation_error( workflow_mock: AsyncMock, train_model_workflow: TrainModel ): """Test parameter validation handles errors correctly.""" input_data = {'experiment_run_id': 123} metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ ValueError('Missing required field'), # validate_train_params fails None, # update_experiment_run with error ] ) with pytest.raises(ValueError, match='Missing required field'): await train_model_workflow._validate_training_parameters(input_data, 123, metadata) # Verify error status was updated assert workflow_mock.execute_activity_method.call_count == 2 # ============================================================================ # Tests for _download_and_train_model() # ============================================================================ @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_download_and_train_model_success( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params, mock_train_result, ): """Test successful download and training.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ b'file_content', # fetch_file_from_minio mock_train_result, # train_model None, # update_experiment_run ] ) result = await train_model_workflow._download_and_train_model( train_params=mock_train_params, experiment_run_id=123, metadata=metadata ) assert result == mock_train_result assert workflow_mock.execute_activity_method.call_count == 3 @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_download_and_train_model_download_error( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params ): """Test download error is handled correctly.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ Exception('MinIO connection failed'), # fetch_file_from_minio fails None, # update_experiment_run with error ] ) with pytest.raises(Exception, match='MinIO connection failed'): await train_model_workflow._download_and_train_model( train_params=mock_train_params, experiment_run_id=123, metadata=metadata ) # Verify error status was updated assert workflow_mock.execute_activity_method.call_count == 2 @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_download_and_train_model_training_error( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params ): """Test training error is handled correctly.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ b'file_content', # fetch_file_from_minio succeeds Exception('Training failed'), # train_model fails None, # update_experiment_run with error ] ) with pytest.raises(Exception, match='Training failed'): await train_model_workflow._download_and_train_model( train_params=mock_train_params, experiment_run_id=123, metadata=metadata ) # Verify error status was updated assert workflow_mock.execute_activity_method.call_count == 3 @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_download_and_train_model_closes_bytesio_on_success( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params, mock_train_result ): """Test that BytesIO is closed in finally block on success.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} # Create a mock BytesIO with close method mock_file = MagicMock() mock_file.close = MagicMock() workflow_mock.execute_activity_method = AsyncMock( side_effect=[ mock_file, # fetch_file_from_minio returns BytesIO mock_train_result, # train_model succeeds None, # update_experiment_run ] ) result = await train_model_workflow._download_and_train_model( train_params=mock_train_params, experiment_run_id=123, metadata=metadata ) # Verify BytesIO.close() was called in finally block mock_file.close.assert_called_once() assert result == mock_train_result @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_download_and_train_model_closes_bytesio_on_error( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params ): """Test that BytesIO is closed in finally block even on error.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} # Create a mock BytesIO with close method mock_file = MagicMock() mock_file.close = MagicMock() workflow_mock.execute_activity_method = AsyncMock( side_effect=[ mock_file, # fetch_file_from_minio returns BytesIO Exception('Training failed'), # train_model fails None, # update_experiment_run with error ] ) with pytest.raises(Exception, match='Training failed'): await train_model_workflow._download_and_train_model( train_params=mock_train_params, experiment_run_id=123, metadata=metadata ) # Verify BytesIO.close() was called in finally block even after exception mock_file.close.assert_called_once() @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_download_and_train_model_handles_file_without_close( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_params, mock_train_result ): """Test that workflow handles file objects without close method gracefully.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} # Create a mock file without close method mock_file = MagicMock(spec=[]) workflow_mock.execute_activity_method = AsyncMock( side_effect=[ mock_file, # fetch_file_from_minio returns object without close mock_train_result, # train_model succeeds None, # update_experiment_run ] ) # Should not raise error even if file doesn't have close method result = await train_model_workflow._download_and_train_model( train_params=mock_train_params, experiment_run_id=123, metadata=metadata ) assert result == mock_train_result # ============================================================================ # Tests for _save_model_to_mlflow() # ============================================================================ @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_save_model_to_mlflow_success( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_result ): """Test successful model saving to MLFlow.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ mock_train_result, # save_model None, # update_experiment_run with MODEL_SAVED ] ) result = await train_model_workflow._save_model_to_mlflow( train_result=mock_train_result, experiment_run_id=123, metadata=metadata ) assert result == mock_train_result assert workflow_mock.execute_activity_method.call_count == 2 @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_save_model_to_mlflow_error( workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_result ): """Test MLFlow save error is handled correctly.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ Exception('MLFlow connection failed'), # save_model fails None, # update_experiment_run with error ] ) with pytest.raises(Exception, match='MLFlow connection failed'): await train_model_workflow._save_model_to_mlflow( train_result=mock_train_result, experiment_run_id=123, metadata=metadata ) # Verify error status was updated assert workflow_mock.execute_activity_method.call_count == 2 # ============================================================================ # Tests for _cleanup_resources() # ============================================================================ @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) @patch('model_manager.workflows.train_model.shutil') async def test_cleanup_resources_success( mock_shutil: MagicMock, workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_result, ): """Test successful resource cleanup.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ None, # delete_file_from_minio None, # update_experiment_run ] ) await train_model_workflow._cleanup_resources( saved_result=mock_train_result, experiment_run_id=123, metadata=metadata ) # Verify shutil.rmtree was called mock_shutil.rmtree.assert_called_once_with('test_run_dir') # Verify activities were called assert workflow_mock.execute_activity_method.call_count == 2 @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) @patch('model_manager.workflows.train_model.shutil') async def test_cleanup_resources_delete_error( mock_shutil: MagicMock, workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_result, ): """Test cleanup handles delete errors correctly.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} workflow_mock.execute_activity_method = AsyncMock( side_effect=[ Exception('MinIO delete failed'), # delete_file_from_minio fails None, # update_experiment_run with error ] ) with pytest.raises(Exception, match='MinIO delete failed'): await train_model_workflow._cleanup_resources( saved_result=mock_train_result, experiment_run_id=123, metadata=metadata ) # Verify error status was updated assert workflow_mock.execute_activity_method.call_count == 2 @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) @patch('model_manager.workflows.train_model.shutil') async def test_cleanup_resources_without_run_dir( mock_shutil: MagicMock, workflow_mock: AsyncMock, train_model_workflow: TrainModel, mock_train_result, ): """Test cleanup works when run_dir is not set.""" metadata = {'metadata': {'experiment_run_id': 123, 'workflow_name': 'train_model'}} # Mock result without run_dir mock_train_result.run_dir = None workflow_mock.execute_activity_method = AsyncMock( side_effect=[ None, # delete_file_from_minio None, # update_experiment_run ] ) await train_model_workflow._cleanup_resources( saved_result=mock_train_result, experiment_run_id=123, metadata=metadata ) # Verify shutil.rmtree was NOT called mock_shutil.rmtree.assert_not_called() # Verify activities were still called assert workflow_mock.execute_activity_method.call_count == 2 # ============================================================================ # Tests for _update_experiment_run() # ============================================================================ @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_update_experiment_run_status_only( workflow_mock: AsyncMock, train_model_workflow: TrainModel ): """Test updating experiment run with status only.""" from model_manager.activities.experiment_tracking import UpdateType metadata = {'metadata': {'experiment_run_id': 123}} workflow_mock.execute_activity_method = AsyncMock(return_value=None) await train_model_workflow._update_experiment_run( metadata=metadata, experiment_run_id=123, update_type=UpdateType.STATUS, status=ExperimentStatus.TRAINING_SUCCESS, ) workflow_mock.execute_activity_method.assert_called_once() @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_update_experiment_run_with_error( workflow_mock: AsyncMock, train_model_workflow: TrainModel ): """Test updating experiment run with error message.""" from model_manager.activities.experiment_tracking import UpdateType metadata = {'metadata': {'experiment_run_id': 123}} workflow_mock.execute_activity_method = AsyncMock(return_value=None) await train_model_workflow._update_experiment_run( metadata=metadata, experiment_run_id=123, update_type=UpdateType.STATUS_WITH_ERROR, status=ExperimentStatus.TRAINING_ERROR, error_message='Training failed', ) workflow_mock.execute_activity_method.assert_called_once() call_args = workflow_mock.execute_activity_method.call_args[0][1] assert call_args['error_message'] == 'Training failed' @mark.asyncio @patch('model_manager.workflows.train_model.workflow', new_callable=AsyncMock) async def test_update_experiment_run_with_run_name( workflow_mock: AsyncMock, train_model_workflow: TrainModel ): """Test updating experiment run with run_name.""" from model_manager.activities.experiment_tracking import UpdateType metadata = {'metadata': {'experiment_run_id': 123}} workflow_mock.execute_activity_method = AsyncMock(return_value=None) await train_model_workflow._update_experiment_run( metadata=metadata, experiment_run_id=123, update_type=UpdateType.MODEL_SAVED, status=ExperimentStatus.MLFLOW_SENT, run_name='test_experiment-1', ) workflow_mock.execute_activity_method.assert_called_once() call_args = workflow_mock.execute_activity_method.call_args[0][1] assert call_args['run_name'] == 'test_experiment-1'