diff --git a/tests/activities/test_experiment_tracking.py b/tests/activities/test_experiment_tracking.py new file mode 100644 index 0000000..817369c --- /dev/null +++ b/tests/activities/test_experiment_tracking.py @@ -0,0 +1,628 @@ +"""Unit tests for ExperimentTracking class with 100% coverage.""" + +import asyncio +from unittest.mock import MagicMock, patch + +import pytest + + +@pytest.fixture +def mock_logger(): + """Create a mock logger.""" + return MagicMock() + + +@pytest.fixture +def mock_notification_handler(): + """Create a mock notification handler.""" + return MagicMock() + + +@pytest.fixture +def db_config(): + """Create a valid database configuration.""" + return { + 'host': 'localhost', + 'port': 5432, + 'user': 'testuser', + 'password': 'testpass', + 'dbname': 'testdb', + 'min_connections': 1, + 'max_connections': 10, + } + + +@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 +): + """Test ExperimentTracking initialization.""" + from model_manager.activities.experiment_tracking import ExperimentTracking + + ExperimentTracking( + 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, + ) + + 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, + ) + + +@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 +): + """Test __del__ when engine attribute does not exist.""" + from model_manager.activities.experiment_tracking import ExperimentTracking + + et = ExperimentTracking( + 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, + ) + + if hasattr(et, 'engine'): + delattr(et, '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 +): + """Test __del__ when engine exists.""" + from model_manager.activities.experiment_tracking import ExperimentTracking + + et = ExperimentTracking( + 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, + ) + + et.engine = MagicMock() + + class MockSuper: + def __del__(self): + pass + + with patch('builtins.super', return_value=MockSuper()): + 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 +): + """Test __del__ catches exceptions.""" + from model_manager.activities.experiment_tracking import ExperimentTracking + + et = ExperimentTracking( + 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, + ) + + et.engine = MagicMock() + + class MockSuperWithError: + def __del__(self): + raise RuntimeError('Test error') + + with patch('builtins.super', return_value=MockSuperWithError()): + et.__del__() + + +@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 +): + """Test _execute_update executes query successfully.""" + from model_manager.activities.experiment_tracking import ExperimentTracking + + et = ExperimentTracking( + 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, + ) + + mock_connection = MagicMock() + mock_result = MagicMock() + mock_result.rowcount = 1 + mock_connection.execute.return_value = mock_result + mock_engine = MagicMock() + mock_engine.begin.return_value.__enter__.return_value = mock_connection + et.engine = mock_engine + + result = asyncio.run(et._execute_update('UPDATE test SET x = :x', {'x': 1})) + + assert result == {'rowcount': 1} + 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 +): + """Test update_experiment_run with STATUS update type.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + mock_execute = MagicMock() + + async def mock_execute_update(*args, **kwargs): + mock_execute(*args, **kwargs) + return {'rowcount': 1} + + et._execute_update = mock_execute_update + et.info = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.STATUS, + 'status': 'running', + } + + asyncio.run(et.update_experiment_run(input_data)) + + mock_execute.assert_called_once() + call_args = mock_execute.call_args + assert 'status' in call_args[0][1] + assert call_args[0][1]['status'] == 'running' + assert call_args[0][1]['experiment_run_id'] == 1 + 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 +): + """Test update_experiment_run with STATUS but missing status parameter.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + et.send_notification = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.STATUS, + } + + with pytest.raises(RuntimeError): + asyncio.run(et.update_experiment_run(input_data)) + + 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 +): + """Test update_experiment_run with STATUS_WITH_ERROR update type.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + mock_execute = MagicMock() + + async def mock_execute_update(*args, **kwargs): + mock_execute(*args, **kwargs) + return {'rowcount': 1} + + et._execute_update = mock_execute_update + et.info = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.STATUS_WITH_ERROR, + 'status': 'failed', + 'error_message': 'Test error', + } + + asyncio.run(et.update_experiment_run(input_data)) + + mock_execute.assert_called_once() + call_args = mock_execute.call_args + assert 'status' in call_args[0][1] + assert call_args[0][1]['status'] == 'failed' + assert call_args[0][1]['error_message'] == 'Test error' + 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 +): + """Test update_experiment_run truncates error message if too long.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + mock_execute = MagicMock() + + async def mock_execute_update(*args, **kwargs): + mock_execute(*args, **kwargs) + return {'rowcount': 1} + + et._execute_update = mock_execute_update + et.info = MagicMock() + + long_error = 'x' * 2000 + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.STATUS_WITH_ERROR, + 'status': 'failed', + 'error_message': long_error, + } + + asyncio.run(et.update_experiment_run(input_data)) + + call_args = mock_execute.call_args + 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 +): + """Test update_experiment_run with STATUS_WITH_ERROR but missing error_message.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + et.send_notification = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.STATUS_WITH_ERROR, + 'status': 'failed', + } + + with pytest.raises(RuntimeError): + asyncio.run(et.update_experiment_run(input_data)) + + 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 +): + """Test update_experiment_run with MODEL_SAVED update type.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + mock_execute = MagicMock() + + async def mock_execute_update(*args, **kwargs): + mock_execute(*args, **kwargs) + return {'rowcount': 1} + + et._execute_update = mock_execute_update + et.info = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.MODEL_SAVED, + 'status': 'completed', + 'run_name': 'run_001', + } + + asyncio.run(et.update_experiment_run(input_data)) + + mock_execute.assert_called_once() + call_args = mock_execute.call_args + assert 'run_name' in call_args[0][1] + assert call_args[0][1]['run_name'] == 'run_001' + assert call_args[0][1]['status'] == 'completed' + 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 +): + """Test update_experiment_run with MODEL_SAVED but missing run_name.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + et.send_notification = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.MODEL_SAVED, + 'status': 'completed', + } + + with pytest.raises(RuntimeError): + asyncio.run(et.update_experiment_run(input_data)) + + 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 +): + """Test update_experiment_run with invalid update_type.""" + from model_manager.activities.experiment_tracking import ExperimentTracking + + et = ExperimentTracking( + 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, + ) + + et.send_notification = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': 'invalid_type', + } + + with pytest.raises(RuntimeError): + asyncio.run(et.update_experiment_run(input_data)) + + 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 +): + """Test update_experiment_run raises error when no rows are updated.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + async def mock_execute_update(*args, **kwargs): + return {'rowcount': 0} + + et._execute_update = mock_execute_update + et.send_notification = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 999, + 'update_type': UpdateType.STATUS, + 'status': 'running', + } + + with pytest.raises(RuntimeError): + asyncio.run(et.update_experiment_run(input_data)) + + 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 +): + """Test update_experiment_run with STATUS_WITH_ERROR but missing status - covers line 179.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + et.send_notification = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.STATUS_WITH_ERROR, + 'error_message': 'Some error', + } + + with pytest.raises(RuntimeError): + asyncio.run(et.update_experiment_run(input_data)) + + 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 +): + """Test update_experiment_run with MODEL_SAVED but missing status - covers line 204.""" + from model_manager.activities.experiment_tracking import ExperimentTracking, UpdateType + + et = ExperimentTracking( + 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, + ) + + et.send_notification = MagicMock() + + input_data = { + 'metadata': {'workflow_id': 'test-123'}, + 'experiment_run_id': 1, + 'update_type': UpdateType.MODEL_SAVED, + 'run_name': 'run_001', + } + + with pytest.raises(RuntimeError): + asyncio.run(et.update_experiment_run(input_data)) + + 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 +): + """Test __del__ when engine exists but super has no __del__ - covers line 103.""" + from model_manager.activities.experiment_tracking import ExperimentTracking + + et = ExperimentTracking( + 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, + ) + + et.engine = MagicMock() + + class MockSuperNoDel: + pass + + with patch('builtins.super', return_value=MockSuperNoDel()): + et.__del__()