Files
sientia-dataops-model-manager/tests/activities/test_experiment_tracking.py

386 lines
12 KiB
Python

"""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'