diff --git a/tests/laborious/worker/test_worker.py b/tests/laborious/worker/test_worker.py new file mode 100644 index 0000000..0c43222 --- /dev/null +++ b/tests/laborious/worker/test_worker.py @@ -0,0 +1,499 @@ +from unittest.mock import ANY, AsyncMock, MagicMock, patch + +import pytest + + +@pytest.fixture +def mock_env_vars(monkeypatch): + """Fixture to set up environment variables for tests.""" + monkeypatch.setenv('POD_ID', 'test-pod-123') + monkeypatch.setenv('TEMPORAL_HOST', 'test-temporal:7233') + monkeypatch.setenv('TEMPORAL_NAMESPACE', 'test-namespace') + monkeypatch.setenv('HTTP_METRICS_PORT', '9090') + monkeypatch.setenv('HTTP_SDK_METRICS_PORT', '9091') + monkeypatch.setenv('PROJECT_NAME', 'test-project') + monkeypatch.setenv('POSTGRES_HOST', 'localhost') + monkeypatch.setenv('POSTGRES_PORT', '5432') + monkeypatch.setenv('POSTGRES_USER', 'test') + monkeypatch.setenv('POSTGRES_PASSWORD', 'test') + monkeypatch.setenv('POSTGRES_DBNAME', 'test') + monkeypatch.setenv('MLFLOW_HOST', 'http://localhost') + monkeypatch.setenv('MLFLOW_PORT', '5000') + monkeypatch.setenv('MLFLOW_USERNAME', 'test') + monkeypatch.setenv('MLFLOW_PASSWORD', 'test') + monkeypatch.setenv('MINIO_ENDPOINT_URL', 'http://localhost:9000') + monkeypatch.setenv('MINIO_ACCESS_KEY', 'test') + monkeypatch.setenv('MINIO_SECRET_KEY', 'test') + monkeypatch.setenv('MONGODB_USERNAME', 'test') + monkeypatch.setenv('MONGODB_PASSWORD', 'test') + monkeypatch.setenv('MONGODB_URL', 'localhost:27017') + monkeypatch.setenv('MONGODB_DATABASE_NAME', 'test') + + +@patch('model_manager.worker.worker.start_http_server') +@patch('model_manager.worker.worker.metrics') +def test_start_prometheus_server_success(mock_metrics, mock_start_http_server, mock_env_vars): + """Test successful Prometheus server startup.""" + # Import after patching to ensure mocks are in place + from model_manager.worker.worker import start_prometheus_server + + # Act + start_prometheus_server() + + # Assert + mock_start_http_server.assert_called_once_with(9090) + mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123') + mock_metrics.APP_UP.labels.return_value.set.assert_called_once_with(1) + + +@patch('model_manager.worker.worker.start_http_server') +@patch('model_manager.worker.worker.os._exit') +def test_start_prometheus_server_failure(mock_exit, mock_start_http_server, mock_env_vars): + """Test Prometheus server startup failure.""" + # Arrange + mock_start_http_server.side_effect = Exception('Port already in use') + + # Import after patching + from model_manager.worker.worker import start_prometheus_server + + # Act + start_prometheus_server() + + # Assert + mock_start_http_server.assert_called_once_with(9090) + mock_exit.assert_called_once_with(1) + + +@pytest.mark.asyncio +@patch('model_manager.worker.worker.sys.exit') +@patch('model_manager.worker.worker.metrics') +@patch('model_manager.worker.worker.asyncio.gather') +@patch('model_manager.worker.worker.Worker') +@patch('model_manager.worker.worker.client.Client.connect') +@patch('model_manager.worker.worker.Runtime') +@patch('model_manager.worker.worker.Activities') +@patch('model_manager.worker.worker.NotificationHandler') +@patch('model_manager.worker.worker.start_prometheus_server') +@patch('model_manager.worker.worker.get_logger') +async def test_main_success( + mock_get_logger, + mock_start_prometheus, + mock_notification_handler, + mock_activities, + mock_runtime, + mock_client_connect, + mock_worker, + mock_gather, + mock_metrics, + mock_sys_exit, + mock_env_vars, +): + """Test successful main function execution.""" + # Arrange + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + mock_handler = MagicMock() + mock_notification_handler.return_value = mock_handler + + mock_activities_instance = MagicMock() + mock_activities_instance.shutdown = AsyncMock() + mock_activities.return_value = mock_activities_instance + + mock_temporal_client = AsyncMock() + mock_client_connect.return_value = mock_temporal_client + + mock_worker_instance = MagicMock() + mock_worker_instance.run = MagicMock(return_value=AsyncMock()) + mock_worker.return_value = mock_worker_instance + + # Mock gather to complete successfully + mock_gather.return_value = None + + # Import and run + from model_manager.worker.worker import main + + # Act + await main() + + # Assert + mock_start_prometheus.assert_called_once() + mock_notification_handler.assert_called_once() + mock_activities.assert_called_once() + mock_client_connect.assert_called_once_with( + target_host='test-temporal:7233', + namespace='test-namespace', + runtime=ANY, + ) + assert mock_worker.call_count == 2 # Two workers created + mock_gather.assert_called_once() + mock_sys_exit.assert_called_once_with(1) + + +@pytest.mark.asyncio +@patch('model_manager.worker.worker.asyncio.gather') +@patch('model_manager.worker.worker.Worker') +@patch('model_manager.worker.worker.client.Client.connect') +@patch('model_manager.worker.worker.Runtime') +@patch('model_manager.worker.worker.Activities') +@patch('model_manager.worker.worker.NotificationHandler') +@patch('model_manager.worker.worker.start_prometheus_server') +@patch('model_manager.worker.worker.get_logger') +@patch('model_manager.worker.worker.sys.exit') +@patch('model_manager.worker.worker.metrics') +async def test_main_exception_handling( + mock_metrics, + mock_sys_exit, + mock_get_logger, + mock_start_prometheus, + mock_notification_handler, + mock_activities, + mock_runtime, + mock_client_connect, + mock_worker, + mock_gather, + mock_env_vars, +): + """Test main function exception handling and cleanup.""" + # Arrange + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + mock_handler = MagicMock() + mock_notification_handler.return_value = mock_handler + + mock_activities_instance = MagicMock() + mock_activities_instance.shutdown = AsyncMock() + mock_activities.return_value = mock_activities_instance + + mock_temporal_client = AsyncMock() + mock_client_connect.return_value = mock_temporal_client + + mock_worker_instance = MagicMock() + mock_worker_instance.run = MagicMock(return_value=AsyncMock()) + mock_worker.return_value = mock_worker_instance + + # Mock gather to raise an exception + mock_gather.side_effect = Exception('Worker failed') + + # Import and run + from model_manager.worker.worker import main + + # Act + await main() + + # Assert - Verify cleanup was performed + mock_logger.custom_error.assert_called_once() + mock_handler.shutdown.assert_called_once() + mock_activities_instance.shutdown.assert_called_once() + mock_metrics.APP_UP.labels.assert_called_once_with(pod_id='test-pod-123') + mock_metrics.APP_UP.labels.return_value.set.assert_called_once_with(0) + mock_sys_exit.assert_called_once_with(1) + + +@pytest.mark.asyncio +@patch('model_manager.worker.worker.sys.exit') +@patch('model_manager.worker.worker.metrics') +@patch('model_manager.worker.worker.Worker') +@patch('model_manager.worker.worker.client.Client.connect') +@patch('model_manager.worker.worker.Runtime') +@patch('model_manager.worker.worker.Activities') +@patch('model_manager.worker.worker.NotificationHandler') +@patch('model_manager.worker.worker.start_prometheus_server') +@patch('model_manager.worker.worker.get_logger') +async def test_main_creates_two_workers( + mock_get_logger, + mock_start_prometheus, + mock_notification_handler, + mock_activities, + mock_runtime, + mock_client_connect, + mock_worker, + mock_metrics, + mock_sys_exit, + mock_env_vars, +): + """Test that main creates two workers with correct configurations.""" + # Arrange + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + mock_handler = MagicMock() + mock_notification_handler.return_value = mock_handler + + mock_activities_instance = MagicMock() + mock_activities_instance.shutdown = AsyncMock() + mock_activities.return_value = mock_activities_instance + + mock_temporal_client = AsyncMock() + mock_client_connect.return_value = mock_temporal_client + + mock_worker_instance = MagicMock() + mock_worker_instance.run = MagicMock(return_value=AsyncMock()) + mock_worker.return_value = mock_worker_instance + + # Import + from model_manager.worker.worker import main + + # Mock gather to prevent infinite wait + with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock): + # Act + await main() + + # Assert - Verify two workers were created + assert mock_worker.call_count == 2 + + # Verify first worker (minimal_retrain-queue) + first_call = mock_worker.call_args_list[0] + assert first_call[1]['task_queue'] == 'minimal_retrain-queue' + assert 'MinimalRetrain' in str(first_call[1]['workflows']) + + # Verify second worker (predictions_batch-queue) + second_call = mock_worker.call_args_list[1] + assert second_call[1]['task_queue'] == 'predictions_batch-queue' + assert 'PredictionsBatch' in str(second_call[1]['workflows']) + + +@pytest.mark.asyncio +@patch('model_manager.worker.worker.sys.exit') +@patch('model_manager.worker.worker.metrics') +@patch('model_manager.worker.worker.Worker') +@patch('model_manager.worker.worker.client.Client.connect') +@patch('model_manager.worker.worker.Runtime') +@patch('model_manager.worker.worker.Activities') +@patch('model_manager.worker.worker.NotificationHandler') +@patch('model_manager.worker.worker.start_prometheus_server') +@patch('model_manager.worker.worker.get_logger') +async def test_main_initializes_activities_with_configs( + mock_get_logger, + mock_start_prometheus, + mock_notification_handler, + mock_activities, + mock_runtime, + mock_client_connect, + mock_worker, + mock_metrics, + mock_sys_exit, + mock_env_vars, +): + """Test that main initializes Activities with correct configurations.""" + # Arrange + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + mock_handler = MagicMock() + mock_notification_handler.return_value = mock_handler + + mock_activities_instance = MagicMock() + mock_activities_instance.shutdown = AsyncMock() + mock_activities.return_value = mock_activities_instance + + mock_temporal_client = AsyncMock() + mock_client_connect.return_value = mock_temporal_client + + mock_worker_instance = MagicMock() + mock_worker_instance.run = MagicMock(return_value=AsyncMock()) + mock_worker.return_value = mock_worker_instance + + # Import + from model_manager.worker.worker import main + + # Mock gather to prevent infinite wait + with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock): + # Act + await main() + + # Assert - Verify Activities was initialized with correct parameters + mock_activities.assert_called_once() + call_kwargs = mock_activities.call_args[1] + assert 'postgres_config' in call_kwargs + assert 'mlflow_config' in call_kwargs + assert 'minio_config' in call_kwargs + assert call_kwargs['logger'] == mock_logger + assert call_kwargs['notification_handler'] == mock_handler + + +@pytest.mark.asyncio +@patch('model_manager.worker.worker.sys.exit') +@patch('model_manager.worker.worker.metrics') +@patch('model_manager.worker.worker.Worker') +@patch('model_manager.worker.worker.client.Client.connect') +@patch('model_manager.worker.worker.Runtime') +@patch('model_manager.worker.worker.Activities') +@patch('model_manager.worker.worker.NotificationHandler') +@patch('model_manager.worker.worker.start_prometheus_server') +@patch('model_manager.worker.worker.get_logger') +async def test_main_uses_environment_variables( + mock_get_logger, + mock_start_prometheus, + mock_notification_handler, + mock_activities, + mock_runtime, + mock_client_connect, + mock_worker, + mock_metrics, + mock_sys_exit, + mock_env_vars, +): + """Test that main uses environment variables correctly.""" + # Arrange + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + mock_handler = MagicMock() + mock_notification_handler.return_value = mock_handler + + mock_activities_instance = MagicMock() + mock_activities_instance.shutdown = AsyncMock() + mock_activities.return_value = mock_activities_instance + + mock_temporal_client = AsyncMock() + mock_client_connect.return_value = mock_temporal_client + + mock_worker_instance = MagicMock() + mock_worker_instance.run = MagicMock(return_value=AsyncMock()) + mock_worker.return_value = mock_worker_instance + + # Import + from model_manager.worker.worker import main + + # Mock gather to prevent infinite wait + with patch('model_manager.worker.worker.asyncio.gather', new_callable=AsyncMock): + # Act + await main() + + # Assert - Verify environment variables were used + mock_client_connect.assert_called_once_with( + target_host='test-temporal:7233', + namespace='test-namespace', + runtime=ANY, + ) + + mock_notification_handler.assert_called_once() + notification_call_kwargs = mock_notification_handler.call_args[1] + assert notification_call_kwargs['project_name'] == 'test-project' + + +@pytest.mark.asyncio +@patch('model_manager.worker.worker.sys.exit') +@patch('model_manager.worker.worker.metrics') +@patch('model_manager.worker.worker.asyncio.gather') +@patch('model_manager.worker.worker.Worker') +@patch('model_manager.worker.worker.client.Client.connect') +@patch('model_manager.worker.worker.Runtime') +@patch('model_manager.worker.worker.Activities') +@patch('model_manager.worker.worker.NotificationHandler') +@patch('model_manager.worker.worker.start_prometheus_server') +@patch('model_manager.worker.worker.get_logger') +async def test_main_cleanup_with_none_notification_handler( + mock_get_logger, + mock_start_prometheus, + mock_notification_handler, + mock_activities, + mock_runtime, + mock_client_connect, + mock_worker, + mock_gather, + mock_metrics, + mock_sys_exit, + mock_env_vars, +): + """Test cleanup when notification_handler is None (line 194 branch False).""" + # Arrange + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + # Return None for notification_handler + mock_notification_handler.return_value = None + + mock_activities_instance = MagicMock() + mock_activities_instance.shutdown = AsyncMock() + mock_activities.return_value = mock_activities_instance + + mock_temporal_client = AsyncMock() + mock_client_connect.return_value = mock_temporal_client + + mock_worker_instance = MagicMock() + mock_worker_instance.run = MagicMock(return_value=AsyncMock()) + mock_worker.return_value = mock_worker_instance + + # Mock gather to complete + mock_gather.return_value = None + + # Import and run + from model_manager.worker.worker import main + + # Act + await main() + + # Assert - notification_handler.shutdown() should NOT be called (line 194 False) + # Since notification_handler is None, we can't call shutdown on it + mock_activities_instance.shutdown.assert_called_once() + mock_sys_exit.assert_called_once_with(1) + + +@pytest.mark.asyncio +@patch('model_manager.worker.worker.sys.exit') +@patch('model_manager.worker.worker.metrics') +@patch('model_manager.worker.worker.asyncio.gather') +@patch('model_manager.worker.worker.Worker') +@patch('model_manager.worker.worker.client.Client.connect') +@patch('model_manager.worker.worker.Runtime') +@patch('model_manager.worker.worker.Activities') +@patch('model_manager.worker.worker.NotificationHandler') +@patch('model_manager.worker.worker.start_prometheus_server') +@patch('model_manager.worker.worker.get_logger') +async def test_main_cleanup_with_falsy_activities( + mock_get_logger, + mock_start_prometheus, + mock_notification_handler, + mock_activities, + mock_runtime, + mock_client_connect, + mock_worker, + mock_gather, + mock_metrics, + mock_sys_exit, + mock_env_vars, +): + """Test cleanup when activities evaluates to False (line 196 branch False).""" + # Arrange + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + mock_handler = MagicMock() + mock_notification_handler.return_value = mock_handler + + # Create a falsy activities object (empty list, 0, False, etc.) + # Using an object that evaluates to False but doesn't cause AttributeError + class FalsyActivities: + def __bool__(self): + return False + + def __getattr__(self, name): + # Return mock methods to avoid AttributeError during worker creation + return MagicMock() + + falsy_activities = FalsyActivities() + mock_activities.return_value = falsy_activities + + mock_temporal_client = AsyncMock() + mock_client_connect.return_value = mock_temporal_client + + mock_worker_instance = MagicMock() + mock_worker_instance.run = MagicMock(return_value=AsyncMock()) + mock_worker.return_value = mock_worker_instance + + # Mock gather to complete + mock_gather.return_value = None + + # Import and run + from model_manager.worker.worker import main + + # Act + await main() + + # Assert - notification_handler.shutdown() is called, but activities.shutdown() is NOT + mock_handler.shutdown.assert_called_once() + # activities is falsy, so shutdown should NOT be called + mock_sys_exit.assert_called_once_with(1)