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:
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user