"""Unit tests for Training class with 100% coverage.""" import asyncio from io import BytesIO from unittest.mock import MagicMock, patch import pytest @pytest.fixture def mock_logger(): """Create a mock logger.""" return MagicMock() @pytest.fixture def mock_notification_handler(): """Create a mock notification handler.""" return MagicMock() @pytest.fixture def mock_metrics_controller(): """Create a mock metrics controller.""" controller = MagicMock() # Make shutdown an async coroutine async def mock_shutdown(): pass # Make emit an async coroutine async def mock_emit(*args, **kwargs): pass controller.shutdown = mock_shutdown controller.emit = mock_emit return controller @pytest.fixture def mock_model_repository(): """Create a mock model repository.""" return MagicMock() @pytest.fixture def mock_storage_repository(): """Create a mock storage repository.""" return MagicMock() @pytest.fixture def mock_train_params(): """Create a mock TrainModelParams.""" params = MagicMock() params.experiment_run_id = 1 params.target_variable = 'target' params.experiment_name = 'test_experiment' params.bucket_name = 'test-bucket' params.file_name = 'test-file.csv' params.validate_business_rules = MagicMock() return params @patch('model_manager.activities.training.TrainingRepository') def test_training_init( mock_training_repo, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_metrics_controller, ): """Test Training initialization.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) mock_training_repo.assert_called_once_with(mock_logger) assert training.model_repository is mock_model_repository assert training.storage_repository is mock_storage_repository assert training.metrics_controller is mock_metrics_controller @patch('model_manager.activities.training.TrainingRepository') @patch('model_manager.activities.training.TrainModelParams') def test_validate_train_params_success( mock_train_params_class, mock_training_repo, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_train_params, mock_metrics_controller, ): """Test validate_train_params with valid parameters.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) training.info = MagicMock() mock_train_params_class.from_dict.return_value = mock_train_params input_data = { 'metadata': {'workflow_id': 'test-123'}, 'experiment_run_id': 1, 'target_variable': 'target', } result = asyncio.run(training.validate_train_params(input_data)) assert result is mock_train_params mock_train_params_class.from_dict.assert_called_once_with(input_data) mock_train_params.validate_business_rules.assert_called_once() training.info.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') @patch('model_manager.activities.training.TrainModelParams') def test_validate_train_params_value_error( mock_train_params_class, mock_training_repo, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_metrics_controller, ): """Test validate_train_params with ValueError.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) training.send_notification = MagicMock() mock_train_params_class.from_dict.side_effect = ValueError('Invalid parameter') input_data = { 'metadata': {'workflow_id': 'test-123'}, 'experiment_run_id': 1, } with pytest.raises(ValueError): asyncio.run(training.validate_train_params(input_data)) training.send_notification.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') @patch('model_manager.activities.training.TrainModelParams') def test_validate_train_params_type_error( mock_train_params_class, mock_training_repo, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_metrics_controller, ): """Test validate_train_params with TypeError.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) training.send_notification = MagicMock() mock_train_params_class.from_dict.side_effect = TypeError('Type mismatch') input_data = { 'metadata': {'workflow_id': 'test-123'}, 'experiment_run_id': 1, } with pytest.raises(TypeError): asyncio.run(training.validate_train_params(input_data)) training.send_notification.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') @patch('model_manager.activities.training.TrainModelParams') def test_validate_train_params_key_error( mock_train_params_class, mock_training_repo, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_metrics_controller, ): """Test validate_train_params with KeyError.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) training.send_notification = MagicMock() mock_train_params_class.from_dict.side_effect = KeyError('missing_key') input_data = { 'metadata': {'workflow_id': 'test-123'}, 'experiment_run_id': 1, } with pytest.raises(KeyError): asyncio.run(training.validate_train_params(input_data)) training.send_notification.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') def test_train_model_success_with_params_object( mock_training_repo_class, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_train_params, mock_metrics_controller, ): """Test train_model with TrainModelParams object.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) mock_file = BytesIO(b'test data') mock_storage_repository.fetch_file.return_value.__enter__.return_value = mock_file mock_train_result = MagicMock() mock_train_result.run_name = 'run_001' mock_train_result.run_dir = '/tmp/run_001' # noqa: S108 training.training_repository.train.return_value = mock_train_result training.training_repository.after_train_calculation.return_value = mock_train_result mock_model_repository.save_model.return_value = mock_train_result input_data = { 'metadata': {'workflow_id': 'test-123'}, 'train_params': mock_train_params, } result = asyncio.run(training.train_model(input_data)) assert result == {'run_name': 'run_001', 'run_dir': '/tmp/run_001'} # noqa: S108 mock_storage_repository.fetch_file.assert_called_once_with('test-bucket', 'test-file.csv') training.training_repository.train.assert_called_once_with(mock_file, mock_train_params) training.training_repository.after_train_calculation.assert_called_once_with( mock_train_params, mock_train_result ) mock_model_repository.save_model.assert_called_once_with(mock_train_result) @patch('model_manager.activities.training.TrainingRepository') @patch('model_manager.activities.training.TrainModelParams') def test_train_model_success_with_params_dict( mock_train_params_class, mock_training_repo_class, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_train_params, mock_metrics_controller, ): """Test train_model with dict parameters.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) mock_train_params_class.from_dict.return_value = mock_train_params mock_file = BytesIO(b'test data') mock_storage_repository.fetch_file.return_value.__enter__.return_value = mock_file mock_train_result = MagicMock() mock_train_result.run_name = 'run_002' mock_train_result.run_dir = '/tmp/run_002' # noqa: S108 training.training_repository.train.return_value = mock_train_result training.training_repository.after_train_calculation.return_value = mock_train_result mock_model_repository.save_model.return_value = mock_train_result input_data = { 'metadata': {'workflow_id': 'test-123'}, 'train_params': {'experiment_run_id': 1, 'target_variable': 'target'}, } result = asyncio.run(training.train_model(input_data)) assert result == {'run_name': 'run_002', 'run_dir': '/tmp/run_002'} # noqa: S108 mock_train_params_class.from_dict.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') def test_train_model_training_fails( mock_training_repo_class, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_train_params, mock_metrics_controller, ): """Test train_model when training fails.""" from model_manager.activities.training import Training from model_manager.utils.exceptions import ModelTrainingError training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) training.send_notification = MagicMock() mock_file = BytesIO(b'test data') mock_storage_repository.fetch_file.return_value.__enter__.return_value = mock_file training.training_repository.train.side_effect = RuntimeError('Training failed') input_data = { 'metadata': {'workflow_id': 'test-123'}, 'train_params': mock_train_params, } with pytest.raises(ModelTrainingError) as exc_info: asyncio.run(training.train_model(input_data)) assert exc_info.value.model_trained is False assert exc_info.value.model_saved is False training.send_notification.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') def test_train_model_save_fails( mock_training_repo_class, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_train_params, mock_metrics_controller, ): """Test train_model when model saving fails.""" from model_manager.activities.training import Training from model_manager.utils.exceptions import ModelTrainingError training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) training.send_notification = MagicMock() mock_file = BytesIO(b'test data') mock_storage_repository.fetch_file.return_value.__enter__.return_value = mock_file mock_train_result = MagicMock() training.training_repository.train.return_value = mock_train_result training.training_repository.after_train_calculation.return_value = mock_train_result mock_model_repository.save_model.side_effect = RuntimeError('Save failed') input_data = { 'metadata': {'workflow_id': 'test-123'}, 'train_params': mock_train_params, } with pytest.raises(ModelTrainingError) as exc_info: asyncio.run(training.train_model(input_data)) assert exc_info.value.model_trained is True assert exc_info.value.model_saved is False training.send_notification.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') def test_cleanup_resources_success( mock_training_repo_class, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_metrics_controller, ): """Test cleanup_resources successfully.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) input_data = { 'metadata': {'workflow_id': 'test-123'}, 'run_dir': '/tmp/run_001', # noqa: S108 'bucket_name': 'test-bucket', 'file_name': 'test-file.csv', } asyncio.run(training.cleanup_resources(input_data)) mock_model_repository.cleanup_run_directory.assert_called_once_with('/tmp/run_001') # noqa: S108 mock_storage_repository.delete_file.assert_called_once_with('test-bucket', 'test-file.csv') @patch('model_manager.activities.training.TrainingRepository') def test_cleanup_resources_cleanup_fails( mock_training_repo_class, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_metrics_controller, ): """Test cleanup_resources when cleanup fails.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) training.send_notification = MagicMock() mock_model_repository.cleanup_run_directory.side_effect = RuntimeError('Cleanup failed') input_data = { 'metadata': {'workflow_id': 'test-123'}, 'run_dir': '/tmp/run_001', # noqa: S108 'bucket_name': 'test-bucket', 'file_name': 'test-file.csv', } with pytest.raises(RuntimeError): asyncio.run(training.cleanup_resources(input_data)) training.send_notification.assert_called_once() @patch('model_manager.activities.training.TrainingRepository') def test_cleanup_resources_with_empty_values( mock_training_repo_class, mock_model_repository, mock_storage_repository, mock_logger, mock_notification_handler, mock_metrics_controller, ): """Test cleanup_resources with empty values.""" from model_manager.activities.training import Training training = Training( model_repository=mock_model_repository, storage_repository=mock_storage_repository, logger=mock_logger, notification_handler=mock_notification_handler, metrics_controller=mock_metrics_controller, ) input_data = { 'metadata': {'workflow_id': 'test-123'}, } asyncio.run(training.cleanup_resources(input_data)) mock_model_repository.cleanup_run_directory.assert_called_once_with('') mock_storage_repository.delete_file.assert_called_once_with('', '')