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