fix by adding init lock (#13666)

This commit is contained in:
Jugal D. Bhatt 2025-08-22 11:42:43 -07:00 • committed by GitHub
parent 694a9f0f4c
commit 3c1dcb64cc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 81 additions and 11 deletions

View file

@ -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."""

View file

@ -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