509 lines
16 KiB
Python
509 lines
16 KiB
Python
"""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
|
|
|
|
controller.shutdown = mock_shutdown
|
|
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('', '')
|