SIENTIAPDE-1241: Add unit tests for the Training activity with 100% coverage.
This commit is contained in:
497
tests/activities/test_training.py
Normal file
497
tests/activities/test_training.py
Normal file
@@ -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('', '')
|
||||
Reference in New Issue
Block a user