SIENTIAPDE-1307: Integrate metrics controller and update sientia-dataops-library to 1.5.1. This change adds metrics collection capabilities to activities and updates the dataops library dependency. (18 files changed, 143 insertions(+), 16 deletions(-))

This commit is contained in:
Bruno Domingues
2025-10-31 17:06:26 -03:00
parent 326f9f9dc6
commit db581f60f9
6 changed files with 102 additions and 44 deletions

View File

@@ -5,6 +5,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from model_manager.activities.experiment_tracking import ExperimentTracking
from model_manager.activities.training import Training
@@ -61,7 +62,10 @@ class Activities(ExperimentTracking, Training):
Raises:
Exception: If any parent class initialization fails
"""
# Initialize parent classes
metrics_controller = MetricsController(
logger=logger,
)
ExperimentTracking.__init__(
self,
host=postgres_config['host'],
@@ -73,6 +77,7 @@ class Activities(ExperimentTracking, Training):
max_connections=postgres_config['max_connections'],
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
)
self.model_repository = ModelRepository(
@@ -101,6 +106,7 @@ class Activities(ExperimentTracking, Training):
storage_repository=self.storage_repository,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
)
def __del__(self):

View File

@@ -19,6 +19,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.temporal.activities.postgres import Postgres
from sqlalchemy import text
@@ -42,10 +43,6 @@ class ExperimentTracking(Postgres):
The activity uses the existing Postgres connection pool and adds experiment-specific
operations with proper error handling and notifications.
Attributes:
logger (Logger): Logger instance for observability
notification_handler (NotificationHandler): Handler for sending notifications
"""
def __init__(
@@ -59,6 +56,7 @@ class ExperimentTracking(Postgres):
max_connections: int,
logger: Logger,
notification_handler: NotificationHandler,
metrics_controller: MetricsController,
):
"""
Initialize ExperimentTracking activity with database configuration.
@@ -73,6 +71,7 @@ class ExperimentTracking(Postgres):
max_connections: Maximum connections in pool
logger: Logger instance for observability
notification_handler: Notification handler for alerts
metrics_controller: Metrics controller for observability
Raises:
ConnectionError: If database connection cannot be established
@@ -87,6 +86,7 @@ class ExperimentTracking(Postgres):
max_connections=max_connections,
logger=logger,
notification_handler=notification_handler,
metrics_controller=metrics_controller,
)
def __del__(self):

View File

@@ -15,7 +15,8 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
from sientia_do.temporal.activities.base import BaseActivity
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from model_manager.utils.exceptions import ModelTrainingError
from model_manager.utils.models.train_model_params import TrainModelParams
@@ -24,18 +25,14 @@ with workflow.unsafe.imports_passed_through():
from model_manager.utils.repository.training_repository import TrainingRepository
class Training(BaseActivity):
class Training(SientiaMonitoring):
"""
Activity for ML model training operations.
This activity extends BaseActivity and handles machine learning model
This activity extends SientiaMonitoring and handles machine learning model
training with comprehensive error handling. It receives pre-downloaded
files from the workflow and returns success/failure status without
raising exceptions.
Attributes:
logger (Logger): Logger instance for observability (inherited from BaseActivity)
notification_handler (NotificationHandler): Handler for sending notifications (inherited)
"""
def __init__(
@@ -44,6 +41,7 @@ class Training(BaseActivity):
storage_repository: StorageRepository,
logger: Logger,
notification_handler: NotificationHandler,
metrics_controller: MetricsController,
):
"""
Initialize Training activity.
@@ -52,7 +50,7 @@ class Training(BaseActivity):
logger: Logger instance for observability
notification_handler: Handler for sending notifications
"""
super().__init__(logger, notification_handler, set_error_counter=True)
super().__init__(logger, notification_handler, metrics_controller, set_error_counter=True)
self.training_repository = TrainingRepository(logger)
self.model_repository = model_repository
self.storage_repository = storage_repository

View File

@@ -3,7 +3,7 @@ psycopg2-binary==2.9.11
sqlalchemy==2.0.44
boto3==1.40.55
botocore==1.40.55
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.4.6
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.5.1
prometheus-client==0.23.1
mlflow==2.10.1
evidently==0.4.21

View File

@@ -18,6 +18,12 @@ def mock_notification_handler():
return MagicMock()
@pytest.fixture
def mock_metrics_controller():
"""Create a mock metrics controller."""
return MagicMock()
@pytest.fixture
def db_config():
"""Create a valid database configuration."""
@@ -34,7 +40,7 @@ def db_config():
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_experiment_tracking_init(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test ExperimentTracking initialization."""
from model_manager.activities.experiment_tracking import ExperimentTracking
@@ -49,6 +55,7 @@ def test_experiment_tracking_init(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_postgres_init.assert_called_once_with(
@@ -61,12 +68,13 @@ def test_experiment_tracking_init(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_experiment_tracking_del_without_engine(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test __del__ when engine attribute does not exist."""
from model_manager.activities.experiment_tracking import ExperimentTracking
@@ -81,6 +89,7 @@ def test_experiment_tracking_del_without_engine(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
if hasattr(et, 'engine'):
@@ -91,7 +100,7 @@ def test_experiment_tracking_del_without_engine(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_experiment_tracking_del_with_engine(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test __del__ when engine exists."""
from model_manager.activities.experiment_tracking import ExperimentTracking
@@ -106,6 +115,7 @@ def test_experiment_tracking_del_with_engine(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.engine = MagicMock()
@@ -120,7 +130,7 @@ def test_experiment_tracking_del_with_engine(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_experiment_tracking_del_with_engine_exception(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test __del__ catches exceptions."""
from model_manager.activities.experiment_tracking import ExperimentTracking
@@ -135,6 +145,7 @@ def test_experiment_tracking_del_with_engine_exception(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.engine = MagicMock()
@@ -154,7 +165,7 @@ def test_experiment_tracking_del_with_engine_exception(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_execute_update_success(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test _execute_update executes query successfully."""
from model_manager.activities.experiment_tracking import ExperimentTracking
@@ -169,6 +180,7 @@ def test_execute_update_success(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_connection = MagicMock()
@@ -187,7 +199,7 @@ def test_execute_update_success(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_status_success(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with STATUS update type."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -202,6 +214,7 @@ def test_update_experiment_run_status_success(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_execute = MagicMock()
@@ -232,7 +245,7 @@ def test_update_experiment_run_status_success(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_status_missing_status(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with STATUS but missing status parameter."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -247,6 +260,7 @@ def test_update_experiment_run_status_missing_status(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
@@ -265,7 +279,7 @@ def test_update_experiment_run_status_missing_status(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_status_with_error_success(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with STATUS_WITH_ERROR update type."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -280,6 +294,7 @@ def test_update_experiment_run_status_with_error_success(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_execute = MagicMock()
@@ -311,7 +326,7 @@ def test_update_experiment_run_status_with_error_success(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_status_with_error_truncate_message(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run truncates error message if too long."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -326,6 +341,7 @@ def test_update_experiment_run_status_with_error_truncate_message(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_execute = MagicMock()
@@ -354,7 +370,7 @@ def test_update_experiment_run_status_with_error_truncate_message(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_status_with_error_missing_error_message(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with STATUS_WITH_ERROR but missing error_message."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -369,6 +385,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
@@ -388,7 +405,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_model_saved_success(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with MODEL_SAVED update type."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -403,6 +420,7 @@ def test_update_experiment_run_model_saved_success(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_execute = MagicMock()
@@ -434,7 +452,7 @@ def test_update_experiment_run_model_saved_success(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_model_saved_missing_run_name(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with MODEL_SAVED but missing run_name."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -449,6 +467,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
@@ -468,7 +487,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_invalid_update_type(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with invalid update_type."""
from model_manager.activities.experiment_tracking import ExperimentTracking
@@ -483,6 +502,7 @@ def test_update_experiment_run_invalid_update_type(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
@@ -501,7 +521,7 @@ def test_update_experiment_run_invalid_update_type(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_no_rows_updated(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run raises error when no rows are updated."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -516,6 +536,7 @@ def test_update_experiment_run_no_rows_updated(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
async def mock_execute_update(*args, **kwargs):
@@ -539,7 +560,7 @@ def test_update_experiment_run_no_rows_updated(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_status_with_error_missing_status(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with STATUS_WITH_ERROR but missing status - covers line 179."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -554,6 +575,7 @@ def test_update_experiment_run_status_with_error_missing_status(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
@@ -573,7 +595,7 @@ def test_update_experiment_run_status_with_error_missing_status(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_update_experiment_run_model_saved_missing_status(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test update_experiment_run with MODEL_SAVED but missing status - covers line 204."""
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
@@ -588,6 +610,7 @@ def test_update_experiment_run_model_saved_missing_status(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
@@ -607,7 +630,7 @@ def test_update_experiment_run_model_saved_missing_status(
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
def test_experiment_tracking_del_with_engine_no_super_del(
mock_postgres_init, db_config, mock_logger, mock_notification_handler
mock_postgres_init, db_config, mock_logger, mock_notification_handler, mock_metrics_controller
):
"""Test __del__ when engine exists but super has no __del__ - covers line 103."""
from model_manager.activities.experiment_tracking import ExperimentTracking
@@ -622,6 +645,7 @@ def test_experiment_tracking_del_with_engine_no_super_del(
max_connections=db_config['max_connections'],
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
et.engine = MagicMock()

View File

@@ -19,6 +19,12 @@ def mock_notification_handler():
return MagicMock()
@pytest.fixture
def mock_metrics_controller():
"""Create a mock metrics controller."""
return MagicMock()
@pytest.fixture
def mock_model_repository():
"""Create a mock model repository."""
@@ -44,7 +50,7 @@ def mock_train_params():
return params
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
def test_training_init(
mock_training_repo,
@@ -53,6 +59,7 @@ def test_training_init(
mock_storage_repository,
mock_logger,
mock_notification_handler,
mock_metrics_controller,
):
"""Test Training initialization."""
from model_manager.activities.training import Training
@@ -62,17 +69,18 @@ def test_training_init(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_base_init.assert_called_once_with(
mock_logger, mock_notification_handler, set_error_counter=True
mock_logger, mock_notification_handler, mock_metrics_controller, 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.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
@patch('model_manager.activities.training.TrainModelParams')
def test_validate_train_params_success(
@@ -84,6 +92,7 @@ def test_validate_train_params_success(
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
@@ -93,6 +102,7 @@ def test_validate_train_params_success(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
training.info = MagicMock()
@@ -112,7 +122,7 @@ def test_validate_train_params_success(
training.info.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
@patch('model_manager.activities.training.TrainModelParams')
def test_validate_train_params_value_error(
@@ -123,6 +133,7 @@ def test_validate_train_params_value_error(
mock_storage_repository,
mock_logger,
mock_notification_handler,
mock_metrics_controller,
):
"""Test validate_train_params with ValueError."""
from model_manager.activities.training import Training
@@ -132,6 +143,7 @@ def test_validate_train_params_value_error(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
training.send_notification = MagicMock()
@@ -148,7 +160,7 @@ def test_validate_train_params_value_error(
training.send_notification.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
@patch('model_manager.activities.training.TrainModelParams')
def test_validate_train_params_type_error(
@@ -159,6 +171,7 @@ def test_validate_train_params_type_error(
mock_storage_repository,
mock_logger,
mock_notification_handler,
mock_metrics_controller,
):
"""Test validate_train_params with TypeError."""
from model_manager.activities.training import Training
@@ -168,6 +181,7 @@ def test_validate_train_params_type_error(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
training.send_notification = MagicMock()
@@ -184,7 +198,7 @@ def test_validate_train_params_type_error(
training.send_notification.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
@patch('model_manager.activities.training.TrainModelParams')
def test_validate_train_params_key_error(
@@ -195,6 +209,7 @@ def test_validate_train_params_key_error(
mock_storage_repository,
mock_logger,
mock_notification_handler,
mock_metrics_controller,
):
"""Test validate_train_params with KeyError."""
from model_manager.activities.training import Training
@@ -204,6 +219,7 @@ def test_validate_train_params_key_error(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
training.send_notification = MagicMock()
@@ -220,7 +236,7 @@ def test_validate_train_params_key_error(
training.send_notification.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
def test_train_model_success_with_params_object(
mock_training_repo_class,
@@ -230,6 +246,7 @@ def test_train_model_success_with_params_object(
mock_logger,
mock_notification_handler,
mock_train_params,
mock_metrics_controller,
):
"""Test train_model with TrainModelParams object."""
from model_manager.activities.training import Training
@@ -239,6 +256,7 @@ def test_train_model_success_with_params_object(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
mock_file = BytesIO(b'test data')
@@ -268,7 +286,7 @@ def test_train_model_success_with_params_object(
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.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
@patch('model_manager.activities.training.TrainModelParams')
def test_train_model_success_with_params_dict(
@@ -280,6 +298,7 @@ def test_train_model_success_with_params_dict(
mock_logger,
mock_notification_handler,
mock_train_params,
mock_metrics_controller,
):
"""Test train_model with dict parameters."""
from model_manager.activities.training import Training
@@ -289,6 +308,7 @@ def test_train_model_success_with_params_dict(
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
@@ -314,7 +334,7 @@ def test_train_model_success_with_params_dict(
mock_train_params_class.from_dict.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
def test_train_model_training_fails(
mock_training_repo_class,
@@ -324,6 +344,7 @@ def test_train_model_training_fails(
mock_logger,
mock_notification_handler,
mock_train_params,
mock_metrics_controller,
):
"""Test train_model when training fails."""
from model_manager.activities.training import Training
@@ -334,6 +355,7 @@ def test_train_model_training_fails(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
training.send_notification = MagicMock()
@@ -354,7 +376,7 @@ def test_train_model_training_fails(
training.send_notification.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
def test_train_model_save_fails(
mock_training_repo_class,
@@ -364,6 +386,7 @@ def test_train_model_save_fails(
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
@@ -374,6 +397,7 @@ def test_train_model_save_fails(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
training.send_notification = MagicMock()
@@ -398,7 +422,7 @@ def test_train_model_save_fails(
training.send_notification.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
def test_cleanup_resources_success(
mock_training_repo_class,
@@ -407,6 +431,7 @@ def test_cleanup_resources_success(
mock_storage_repository,
mock_logger,
mock_notification_handler,
mock_metrics_controller,
):
"""Test cleanup_resources successfully."""
from model_manager.activities.training import Training
@@ -416,6 +441,7 @@ def test_cleanup_resources_success(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
input_data = {
@@ -431,7 +457,7 @@ def test_cleanup_resources_success(
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.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
def test_cleanup_resources_cleanup_fails(
mock_training_repo_class,
@@ -440,6 +466,7 @@ def test_cleanup_resources_cleanup_fails(
mock_storage_repository,
mock_logger,
mock_notification_handler,
mock_metrics_controller,
):
"""Test cleanup_resources when cleanup fails."""
from model_manager.activities.training import Training
@@ -449,6 +476,7 @@ def test_cleanup_resources_cleanup_fails(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
training.send_notification = MagicMock()
@@ -467,7 +495,7 @@ def test_cleanup_resources_cleanup_fails(
training.send_notification.assert_called_once()
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
@patch('model_manager.activities.training.SientiaMonitoring.__init__', return_value=None)
@patch('model_manager.activities.training.TrainingRepository')
def test_cleanup_resources_with_empty_values(
mock_training_repo_class,
@@ -476,6 +504,7 @@ def test_cleanup_resources_with_empty_values(
mock_storage_repository,
mock_logger,
mock_notification_handler,
mock_metrics_controller,
):
"""Test cleanup_resources with empty values."""
from model_manager.activities.training import Training
@@ -485,6 +514,7 @@ def test_cleanup_resources_with_empty_values(
storage_repository=mock_storage_repository,
logger=mock_logger,
notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller,
)
input_data = {