386 lines
12 KiB
Python
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'
|