Files
sientia-dataops-model-manager/tests/activities/test_training.py
vitor-aignosi 342a02d6f7 feat: enhance training workflow with model metadata loading and refactor data handling
- Introduced a new activity to load model metadata from the model store.
- Refactored training logic to utilize new model metadata and improved parameter handling.
- Updated the `TrainModelParams` class to include additional fields for model configuration.
- Replaced deprecated utility functions with a custom train-test split implementation.
- Removed unused utility functions and cleaned up the data manager repository.
- Adjusted experiment tracking to include model-specific metadata in notifications.
2026-03-24 14:39:28 -03:00

514 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
# 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.model_name = 'Linear Regression'
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('', '')