SIENTIAPDE-1250: Implement experiment tracking activity with status updates, error handling, and model registration. This includes a unified update method, error message truncation, and integration into the main activities orchestrator. (236+, 2-)
This commit is contained in:
@@ -1,18 +1,20 @@
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
from pytest import mark
|
||||
from sientia_do.temporal.activities.postgres import Postgres
|
||||
|
||||
from model_manager.activities.activities import Activities
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
from model_manager.activities.gates import Gates
|
||||
from model_manager.activities.mlflow import MLFlow
|
||||
|
||||
|
||||
@patch('model_manager.activities.activities.Postgres.__init__')
|
||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
|
||||
@patch('model_manager.activities.activities.MLFlow.__init__')
|
||||
@patch('model_manager.activities.activities.MinIO.__init__')
|
||||
@patch('model_manager.activities.activities.Gates.__init__')
|
||||
def test___init__(mock_gates_init, mock_minio_init, mock_mlflow_init, mock_postgres_init):
|
||||
def test___init__(
|
||||
mock_gates_init, mock_minio_init, mock_mlflow_init, mock_experiment_tracking_init
|
||||
):
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
@@ -49,11 +51,11 @@ def test___init__(mock_gates_init, mock_minio_init, mock_mlflow_init, mock_postg
|
||||
)
|
||||
|
||||
assert isinstance(activities, Activities)
|
||||
assert isinstance(activities, Postgres)
|
||||
assert isinstance(activities, ExperimentTracking)
|
||||
assert isinstance(activities, MLFlow)
|
||||
assert isinstance(activities, Gates)
|
||||
|
||||
mock_postgres_init.assert_called_once_with(
|
||||
mock_experiment_tracking_init.assert_called_once_with(
|
||||
ANY,
|
||||
host=postgres_config['host'],
|
||||
port=postgres_config['port'],
|
||||
@@ -97,9 +99,9 @@ def test___init__(mock_gates_init, mock_minio_init, mock_mlflow_init, mock_postg
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.activities.Postgres', return_value=MagicMock())
|
||||
@patch('model_manager.activities.activities.ExperimentTracking', return_value=MagicMock())
|
||||
@patch('model_manager.activities.activities.MLFlow', return_value=MagicMock())
|
||||
async def test_shutdown(_mock_mlflow_init, mock_postgres_init):
|
||||
async def test_shutdown(_mock_mlflow_init, mock_experiment_tracking_init):
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
@@ -136,4 +138,4 @@ async def test_shutdown(_mock_mlflow_init, mock_postgres_init):
|
||||
)
|
||||
|
||||
await activities.shutdown()
|
||||
mock_postgres_init.close.assert_called_once()
|
||||
mock_experiment_tracking_init.close.assert_called_once()
|
||||
|
||||
385
tests/activities/test_experiment_tracking.py
Normal file
385
tests/activities/test_experiment_tracking.py
Normal file
@@ -0,0 +1,385 @@
|
||||
"""Unit tests for ExperimentTracking activity."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from pytest import mark, raises
|
||||
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_status_success(mock_postgres_init):
|
||||
"""Test successful status update."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
# Create instance
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
# Mock methods
|
||||
tracking.info = MagicMock()
|
||||
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
|
||||
|
||||
# Test data
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.STATUS,
|
||||
'status': 'TRAINING_SUCCESS',
|
||||
}
|
||||
|
||||
# Execute
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
# Assertions
|
||||
tracking.info.assert_called()
|
||||
tracking.execute_query.assert_called_once()
|
||||
call_args = tracking.execute_query.call_args
|
||||
assert 'UPDATE experiment_run' in call_args[0][0]
|
||||
assert 'SET status = %s' in call_args[0][0]
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_status_with_error_success(mock_postgres_init):
|
||||
"""Test successful status update with error message."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.info = MagicMock()
|
||||
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.STATUS_WITH_ERROR,
|
||||
'status': 'TRAINING_ERROR',
|
||||
'error_message': 'Model training failed due to insufficient data',
|
||||
}
|
||||
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
tracking.info.assert_called()
|
||||
tracking.execute_query.assert_called_once()
|
||||
call_args = tracking.execute_query.call_args
|
||||
assert 'UPDATE experiment_run' in call_args[0][0]
|
||||
assert 'SET status = %s, error_message = %s' in call_args[0][0]
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_error_message_truncation(mock_postgres_init):
|
||||
"""Test that error messages longer than 1024 chars are truncated."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.info = MagicMock()
|
||||
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
|
||||
|
||||
# Create error message longer than 1024 characters
|
||||
long_error = 'A' * 2000
|
||||
|
||||
input_data = {
|
||||
'metadata': {},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.STATUS_WITH_ERROR,
|
||||
'status': 'TRAINING_ERROR',
|
||||
'error_message': long_error,
|
||||
}
|
||||
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
# Check that error message was truncated to 1024 chars
|
||||
call_args = tracking.execute_query.call_args
|
||||
query_params = call_args[0][1]
|
||||
assert len(query_params[1]) == 1024 # error_message is second param
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_model_saved_success(mock_postgres_init):
|
||||
"""Test successful model saved update."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.info = MagicMock()
|
||||
tracking.execute_query = AsyncMock(return_value={'rowcount': 1})
|
||||
|
||||
input_data = {
|
||||
'metadata': {},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.MODEL_SAVED,
|
||||
'run_name': 'experiment-model-123',
|
||||
'status': 'MLFLOW_SENT',
|
||||
}
|
||||
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
tracking.info.assert_called()
|
||||
tracking.execute_query.assert_called_once()
|
||||
call_args = tracking.execute_query.call_args
|
||||
assert 'UPDATE experiment_run' in call_args[0][0]
|
||||
assert 'SET run_name = %s, status = %s' in call_args[0][0]
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_missing_status_raises_error(mock_postgres_init):
|
||||
"""Test that missing status parameter raises ValueError."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.send_notification = MagicMock()
|
||||
tracking.error = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.STATUS,
|
||||
# Missing 'status' parameter
|
||||
}
|
||||
|
||||
with raises(RuntimeError, match='Error updating experiment run'):
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
tracking.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_missing_error_message_raises_error(mock_postgres_init):
|
||||
"""Test that missing error_message parameter raises ValueError."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.send_notification = MagicMock()
|
||||
tracking.error = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.STATUS_WITH_ERROR,
|
||||
'status': 'TRAINING_ERROR',
|
||||
# Missing 'error_message' parameter
|
||||
}
|
||||
|
||||
with raises(RuntimeError, match='Error updating experiment run'):
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
tracking.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_missing_run_name_raises_error(mock_postgres_init):
|
||||
"""Test that missing run_name parameter raises ValueError."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.send_notification = MagicMock()
|
||||
tracking.error = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.MODEL_SAVED,
|
||||
# Missing 'run_name' parameter
|
||||
}
|
||||
|
||||
with raises(RuntimeError, match='Error updating experiment run'):
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
tracking.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_invalid_update_type_raises_error(mock_postgres_init):
|
||||
"""Test that invalid update_type raises ValueError."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.send_notification = MagicMock()
|
||||
tracking.error = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': 'INVALID_TYPE',
|
||||
'status': 'TRAINING_SUCCESS',
|
||||
}
|
||||
|
||||
with raises(RuntimeError, match='Error updating experiment run'):
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
tracking.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_no_rows_updated_raises_error(mock_postgres_init):
|
||||
"""Test that zero rows updated raises ValueError."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.info = MagicMock()
|
||||
tracking.execute_query = AsyncMock(return_value={'rowcount': 0})
|
||||
tracking.send_notification = MagicMock()
|
||||
tracking.error = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {},
|
||||
'experiment_run_id': 999, # Non-existent ID
|
||||
'update_type': UpdateType.STATUS,
|
||||
'status': 'TRAINING_SUCCESS',
|
||||
}
|
||||
|
||||
with raises(RuntimeError, match='Error updating experiment run'):
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
tracking.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('model_manager.activities.experiment_tracking.Postgres.__init__')
|
||||
async def test_update_experiment_run_sends_notification_on_error(mock_postgres_init):
|
||||
"""Test that notification is sent when update fails."""
|
||||
mock_postgres_init.return_value = None
|
||||
|
||||
tracking = ExperimentTracking(
|
||||
host='localhost',
|
||||
port=5432,
|
||||
user='test',
|
||||
password='test',
|
||||
dbname='test',
|
||||
min_connections=1,
|
||||
max_connections=10,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
)
|
||||
|
||||
tracking.info = MagicMock()
|
||||
tracking.execute_query = AsyncMock(side_effect=Exception('Database error'))
|
||||
tracking.send_notification = MagicMock()
|
||||
tracking.error = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
'experiment_run_id': 456,
|
||||
'update_type': UpdateType.STATUS,
|
||||
'status': 'TRAINING_SUCCESS',
|
||||
}
|
||||
|
||||
with raises(RuntimeError):
|
||||
await tracking.update_experiment_run(input_data)
|
||||
|
||||
# Verify notification was sent
|
||||
tracking.send_notification.assert_called_once()
|
||||
call_args = tracking.send_notification.call_args
|
||||
assert call_args[1]['notification_id'] == 'UPDATE_EXPERIMENT_RUN_ERROR'
|
||||
assert call_args[1]['metadata'] == {'workflow_id': 'test-123'}
|
||||
|
||||
|
||||
def test_update_type_enum_values():
|
||||
"""Test UpdateType enum has correct values."""
|
||||
assert UpdateType.STATUS == 'status'
|
||||
assert UpdateType.STATUS_WITH_ERROR == 'status_with_error'
|
||||
assert UpdateType.MODEL_SAVED == 'model_saved'
|
||||
Reference in New Issue
Block a user