diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index aa2a6ee1a75..38619112ccc 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -40,6 +40,7 @@ except ImportError as e: # Global variables to track initialization _SESSION_MANAGERS_INITIALIZED = False +_INITIALIZATION_LOCK = asyncio.Lock() if MCP_AVAILABLE: from mcp.server import Server @@ -113,21 +114,23 @@ if MCP_AVAILABLE: """Initialize the session managers. Can be called from main app lifespan.""" global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm - if _SESSION_MANAGERS_INITIALIZED: - return + # Use async lock to prevent concurrent initialization + async with _INITIALIZATION_LOCK: + if _SESSION_MANAGERS_INITIALIZED: + return - verbose_logger.info("Initializing MCP session managers...") + verbose_logger.info("Initializing MCP session managers...") - # Start the session managers with context managers - _session_manager_cm = session_manager.run() - _sse_session_manager_cm = sse_session_manager.run() + # Start the session managers with context managers + _session_manager_cm = session_manager.run() + _sse_session_manager_cm = sse_session_manager.run() - # Enter the context managers - await _session_manager_cm.__aenter__() - await _sse_session_manager_cm.__aenter__() + # Enter the context managers + await _session_manager_cm.__aenter__() + await _sse_session_manager_cm.__aenter__() - _SESSION_MANAGERS_INITIALIZED = True - verbose_logger.info("MCP Server started with StreamableHTTP and SSE session managers!") + _SESSION_MANAGERS_INITIALIZED = True + verbose_logger.info("MCP Server started with StreamableHTTP and SSE session managers!") async def shutdown_session_managers(): """Shutdown the session managers.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index d67e22b6d3b..42c64c15814 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -275,3 +275,70 @@ async def test_mcp_server_tool_call_body_with_none_arguments(): body = captured_data["proxy_server_request"]["body"] assert body["name"] == tool_name assert body["arguments"] == tool_arguments # Should be None + + +@pytest.mark.asyncio +async def test_concurrent_initialize_session_managers(): + """Test that concurrent calls to initialize_session_managers don't cause race conditions.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + initialize_session_managers, + _SESSION_MANAGERS_INITIALIZED, + _INITIALIZATION_LOCK, + ) + except ImportError: + pytest.skip("MCP server not available") + + # Import the module to reset state + import litellm.proxy._experimental.mcp_server.server as mcp_server + + # Reset state before test + original_initialized = mcp_server._SESSION_MANAGERS_INITIALIZED + original_session_cm = mcp_server._session_manager_cm + original_sse_session_cm = mcp_server._sse_session_manager_cm + + try: + mcp_server._SESSION_MANAGERS_INITIALIZED = False + mcp_server._session_manager_cm = None + mcp_server._sse_session_manager_cm = None + + # Mock the session managers to avoid actual MCP initialization + with patch('litellm.proxy._experimental.mcp_server.server.session_manager') as mock_session_manager, \ + patch('litellm.proxy._experimental.mcp_server.server.sse_session_manager') as mock_sse_session_manager, \ + patch('litellm.proxy._experimental.mcp_server.server.verbose_logger'): + + # Mock the run() method to return a mock context manager + mock_cm = AsyncMock() + mock_cm.__aenter__ = AsyncMock() + mock_cm.__aexit__ = AsyncMock() + + mock_session_manager.run.return_value = mock_cm + mock_sse_session_manager.run.return_value = mock_cm + + # Create multiple concurrent tasks that call initialize_session_managers + async def init_task(): + await initialize_session_managers() + return "success" + + # Run 10 concurrent initialization attempts + tasks = [init_task() for _ in range(10)] + results = await asyncio.gather(*tasks, return_exceptions=True) + + # All tasks should complete successfully (no exceptions) + assert all(result == "success" for result in results), f"Some tasks failed: {results}" + + # session_manager.run() should only be called once due to the lock + assert mock_session_manager.run.call_count == 1, f"Expected 1 call to session_manager.run(), got {mock_session_manager.run.call_count}" + assert mock_sse_session_manager.run.call_count == 1, f"Expected 1 call to sse_session_manager.run(), got {mock_sse_session_manager.run.call_count}" + + # The context managers should only be entered once each + assert mock_cm.__aenter__.call_count == 2, f"Expected 2 calls to __aenter__ (one for each session manager), got {mock_cm.__aenter__.call_count}" + + # State should be properly set + assert mcp_server._SESSION_MANAGERS_INITIALIZED is True + + finally: + # Restore original state + mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized + mcp_server._session_manager_cm = original_session_cm + mcp_server._sse_session_manager_cm = original_sse_session_cm