Files
sientia-dataops-model-manager/tests/activities/test_activities.py
Bruno Domingues bee6036205 SIENTIAPDE-1251: Implement ML model training activity and repository
This commit introduces the 'Training' activity and 'TrainingRepository' for handling ML model training operations within the Model Manager system.

- Added model_manager/activities/training.py for the Training activity, which extends BaseActivity and integrates with Temporal workflows.
- Added model_manager/utils/repository/training_repository.py for the TrainingRepository, which encapsulates the core training logic.
- Updated model_manager/activities/activities.py to include the Training activity in the main activities orchestrator.
- Updated README.md to document the new 'Training' component.
- Added unit tests for the new activity and repository.
2025-10-09 15:07:27 -03:00

153 lines
4.8 KiB
Python

from unittest.mock import ANY, MagicMock, patch
from pytest import mark
from model_manager.activities.activities import Activities
from model_manager.activities.experiment_tracking import ExperimentTracking
from model_manager.activities.gates import Gates
from model_manager.activities.mlflow import MLFlow
from model_manager.activities.training import Training
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
@patch('model_manager.activities.activities.MLFlow.__init__')
@patch('model_manager.activities.activities.MinIO.__init__')
@patch('model_manager.activities.activities.Gates.__init__')
@patch('model_manager.activities.activities.Training.__init__')
def test___init__(
mock_training_init,
mock_gates_init,
mock_minio_init,
mock_mlflow_init,
mock_experiment_tracking_init,
):
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
assert isinstance(activities, Activities)
assert isinstance(activities, ExperimentTracking)
assert isinstance(activities, MLFlow)
assert isinstance(activities, Gates)
assert isinstance(activities, Training)
mock_experiment_tracking_init.assert_called_once_with(
ANY,
host=postgres_config['host'],
port=postgres_config['port'],
user=postgres_config['user'],
password=postgres_config['password'],
dbname=postgres_config['dbname'],
min_connections=postgres_config['min_connections'],
max_connections=postgres_config['max_connections'],
logger=logger,
notification_handler=notification_handler,
)
mock_mlflow_init.assert_called_once_with(
ANY,
mlflow_host=mlflow_config['host'],
mlflow_port=mlflow_config['port'],
mlflow_username=mlflow_config['username'],
mlflow_password=mlflow_config['password'],
logger=logger,
notification_handler=notification_handler,
)
mock_minio_init.assert_called_once_with(
ANY,
endpoint_url=minio_config['endpoint_url'],
access_key=minio_config['access_key'],
secret_key=minio_config['secret_key'],
region=minio_config['region'],
use_ssl=minio_config['use_ssl'],
max_retry_attempts=minio_config['max_retry_attempts'],
retry_mode=minio_config['retry_mode'],
connect_timeout=minio_config['connect_timeout'],
read_timeout=minio_config['read_timeout'],
logger=logger,
notification_handler=notification_handler,
)
mock_gates_init.assert_called_once_with(
ANY, logger=logger, notification_handler=notification_handler
)
mock_training_init.assert_called_once_with(
ANY, logger=logger, notification_handler=notification_handler
)
@mark.asyncio
@patch('model_manager.activities.activities.ExperimentTracking', return_value=MagicMock())
@patch('model_manager.activities.activities.MLFlow', return_value=MagicMock())
async def test_shutdown(_mock_mlflow_init, mock_experiment_tracking_init):
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
mlflow_config = {'host': 'localhost', 'port': 5000, 'username': 'mlflow', 'password': 'mlflow'}
minio_config = {
'endpoint_url': 'http://localhost:9000',
'access_key': 'minioadmin',
'secret_key': 'minioadmin',
'region': 'us-east-1',
'use_ssl': False,
'max_retry_attempts': 3,
'retry_mode': 'adaptive',
'connect_timeout': 10,
'read_timeout': 60,
}
logger = MagicMock()
notification_handler = MagicMock()
activities = Activities(
postgres_config=postgres_config,
mlflow_config=mlflow_config,
minio_config=minio_config,
logger=logger,
notification_handler=notification_handler,
)
await activities.shutdown()
mock_experiment_tracking_init.close.assert_called_once()