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:
Bruno Domingues
2025-11-03 15:26:58 -03:00
committed by GitHub
8 changed files with 210 additions and 94 deletions

View File

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

View File

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

View File

@@ -15,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,
},
)

View File

@@ -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
)

View File

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

View File

@@ -296,6 +296,9 @@ def test_activities_del_with_engine_exception_caught(
class MockSuperWithError:
def __del__(self):
# 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
@@ -303,5 +306,11 @@ def test_activities_del_with_engine_exception_caught(
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
with patch('builtins.super', return_value=MockSuperWithError()):
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

View File

@@ -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,12 +135,16 @@ 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):
# 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
@@ -148,13 +152,18 @@ def test_experiment_tracking_del_with_engine_exception(
warnings.filterwarnings('ignore', category=pytest.PytestUnraisableExceptionWarning)
with patch('builtins.super', return_value=MockSuperWithError()):
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()

View File

@@ -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 = {