diff --git a/tests/activities/test_training.py b/tests/activities/test_training.py new file mode 100644 index 0000000..f5ca3ba --- /dev/null +++ b/tests/activities/test_training.py @@ -0,0 +1,497 @@ +"""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_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.BaseActivity.__init__', return_value=None) +@patch('model_manager.activities.training.TrainingRepository') +def test_training_init( + mock_training_repo, + mock_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, +): + """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, + ) + + mock_base_init.assert_called_once_with( + mock_logger, mock_notification_handler, set_error_counter=True + ) + 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 + + +@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None) +@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_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, + mock_train_params, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@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_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@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_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@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_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@patch('model_manager.activities.training.TrainingRepository') +def test_train_model_success_with_params_object( + mock_training_repo_class, + mock_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, + mock_train_params, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@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_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, + mock_train_params, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@patch('model_manager.activities.training.TrainingRepository') +def test_train_model_training_fails( + mock_training_repo_class, + mock_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, + mock_train_params, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@patch('model_manager.activities.training.TrainingRepository') +def test_train_model_save_fails( + mock_training_repo_class, + mock_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, + mock_train_params, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@patch('model_manager.activities.training.TrainingRepository') +def test_cleanup_resources_success( + mock_training_repo_class, + mock_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@patch('model_manager.activities.training.TrainingRepository') +def test_cleanup_resources_cleanup_fails( + mock_training_repo_class, + mock_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, +): + """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, + ) + + 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.BaseActivity.__init__', return_value=None) +@patch('model_manager.activities.training.TrainingRepository') +def test_cleanup_resources_with_empty_values( + mock_training_repo_class, + mock_base_init, + mock_model_repository, + mock_storage_repository, + mock_logger, + mock_notification_handler, +): + """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, + ) + + 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('', '')