Merge pull request #11 from Aignosi/feature/SIENTIAPDE-1307
SIENTIAPDE-1307: Implement and Integrate Metrics Tracking for Training Activities and Workflows
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.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||||
from sientia_do.observability.logger import Logger
|
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.experiment_tracking import ExperimentTracking
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -61,7 +62,10 @@ class Activities(ExperimentTracking, Training):
|
|||||||
Raises:
|
Raises:
|
||||||
Exception: If any parent class initialization fails
|
Exception: If any parent class initialization fails
|
||||||
"""
|
"""
|
||||||
# Initialize parent classes
|
metrics_controller = MetricsController(
|
||||||
|
logger=logger,
|
||||||
|
)
|
||||||
|
|
||||||
ExperimentTracking.__init__(
|
ExperimentTracking.__init__(
|
||||||
self,
|
self,
|
||||||
host=postgres_config['host'],
|
host=postgres_config['host'],
|
||||||
@@ -73,6 +77,7 @@ class Activities(ExperimentTracking, Training):
|
|||||||
max_connections=postgres_config['max_connections'],
|
max_connections=postgres_config['max_connections'],
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler,
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.model_repository = ModelRepository(
|
self.model_repository = ModelRepository(
|
||||||
@@ -101,6 +106,7 @@ class Activities(ExperimentTracking, Training):
|
|||||||
storage_repository=self.storage_repository,
|
storage_repository=self.storage_repository,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler,
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
def __del__(self):
|
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.handlers import CoreNotificationHandler as NotificationHandler
|
||||||
from sientia_do.notifications.models import NotificationLevel
|
from sientia_do.notifications.models import NotificationLevel
|
||||||
from sientia_do.observability.logger import Logger
|
from sientia_do.observability.logger import Logger
|
||||||
|
from sientia_do.observability.metrics_controller import MetricsController
|
||||||
from sientia_do.temporal.activities.postgres import Postgres
|
from sientia_do.temporal.activities.postgres import Postgres
|
||||||
from sqlalchemy import text
|
from sqlalchemy import text
|
||||||
|
|
||||||
@@ -42,10 +43,6 @@ class ExperimentTracking(Postgres):
|
|||||||
|
|
||||||
The activity uses the existing Postgres connection pool and adds experiment-specific
|
The activity uses the existing Postgres connection pool and adds experiment-specific
|
||||||
operations with proper error handling and notifications.
|
operations with proper error handling and notifications.
|
||||||
|
|
||||||
Attributes:
|
|
||||||
logger (Logger): Logger instance for observability
|
|
||||||
notification_handler (NotificationHandler): Handler for sending notifications
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -59,6 +56,7 @@ class ExperimentTracking(Postgres):
|
|||||||
max_connections: int,
|
max_connections: int,
|
||||||
logger: Logger,
|
logger: Logger,
|
||||||
notification_handler: NotificationHandler,
|
notification_handler: NotificationHandler,
|
||||||
|
metrics_controller: MetricsController,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Initialize ExperimentTracking activity with database configuration.
|
Initialize ExperimentTracking activity with database configuration.
|
||||||
@@ -73,6 +71,7 @@ class ExperimentTracking(Postgres):
|
|||||||
max_connections: Maximum connections in pool
|
max_connections: Maximum connections in pool
|
||||||
logger: Logger instance for observability
|
logger: Logger instance for observability
|
||||||
notification_handler: Notification handler for alerts
|
notification_handler: Notification handler for alerts
|
||||||
|
metrics_controller: Metrics controller for observability
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
ConnectionError: If database connection cannot be established
|
ConnectionError: If database connection cannot be established
|
||||||
@@ -87,6 +86,7 @@ class ExperimentTracking(Postgres):
|
|||||||
max_connections=max_connections,
|
max_connections=max_connections,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
notification_handler=notification_handler,
|
notification_handler=notification_handler,
|
||||||
|
metrics_controller=metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
def __del__(self):
|
def __del__(self):
|
||||||
|
|||||||
@@ -15,8 +15,10 @@ with workflow.unsafe.imports_passed_through():
|
|||||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||||
from sientia_do.notifications.models import NotificationLevel
|
from sientia_do.notifications.models import NotificationLevel
|
||||||
from sientia_do.observability.logger import Logger
|
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.metrics import ACTIVITY_EXECUTION_TOTAL, WORKFLOW_EXECUTION_TOTAL
|
||||||
from model_manager.utils.exceptions import ModelTrainingError
|
from model_manager.utils.exceptions import ModelTrainingError
|
||||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||||
from model_manager.utils.repository.model_repository import ModelRepository
|
from model_manager.utils.repository.model_repository import ModelRepository
|
||||||
@@ -24,18 +26,14 @@ with workflow.unsafe.imports_passed_through():
|
|||||||
from model_manager.utils.repository.training_repository import TrainingRepository
|
from model_manager.utils.repository.training_repository import TrainingRepository
|
||||||
|
|
||||||
|
|
||||||
class Training(BaseActivity):
|
class Training(SientiaMonitoring):
|
||||||
"""
|
"""
|
||||||
Activity for ML model training operations.
|
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
|
training with comprehensive error handling. It receives pre-downloaded
|
||||||
files from the workflow and returns success/failure status without
|
files from the workflow and returns success/failure status without
|
||||||
raising exceptions.
|
raising exceptions.
|
||||||
|
|
||||||
Attributes:
|
|
||||||
logger (Logger): Logger instance for observability (inherited from BaseActivity)
|
|
||||||
notification_handler (NotificationHandler): Handler for sending notifications (inherited)
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -44,6 +42,7 @@ class Training(BaseActivity):
|
|||||||
storage_repository: StorageRepository,
|
storage_repository: StorageRepository,
|
||||||
logger: Logger,
|
logger: Logger,
|
||||||
notification_handler: NotificationHandler,
|
notification_handler: NotificationHandler,
|
||||||
|
metrics_controller: MetricsController,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Initialize Training activity.
|
Initialize Training activity.
|
||||||
@@ -52,7 +51,7 @@ class Training(BaseActivity):
|
|||||||
logger: Logger instance for observability
|
logger: Logger instance for observability
|
||||||
notification_handler: Handler for sending notifications
|
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.training_repository = TrainingRepository(logger)
|
||||||
self.model_repository = model_repository
|
self.model_repository = model_repository
|
||||||
self.storage_repository = storage_repository
|
self.storage_repository = storage_repository
|
||||||
@@ -78,6 +77,7 @@ class Training(BaseActivity):
|
|||||||
ValueError, TypeError, KeyError: If validation fails (after sending notification)
|
ValueError, TypeError, KeyError: If validation fails (after sending notification)
|
||||||
"""
|
"""
|
||||||
metadata = input_data.get('metadata', {})
|
metadata = input_data.get('metadata', {})
|
||||||
|
metrics_status = 'success'
|
||||||
|
|
||||||
try:
|
try:
|
||||||
train_params = TrainModelParams.from_dict(input_data)
|
train_params = TrainModelParams.from_dict(input_data)
|
||||||
@@ -92,6 +92,7 @@ class Training(BaseActivity):
|
|||||||
|
|
||||||
return train_params
|
return train_params
|
||||||
except (ValueError, TypeError, KeyError) as e:
|
except (ValueError, TypeError, KeyError) as e:
|
||||||
|
metrics_status = 'error'
|
||||||
error_msg = f'Error validating training parameters: {str(e)}'
|
error_msg = f'Error validating training parameters: {str(e)}'
|
||||||
trace = traceback.format_exc()
|
trace = traceback.format_exc()
|
||||||
|
|
||||||
@@ -104,6 +105,13 @@ class Training(BaseActivity):
|
|||||||
attachment_content=trace,
|
attachment_content=trace,
|
||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
|
finally:
|
||||||
|
await self._emit_metrics(
|
||||||
|
metadata=metadata,
|
||||||
|
metrics_status=metrics_status,
|
||||||
|
activity_name='validate_train_params',
|
||||||
|
emit_workflow_metric=(metrics_status == 'error'),
|
||||||
|
)
|
||||||
|
|
||||||
@activity.defn(name='train_model')
|
@activity.defn(name='train_model')
|
||||||
async def train_model(self, input_data: dict[str, Any]) -> dict[str, str | None]:
|
async def train_model(self, input_data: dict[str, Any]) -> dict[str, str | None]:
|
||||||
@@ -137,6 +145,7 @@ class Training(BaseActivity):
|
|||||||
|
|
||||||
model_trained = False
|
model_trained = False
|
||||||
model_saved = False
|
model_saved = False
|
||||||
|
metrics_status = 'success'
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with self.storage_repository.fetch_file(
|
with self.storage_repository.fetch_file(
|
||||||
@@ -157,6 +166,8 @@ class Training(BaseActivity):
|
|||||||
'run_dir': train_result.run_dir,
|
'run_dir': train_result.run_dir,
|
||||||
}
|
}
|
||||||
except Exception as e: # noqa: BLE001
|
except Exception as e: # noqa: BLE001
|
||||||
|
metrics_status = 'error'
|
||||||
|
|
||||||
error_msg = (
|
error_msg = (
|
||||||
'Error training model - '
|
'Error training model - '
|
||||||
f'model_trained={model_trained}, model_saved={model_saved}, '
|
f'model_trained={model_trained}, model_saved={model_saved}, '
|
||||||
@@ -178,6 +189,13 @@ class Training(BaseActivity):
|
|||||||
model_trained=model_trained,
|
model_trained=model_trained,
|
||||||
model_saved=model_saved,
|
model_saved=model_saved,
|
||||||
) from e
|
) from e
|
||||||
|
finally:
|
||||||
|
await self._emit_metrics(
|
||||||
|
metadata=metadata,
|
||||||
|
metrics_status=metrics_status,
|
||||||
|
activity_name='train_model',
|
||||||
|
emit_workflow_metric=(metrics_status == 'error'),
|
||||||
|
)
|
||||||
|
|
||||||
@activity.defn(name='cleanup_resources')
|
@activity.defn(name='cleanup_resources')
|
||||||
async def cleanup_resources(self, input_data: dict[str, Any]) -> None:
|
async def cleanup_resources(self, input_data: dict[str, Any]) -> None:
|
||||||
@@ -198,11 +216,14 @@ class Training(BaseActivity):
|
|||||||
run_dir = input_data.get('run_dir', '')
|
run_dir = input_data.get('run_dir', '')
|
||||||
bucket_name = input_data.get('bucket_name', '')
|
bucket_name = input_data.get('bucket_name', '')
|
||||||
file_name = input_data.get('file_name', '')
|
file_name = input_data.get('file_name', '')
|
||||||
|
metrics_status = 'success'
|
||||||
|
|
||||||
try:
|
try:
|
||||||
self.model_repository.cleanup_run_directory(run_dir)
|
self.model_repository.cleanup_run_directory(run_dir)
|
||||||
self.storage_repository.delete_file(bucket_name, file_name)
|
self.storage_repository.delete_file(bucket_name, file_name)
|
||||||
except Exception as e: # noqa: BLE001
|
except Exception as e: # noqa: BLE001
|
||||||
|
metrics_status = 'error'
|
||||||
|
|
||||||
error_msg = (
|
error_msg = (
|
||||||
f'Error cleaning up resources - Run directory: {run_dir}, '
|
f'Error cleaning up resources - Run directory: {run_dir}, '
|
||||||
f'File: {bucket_name}/{file_name}, Error: {str(e)}'
|
f'File: {bucket_name}/{file_name}, Error: {str(e)}'
|
||||||
@@ -218,4 +239,46 @@ class Training(BaseActivity):
|
|||||||
level=NotificationLevel.ERROR,
|
level=NotificationLevel.ERROR,
|
||||||
attachment_content=trace,
|
attachment_content=trace,
|
||||||
)
|
)
|
||||||
|
|
||||||
raise
|
raise
|
||||||
|
finally:
|
||||||
|
await self._emit_metrics(
|
||||||
|
metadata=metadata,
|
||||||
|
metrics_status=metrics_status,
|
||||||
|
activity_name='cleanup_resources',
|
||||||
|
emit_workflow_metric=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _emit_metrics(
|
||||||
|
self,
|
||||||
|
metadata: dict[str, Any],
|
||||||
|
metrics_status: str,
|
||||||
|
activity_name: str,
|
||||||
|
emit_workflow_metric: bool,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Emit workflow and activity execution metrics.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
metadata: Activity metadata containing pod_id and workflow_name
|
||||||
|
metrics_status: Execution status ('success' or 'error')
|
||||||
|
activity_name: Name of the activity being executed
|
||||||
|
"""
|
||||||
|
if emit_workflow_metric:
|
||||||
|
await self.emit_metric(
|
||||||
|
metric_object=WORKFLOW_EXECUTION_TOTAL,
|
||||||
|
tags={
|
||||||
|
'pod_id': metadata.get('pod_id'),
|
||||||
|
'workflow_name': metadata.get('workflow_name'),
|
||||||
|
'status': metrics_status,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
await self.emit_metric(
|
||||||
|
metric_object=ACTIVITY_EXECUTION_TOTAL,
|
||||||
|
tags={
|
||||||
|
'pod_id': metadata.get('pod_id'),
|
||||||
|
'activity_name': activity_name,
|
||||||
|
'status': metrics_status,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ Metric Labels:
|
|||||||
- pod_id: Kubernetes pod identifier for multi-instance deployments
|
- pod_id: Kubernetes pod identifier for multi-instance deployments
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from prometheus_client import Gauge
|
from prometheus_client import Counter, Gauge
|
||||||
|
|
||||||
# Application health metric
|
# Application health metric
|
||||||
APP_UP = Gauge(
|
APP_UP = Gauge(
|
||||||
@@ -24,3 +24,15 @@ APP_UP = Gauge(
|
|||||||
'Indicates if the application is running (1) or shutting down (0)',
|
'Indicates if the application is running (1) or shutting down (0)',
|
||||||
['pod_id'],
|
['pod_id'],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
WORKFLOW_EXECUTION_TOTAL = Counter(
|
||||||
|
'model_manager_workflow_executions_total',
|
||||||
|
'Total number of workflow executions',
|
||||||
|
['pod_id', 'workflow_name', 'status'], # status: success, error
|
||||||
|
)
|
||||||
|
|
||||||
|
ACTIVITY_EXECUTION_TOTAL = Counter(
|
||||||
|
'model_manager_activity_executions_total',
|
||||||
|
'Total number of activity executions',
|
||||||
|
['pod_id', 'activity_name', 'status'], # status: success, error
|
||||||
|
)
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ psycopg2-binary==2.9.11
|
|||||||
sqlalchemy==2.0.44
|
sqlalchemy==2.0.44
|
||||||
boto3==1.40.55
|
boto3==1.40.55
|
||||||
botocore==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
|
prometheus-client==0.23.1
|
||||||
mlflow==2.10.1
|
mlflow==2.10.1
|
||||||
evidently==0.4.21
|
evidently==0.4.21
|
||||||
|
|||||||
@@ -296,12 +296,21 @@ def test_activities_del_with_engine_exception_caught(
|
|||||||
|
|
||||||
class MockSuperWithError:
|
class MockSuperWithError:
|
||||||
def __del__(self):
|
def __del__(self):
|
||||||
raise RuntimeError('Test error')
|
# Only raise error if not being cleaned up by garbage collector
|
||||||
|
# This prevents the PytestUnraisableExceptionWarning
|
||||||
|
if hasattr(self, '_should_raise') and self._should_raise:
|
||||||
|
raise RuntimeError('Test error')
|
||||||
|
|
||||||
# Suppress the PytestUnraisableExceptionWarning for this specific test
|
# Suppress the PytestUnraisableExceptionWarning for this specific test
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
||||||
|
|
||||||
with patch('builtins.super', return_value=MockSuperWithError()):
|
mock_super = MockSuperWithError()
|
||||||
activities.__del__()
|
mock_super._should_raise = True
|
||||||
|
try:
|
||||||
|
with patch('builtins.super', return_value=mock_super):
|
||||||
|
activities.__del__()
|
||||||
|
finally:
|
||||||
|
# Prevent the exception from being raised during garbage collection
|
||||||
|
mock_super._should_raise = False
|
||||||
|
|||||||
@@ -18,6 +18,19 @@ def mock_notification_handler():
|
|||||||
return MagicMock()
|
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
|
@pytest.fixture
|
||||||
def db_config():
|
def db_config():
|
||||||
"""Create a valid database configuration."""
|
"""Create a valid database configuration."""
|
||||||
@@ -32,9 +45,8 @@ def db_config():
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_experiment_tracking_init(
|
def test_experiment_tracking_init(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test ExperimentTracking initialization."""
|
"""Test ExperimentTracking initialization."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
@@ -49,24 +61,12 @@ def test_experiment_tracking_init(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
)
|
metrics_controller=mock_metrics_controller,
|
||||||
|
|
||||||
mock_postgres_init.assert_called_once_with(
|
|
||||||
host=db_config['host'],
|
|
||||||
port=db_config['port'],
|
|
||||||
user=db_config['user'],
|
|
||||||
password=db_config['password'],
|
|
||||||
dbname=db_config['dbname'],
|
|
||||||
min_connections=db_config['min_connections'],
|
|
||||||
max_connections=db_config['max_connections'],
|
|
||||||
logger=mock_logger,
|
|
||||||
notification_handler=mock_notification_handler,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_experiment_tracking_del_without_engine(
|
def test_experiment_tracking_del_without_engine(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test __del__ when engine attribute does not exist."""
|
"""Test __del__ when engine attribute does not exist."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
@@ -81,6 +81,7 @@ def test_experiment_tracking_del_without_engine(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
if hasattr(et, 'engine'):
|
if hasattr(et, 'engine'):
|
||||||
@@ -89,9 +90,8 @@ def test_experiment_tracking_del_without_engine(
|
|||||||
et.__del__()
|
et.__del__()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_experiment_tracking_del_with_engine(
|
def test_experiment_tracking_del_with_engine(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test __del__ when engine exists."""
|
"""Test __del__ when engine exists."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
@@ -106,6 +106,7 @@ def test_experiment_tracking_del_with_engine(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.engine = MagicMock()
|
et.engine = MagicMock()
|
||||||
@@ -118,9 +119,8 @@ def test_experiment_tracking_del_with_engine(
|
|||||||
et.__del__()
|
et.__del__()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_experiment_tracking_del_with_engine_exception(
|
def test_experiment_tracking_del_with_engine_exception(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test __del__ catches exceptions."""
|
"""Test __del__ catches exceptions."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
@@ -135,26 +135,35 @@ def test_experiment_tracking_del_with_engine_exception(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.engine = MagicMock()
|
et.engine = MagicMock()
|
||||||
|
|
||||||
class MockSuperWithError:
|
class MockSuperWithError:
|
||||||
def __del__(self):
|
def __del__(self):
|
||||||
raise RuntimeError('Test error')
|
# Only raise error if not being cleaned up by garbage collector
|
||||||
|
# This prevents the PytestUnraisableExceptionWarning
|
||||||
|
if hasattr(self, '_should_raise') and self._should_raise:
|
||||||
|
raise RuntimeError('Test error')
|
||||||
|
|
||||||
# Suppress the PytestUnraisableExceptionWarning for this specific test
|
# Suppress the PytestUnraisableExceptionWarning for this specific test
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
||||||
|
|
||||||
with patch('builtins.super', return_value=MockSuperWithError()):
|
mock_super = MockSuperWithError()
|
||||||
et.__del__()
|
mock_super._should_raise = True
|
||||||
|
try:
|
||||||
|
with patch('builtins.super', return_value=mock_super):
|
||||||
|
et.__del__()
|
||||||
|
finally:
|
||||||
|
# Prevent the exception from being raised during garbage collection
|
||||||
|
mock_super._should_raise = False
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_execute_update_success(
|
def test_execute_update_success(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test _execute_update executes query successfully."""
|
"""Test _execute_update executes query successfully."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
@@ -169,6 +178,7 @@ def test_execute_update_success(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_connection = MagicMock()
|
mock_connection = MagicMock()
|
||||||
@@ -185,9 +195,8 @@ def test_execute_update_success(
|
|||||||
mock_connection.execute.assert_called_once()
|
mock_connection.execute.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_status_success(
|
def test_update_experiment_run_status_success(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with STATUS update type."""
|
"""Test update_experiment_run with STATUS update type."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -202,6 +211,7 @@ def test_update_experiment_run_status_success(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_execute = MagicMock()
|
mock_execute = MagicMock()
|
||||||
@@ -230,9 +240,8 @@ def test_update_experiment_run_status_success(
|
|||||||
et.info.assert_called_once()
|
et.info.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_status_missing_status(
|
def test_update_experiment_run_status_missing_status(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with STATUS but missing status parameter."""
|
"""Test update_experiment_run with STATUS but missing status parameter."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -247,6 +256,7 @@ def test_update_experiment_run_status_missing_status(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.send_notification = MagicMock()
|
et.send_notification = MagicMock()
|
||||||
@@ -263,9 +273,8 @@ def test_update_experiment_run_status_missing_status(
|
|||||||
et.send_notification.assert_called_once()
|
et.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_status_with_error_success(
|
def test_update_experiment_run_status_with_error_success(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with STATUS_WITH_ERROR update type."""
|
"""Test update_experiment_run with STATUS_WITH_ERROR update type."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -280,6 +289,7 @@ def test_update_experiment_run_status_with_error_success(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_execute = MagicMock()
|
mock_execute = MagicMock()
|
||||||
@@ -309,9 +319,8 @@ def test_update_experiment_run_status_with_error_success(
|
|||||||
et.info.assert_called_once()
|
et.info.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_status_with_error_truncate_message(
|
def test_update_experiment_run_status_with_error_truncate_message(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run truncates error message if too long."""
|
"""Test update_experiment_run truncates error message if too long."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -326,6 +335,7 @@ def test_update_experiment_run_status_with_error_truncate_message(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_execute = MagicMock()
|
mock_execute = MagicMock()
|
||||||
@@ -352,9 +362,8 @@ def test_update_experiment_run_status_with_error_truncate_message(
|
|||||||
assert len(call_args[0][1]['error_message']) == 1024
|
assert len(call_args[0][1]['error_message']) == 1024
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_status_with_error_missing_error_message(
|
def test_update_experiment_run_status_with_error_missing_error_message(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with STATUS_WITH_ERROR but missing error_message."""
|
"""Test update_experiment_run with STATUS_WITH_ERROR but missing error_message."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -369,6 +378,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.send_notification = MagicMock()
|
et.send_notification = MagicMock()
|
||||||
@@ -386,9 +396,8 @@ def test_update_experiment_run_status_with_error_missing_error_message(
|
|||||||
et.send_notification.assert_called_once()
|
et.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_model_saved_success(
|
def test_update_experiment_run_model_saved_success(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with MODEL_SAVED update type."""
|
"""Test update_experiment_run with MODEL_SAVED update type."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -403,6 +412,7 @@ def test_update_experiment_run_model_saved_success(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_execute = MagicMock()
|
mock_execute = MagicMock()
|
||||||
@@ -432,9 +442,8 @@ def test_update_experiment_run_model_saved_success(
|
|||||||
et.info.assert_called_once()
|
et.info.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_model_saved_missing_run_name(
|
def test_update_experiment_run_model_saved_missing_run_name(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with MODEL_SAVED but missing run_name."""
|
"""Test update_experiment_run with MODEL_SAVED but missing run_name."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -449,6 +458,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.send_notification = MagicMock()
|
et.send_notification = MagicMock()
|
||||||
@@ -466,9 +476,8 @@ def test_update_experiment_run_model_saved_missing_run_name(
|
|||||||
et.send_notification.assert_called_once()
|
et.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_invalid_update_type(
|
def test_update_experiment_run_invalid_update_type(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with invalid update_type."""
|
"""Test update_experiment_run with invalid update_type."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
@@ -483,6 +492,7 @@ def test_update_experiment_run_invalid_update_type(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.send_notification = MagicMock()
|
et.send_notification = MagicMock()
|
||||||
@@ -499,9 +509,8 @@ def test_update_experiment_run_invalid_update_type(
|
|||||||
et.send_notification.assert_called_once()
|
et.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_no_rows_updated(
|
def test_update_experiment_run_no_rows_updated(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run raises error when no rows are updated."""
|
"""Test update_experiment_run raises error when no rows are updated."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -516,6 +525,7 @@ def test_update_experiment_run_no_rows_updated(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def mock_execute_update(*args, **kwargs):
|
async def mock_execute_update(*args, **kwargs):
|
||||||
@@ -537,9 +547,8 @@ def test_update_experiment_run_no_rows_updated(
|
|||||||
et.send_notification.assert_called_once()
|
et.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_status_with_error_missing_status(
|
def test_update_experiment_run_status_with_error_missing_status(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with STATUS_WITH_ERROR but missing status - covers line 179."""
|
"""Test update_experiment_run with STATUS_WITH_ERROR but missing status - covers line 179."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -554,6 +563,7 @@ def test_update_experiment_run_status_with_error_missing_status(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.send_notification = MagicMock()
|
et.send_notification = MagicMock()
|
||||||
@@ -571,9 +581,8 @@ def test_update_experiment_run_status_with_error_missing_status(
|
|||||||
et.send_notification.assert_called_once()
|
et.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_update_experiment_run_model_saved_missing_status(
|
def test_update_experiment_run_model_saved_missing_status(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test update_experiment_run with MODEL_SAVED but missing status - covers line 204."""
|
"""Test update_experiment_run with MODEL_SAVED but missing status - covers line 204."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||||
@@ -588,6 +597,7 @@ def test_update_experiment_run_model_saved_missing_status(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.send_notification = MagicMock()
|
et.send_notification = MagicMock()
|
||||||
@@ -605,9 +615,8 @@ def test_update_experiment_run_model_saved_missing_status(
|
|||||||
et.send_notification.assert_called_once()
|
et.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__', return_value=None)
|
|
||||||
def test_experiment_tracking_del_with_engine_no_super_del(
|
def test_experiment_tracking_del_with_engine_no_super_del(
|
||||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||||
):
|
):
|
||||||
"""Test __del__ when engine exists but super has no __del__ - covers line 103."""
|
"""Test __del__ when engine exists but super has no __del__ - covers line 103."""
|
||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
@@ -622,6 +631,7 @@ def test_experiment_tracking_del_with_engine_no_super_del(
|
|||||||
max_connections=db_config['max_connections'],
|
max_connections=db_config['max_connections'],
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
et.engine = MagicMock()
|
et.engine = MagicMock()
|
||||||
|
|||||||
@@ -19,6 +19,24 @@ def mock_notification_handler():
|
|||||||
return MagicMock()
|
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
|
@pytest.fixture
|
||||||
def mock_model_repository():
|
def mock_model_repository():
|
||||||
"""Create a mock model repository."""
|
"""Create a mock model repository."""
|
||||||
@@ -44,15 +62,14 @@ def mock_train_params():
|
|||||||
return params
|
return params
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
def test_training_init(
|
def test_training_init(
|
||||||
mock_training_repo,
|
mock_training_repo,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test Training initialization."""
|
"""Test Training initialization."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -62,28 +79,26 @@ def test_training_init(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
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_training_repo.assert_called_once_with(mock_logger)
|
mock_training_repo.assert_called_once_with(mock_logger)
|
||||||
assert training.model_repository is mock_model_repository
|
assert training.model_repository is mock_model_repository
|
||||||
assert training.storage_repository is mock_storage_repository
|
assert training.storage_repository is mock_storage_repository
|
||||||
|
assert training.metrics_controller is mock_metrics_controller
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
@patch('model_manager.activities.training.TrainModelParams')
|
@patch('model_manager.activities.training.TrainModelParams')
|
||||||
def test_validate_train_params_success(
|
def test_validate_train_params_success(
|
||||||
mock_train_params_class,
|
mock_train_params_class,
|
||||||
mock_training_repo,
|
mock_training_repo,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
mock_train_params,
|
mock_train_params,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test validate_train_params with valid parameters."""
|
"""Test validate_train_params with valid parameters."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -93,6 +108,7 @@ def test_validate_train_params_success(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
training.info = MagicMock()
|
training.info = MagicMock()
|
||||||
@@ -112,17 +128,16 @@ def test_validate_train_params_success(
|
|||||||
training.info.assert_called_once()
|
training.info.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
@patch('model_manager.activities.training.TrainModelParams')
|
@patch('model_manager.activities.training.TrainModelParams')
|
||||||
def test_validate_train_params_value_error(
|
def test_validate_train_params_value_error(
|
||||||
mock_train_params_class,
|
mock_train_params_class,
|
||||||
mock_training_repo,
|
mock_training_repo,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test validate_train_params with ValueError."""
|
"""Test validate_train_params with ValueError."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -132,6 +147,7 @@ def test_validate_train_params_value_error(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
training.send_notification = MagicMock()
|
training.send_notification = MagicMock()
|
||||||
@@ -148,17 +164,16 @@ def test_validate_train_params_value_error(
|
|||||||
training.send_notification.assert_called_once()
|
training.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
@patch('model_manager.activities.training.TrainModelParams')
|
@patch('model_manager.activities.training.TrainModelParams')
|
||||||
def test_validate_train_params_type_error(
|
def test_validate_train_params_type_error(
|
||||||
mock_train_params_class,
|
mock_train_params_class,
|
||||||
mock_training_repo,
|
mock_training_repo,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test validate_train_params with TypeError."""
|
"""Test validate_train_params with TypeError."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -168,6 +183,7 @@ def test_validate_train_params_type_error(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
training.send_notification = MagicMock()
|
training.send_notification = MagicMock()
|
||||||
@@ -184,17 +200,16 @@ def test_validate_train_params_type_error(
|
|||||||
training.send_notification.assert_called_once()
|
training.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
@patch('model_manager.activities.training.TrainModelParams')
|
@patch('model_manager.activities.training.TrainModelParams')
|
||||||
def test_validate_train_params_key_error(
|
def test_validate_train_params_key_error(
|
||||||
mock_train_params_class,
|
mock_train_params_class,
|
||||||
mock_training_repo,
|
mock_training_repo,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test validate_train_params with KeyError."""
|
"""Test validate_train_params with KeyError."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -204,6 +219,7 @@ def test_validate_train_params_key_error(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
training.send_notification = MagicMock()
|
training.send_notification = MagicMock()
|
||||||
@@ -220,16 +236,15 @@ def test_validate_train_params_key_error(
|
|||||||
training.send_notification.assert_called_once()
|
training.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
def test_train_model_success_with_params_object(
|
def test_train_model_success_with_params_object(
|
||||||
mock_training_repo_class,
|
mock_training_repo_class,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
mock_train_params,
|
mock_train_params,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test train_model with TrainModelParams object."""
|
"""Test train_model with TrainModelParams object."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -239,6 +254,7 @@ def test_train_model_success_with_params_object(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_file = BytesIO(b'test data')
|
mock_file = BytesIO(b'test data')
|
||||||
@@ -268,18 +284,17 @@ def test_train_model_success_with_params_object(
|
|||||||
mock_model_repository.save_model.assert_called_once_with(mock_train_result)
|
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.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
@patch('model_manager.activities.training.TrainModelParams')
|
@patch('model_manager.activities.training.TrainModelParams')
|
||||||
def test_train_model_success_with_params_dict(
|
def test_train_model_success_with_params_dict(
|
||||||
mock_train_params_class,
|
mock_train_params_class,
|
||||||
mock_training_repo_class,
|
mock_training_repo_class,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
mock_train_params,
|
mock_train_params,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test train_model with dict parameters."""
|
"""Test train_model with dict parameters."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -289,6 +304,7 @@ def test_train_model_success_with_params_dict(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_train_params_class.from_dict.return_value = mock_train_params
|
mock_train_params_class.from_dict.return_value = mock_train_params
|
||||||
@@ -314,16 +330,15 @@ def test_train_model_success_with_params_dict(
|
|||||||
mock_train_params_class.from_dict.assert_called_once()
|
mock_train_params_class.from_dict.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
def test_train_model_training_fails(
|
def test_train_model_training_fails(
|
||||||
mock_training_repo_class,
|
mock_training_repo_class,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
mock_train_params,
|
mock_train_params,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test train_model when training fails."""
|
"""Test train_model when training fails."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -334,6 +349,7 @@ def test_train_model_training_fails(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
training.send_notification = MagicMock()
|
training.send_notification = MagicMock()
|
||||||
@@ -354,16 +370,15 @@ def test_train_model_training_fails(
|
|||||||
training.send_notification.assert_called_once()
|
training.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
def test_train_model_save_fails(
|
def test_train_model_save_fails(
|
||||||
mock_training_repo_class,
|
mock_training_repo_class,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
mock_train_params,
|
mock_train_params,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test train_model when model saving fails."""
|
"""Test train_model when model saving fails."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -374,6 +389,7 @@ def test_train_model_save_fails(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
training.send_notification = MagicMock()
|
training.send_notification = MagicMock()
|
||||||
@@ -398,15 +414,14 @@ def test_train_model_save_fails(
|
|||||||
training.send_notification.assert_called_once()
|
training.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
def test_cleanup_resources_success(
|
def test_cleanup_resources_success(
|
||||||
mock_training_repo_class,
|
mock_training_repo_class,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test cleanup_resources successfully."""
|
"""Test cleanup_resources successfully."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -416,6 +431,7 @@ def test_cleanup_resources_success(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
input_data = {
|
input_data = {
|
||||||
@@ -431,15 +447,14 @@ def test_cleanup_resources_success(
|
|||||||
mock_storage_repository.delete_file.assert_called_once_with('test-bucket', 'test-file.csv')
|
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.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
def test_cleanup_resources_cleanup_fails(
|
def test_cleanup_resources_cleanup_fails(
|
||||||
mock_training_repo_class,
|
mock_training_repo_class,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test cleanup_resources when cleanup fails."""
|
"""Test cleanup_resources when cleanup fails."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -449,6 +464,7 @@ def test_cleanup_resources_cleanup_fails(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
training.send_notification = MagicMock()
|
training.send_notification = MagicMock()
|
||||||
@@ -467,15 +483,14 @@ def test_cleanup_resources_cleanup_fails(
|
|||||||
training.send_notification.assert_called_once()
|
training.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
|
||||||
@patch('model_manager.activities.training.TrainingRepository')
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
def test_cleanup_resources_with_empty_values(
|
def test_cleanup_resources_with_empty_values(
|
||||||
mock_training_repo_class,
|
mock_training_repo_class,
|
||||||
mock_base_init,
|
|
||||||
mock_model_repository,
|
mock_model_repository,
|
||||||
mock_storage_repository,
|
mock_storage_repository,
|
||||||
mock_logger,
|
mock_logger,
|
||||||
mock_notification_handler,
|
mock_notification_handler,
|
||||||
|
mock_metrics_controller,
|
||||||
):
|
):
|
||||||
"""Test cleanup_resources with empty values."""
|
"""Test cleanup_resources with empty values."""
|
||||||
from model_manager.activities.training import Training
|
from model_manager.activities.training import Training
|
||||||
@@ -485,6 +500,7 @@ def test_cleanup_resources_with_empty_values(
|
|||||||
storage_repository=mock_storage_repository,
|
storage_repository=mock_storage_repository,
|
||||||
logger=mock_logger,
|
logger=mock_logger,
|
||||||
notification_handler=mock_notification_handler,
|
notification_handler=mock_notification_handler,
|
||||||
|
metrics_controller=mock_metrics_controller,
|
||||||
)
|
)
|
||||||
|
|
||||||
input_data = {
|
input_data = {
|
||||||
|
|||||||
Reference in New Issue
Block a user