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.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,8 +15,10 @@ 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.metrics import ACTIVITY_EXECUTION_TOTAL, WORKFLOW_EXECUTION_TOTAL
|
||||
from model_manager.utils.exceptions import ModelTrainingError
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
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
|
||||
|
||||
|
||||
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 +42,7 @@ class Training(BaseActivity):
|
||||
storage_repository: StorageRepository,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
"""
|
||||
Initialize Training activity.
|
||||
@@ -52,7 +51,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
|
||||
@@ -78,6 +77,7 @@ class Training(BaseActivity):
|
||||
ValueError, TypeError, KeyError: If validation fails (after sending notification)
|
||||
"""
|
||||
metadata = input_data.get('metadata', {})
|
||||
metrics_status = 'success'
|
||||
|
||||
try:
|
||||
train_params = TrainModelParams.from_dict(input_data)
|
||||
@@ -92,6 +92,7 @@ class Training(BaseActivity):
|
||||
|
||||
return train_params
|
||||
except (ValueError, TypeError, KeyError) as e:
|
||||
metrics_status = 'error'
|
||||
error_msg = f'Error validating training parameters: {str(e)}'
|
||||
trace = traceback.format_exc()
|
||||
|
||||
@@ -104,6 +105,13 @@ class Training(BaseActivity):
|
||||
attachment_content=trace,
|
||||
)
|
||||
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')
|
||||
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_saved = False
|
||||
metrics_status = 'success'
|
||||
|
||||
try:
|
||||
with self.storage_repository.fetch_file(
|
||||
@@ -157,6 +166,8 @@ class Training(BaseActivity):
|
||||
'run_dir': train_result.run_dir,
|
||||
}
|
||||
except Exception as e: # noqa: BLE001
|
||||
metrics_status = 'error'
|
||||
|
||||
error_msg = (
|
||||
'Error training model - '
|
||||
f'model_trained={model_trained}, model_saved={model_saved}, '
|
||||
@@ -178,6 +189,13 @@ class Training(BaseActivity):
|
||||
model_trained=model_trained,
|
||||
model_saved=model_saved,
|
||||
) 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')
|
||||
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', '')
|
||||
bucket_name = input_data.get('bucket_name', '')
|
||||
file_name = input_data.get('file_name', '')
|
||||
metrics_status = 'success'
|
||||
|
||||
try:
|
||||
self.model_repository.cleanup_run_directory(run_dir)
|
||||
self.storage_repository.delete_file(bucket_name, file_name)
|
||||
except Exception as e: # noqa: BLE001
|
||||
metrics_status = 'error'
|
||||
|
||||
error_msg = (
|
||||
f'Error cleaning up resources - Run directory: {run_dir}, '
|
||||
f'File: {bucket_name}/{file_name}, Error: {str(e)}'
|
||||
@@ -218,4 +239,46 @@ class Training(BaseActivity):
|
||||
level=NotificationLevel.ERROR,
|
||||
attachment_content=trace,
|
||||
)
|
||||
|
||||
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
|
||||
"""
|
||||
|
||||
from prometheus_client import Gauge
|
||||
from prometheus_client import Counter, Gauge
|
||||
|
||||
# Application health metric
|
||||
APP_UP = Gauge(
|
||||
@@ -24,3 +24,15 @@ APP_UP = Gauge(
|
||||
'Indicates if the application is running (1) or shutting down (0)',
|
||||
['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
|
||||
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
|
||||
|
||||
@@ -296,12 +296,21 @@ def test_activities_del_with_engine_exception_caught(
|
||||
|
||||
class MockSuperWithError:
|
||||
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
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
||||
|
||||
with patch('builtins.super', return_value=MockSuperWithError()):
|
||||
activities.__del__()
|
||||
mock_super = MockSuperWithError()
|
||||
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()
|
||||
|
||||
|
||||
@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
|
||||
def db_config():
|
||||
"""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(
|
||||
mock_postgres_init, db_config, mock_logger, mock_notification_handler
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test ExperimentTracking initialization."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
@@ -49,24 +61,12 @@ def test_experiment_tracking_init(
|
||||
max_connections=db_config['max_connections'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
)
|
||||
|
||||
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,
|
||||
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
|
||||
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 +81,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'):
|
||||
@@ -89,9 +90,8 @@ def test_experiment_tracking_del_without_engine(
|
||||
et.__del__()
|
||||
|
||||
|
||||
@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
|
||||
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 +106,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()
|
||||
@@ -118,9 +119,8 @@ def test_experiment_tracking_del_with_engine(
|
||||
et.__del__()
|
||||
|
||||
|
||||
@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
|
||||
db_config, mock_logger, mock_notification_handler, mock_metrics_controller
|
||||
):
|
||||
"""Test __del__ catches exceptions."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.engine = MagicMock()
|
||||
|
||||
class MockSuperWithError:
|
||||
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
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
|
||||
|
||||
with patch('builtins.super', return_value=MockSuperWithError()):
|
||||
et.__del__()
|
||||
mock_super = MockSuperWithError()
|
||||
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(
|
||||
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."""
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
@@ -169,6 +178,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()
|
||||
@@ -185,9 +195,8 @@ def test_execute_update_success(
|
||||
mock_connection.execute.assert_called_once()
|
||||
|
||||
|
||||
@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
|
||||
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 +211,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()
|
||||
@@ -230,9 +240,8 @@ def test_update_experiment_run_status_success(
|
||||
et.info.assert_called_once()
|
||||
|
||||
|
||||
@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
|
||||
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 +256,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()
|
||||
@@ -263,9 +273,8 @@ def test_update_experiment_run_status_missing_status(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
mock_execute = MagicMock()
|
||||
@@ -309,9 +319,8 @@ def test_update_experiment_run_status_with_error_success(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
|
||||
@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
|
||||
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 +378,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()
|
||||
@@ -386,9 +396,8 @@ def test_update_experiment_run_status_with_error_missing_error_message(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
mock_execute = MagicMock()
|
||||
@@ -432,9 +442,8 @@ def test_update_experiment_run_model_saved_success(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
@@ -466,9 +476,8 @@ def test_update_experiment_run_model_saved_missing_run_name(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
@@ -499,9 +509,8 @@ def test_update_experiment_run_invalid_update_type(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
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()
|
||||
|
||||
|
||||
@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
|
||||
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 +563,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()
|
||||
@@ -571,9 +581,8 @@ def test_update_experiment_run_status_with_error_missing_status(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
@@ -605,9 +615,8 @@ def test_update_experiment_run_model_saved_missing_status(
|
||||
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(
|
||||
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."""
|
||||
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'],
|
||||
logger=mock_logger,
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.engine = MagicMock()
|
||||
|
||||
@@ -19,6 +19,24 @@ def mock_notification_handler():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_metrics_controller():
|
||||
"""Create a mock metrics controller."""
|
||||
controller = MagicMock()
|
||||
|
||||
# Make shutdown an async coroutine
|
||||
async def mock_shutdown():
|
||||
pass
|
||||
|
||||
# Make emit an async coroutine
|
||||
async def mock_emit(*args, **kwargs):
|
||||
pass
|
||||
|
||||
controller.shutdown = mock_shutdown
|
||||
controller.emit = mock_emit
|
||||
return controller
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_model_repository():
|
||||
"""Create a mock model repository."""
|
||||
@@ -44,15 +62,14 @@ def mock_train_params():
|
||||
return params
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
||||
@patch('model_manager.activities.training.TrainingRepository')
|
||||
def test_training_init(
|
||||
mock_training_repo,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test Training initialization."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -62,28 +79,26 @@ 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_training_repo.assert_called_once_with(mock_logger)
|
||||
assert training.model_repository is mock_model_repository
|
||||
assert training.storage_repository is mock_storage_repository
|
||||
assert training.metrics_controller is mock_metrics_controller
|
||||
|
||||
|
||||
@patch('model_manager.activities.training.BaseActivity.__init__', return_value=None)
|
||||
@patch('model_manager.activities.training.TrainingRepository')
|
||||
@patch('model_manager.activities.training.TrainModelParams')
|
||||
def test_validate_train_params_success(
|
||||
mock_train_params_class,
|
||||
mock_training_repo,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_train_params,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test validate_train_params with valid parameters."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -93,6 +108,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,17 +128,16 @@ 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.TrainingRepository')
|
||||
@patch('model_manager.activities.training.TrainModelParams')
|
||||
def test_validate_train_params_value_error(
|
||||
mock_train_params_class,
|
||||
mock_training_repo,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test validate_train_params with ValueError."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -132,6 +147,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,17 +164,16 @@ 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.TrainingRepository')
|
||||
@patch('model_manager.activities.training.TrainModelParams')
|
||||
def test_validate_train_params_type_error(
|
||||
mock_train_params_class,
|
||||
mock_training_repo,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test validate_train_params with TypeError."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -168,6 +183,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,17 +200,16 @@ 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.TrainingRepository')
|
||||
@patch('model_manager.activities.training.TrainModelParams')
|
||||
def test_validate_train_params_key_error(
|
||||
mock_train_params_class,
|
||||
mock_training_repo,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test validate_train_params with KeyError."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -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,16 +236,15 @@ 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.TrainingRepository')
|
||||
def test_train_model_success_with_params_object(
|
||||
mock_training_repo_class,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_train_params,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test train_model with TrainModelParams object."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -239,6 +254,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,18 +284,17 @@ 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.TrainingRepository')
|
||||
@patch('model_manager.activities.training.TrainModelParams')
|
||||
def test_train_model_success_with_params_dict(
|
||||
mock_train_params_class,
|
||||
mock_training_repo_class,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_train_params,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test train_model with dict parameters."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -289,6 +304,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,16 +330,15 @@ 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.TrainingRepository')
|
||||
def test_train_model_training_fails(
|
||||
mock_training_repo_class,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_train_params,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test train_model when training fails."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -334,6 +349,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,16 +370,15 @@ 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.TrainingRepository')
|
||||
def test_train_model_save_fails(
|
||||
mock_training_repo_class,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_train_params,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test train_model when model saving fails."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -374,6 +389,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,15 +414,14 @@ 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.TrainingRepository')
|
||||
def test_cleanup_resources_success(
|
||||
mock_training_repo_class,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test cleanup_resources successfully."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -416,6 +431,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,15 +447,14 @@ 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.TrainingRepository')
|
||||
def test_cleanup_resources_cleanup_fails(
|
||||
mock_training_repo_class,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test cleanup_resources when cleanup fails."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -449,6 +464,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,15 +483,14 @@ 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.TrainingRepository')
|
||||
def test_cleanup_resources_with_empty_values(
|
||||
mock_training_repo_class,
|
||||
mock_base_init,
|
||||
mock_model_repository,
|
||||
mock_storage_repository,
|
||||
mock_logger,
|
||||
mock_notification_handler,
|
||||
mock_metrics_controller,
|
||||
):
|
||||
"""Test cleanup_resources with empty values."""
|
||||
from model_manager.activities.training import Training
|
||||
@@ -485,6 +500,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