Fix MCP stateful session cleanup

This commit is contained in:
Cursor Agent 2026-05-05 19:06:57 +00:00
parent 17c34f66f5
commit 6ea4046135
No known key found for this signature in database
2 changed files with 119 additions and 36 deletions

View file

@ -70,6 +70,7 @@ from litellm.utils import Rules, client, function_setup
_byok_cred_cache: Dict[Tuple[str, str], Tuple[Optional[str], float]] = {}
_BYOK_CRED_CACHE_TTL = 60 # seconds
_BYOK_CRED_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth
_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS = 30 * 60
def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
@ -251,6 +252,7 @@ if MCP_AVAILABLE:
stateless=False,
)
_stateful_session_auth_contexts: Dict[str, MCPAuthenticatedUser] = {}
_stateful_session_auth_context_last_seen: Dict[str, float] = {}
# Keep this alias so existing references to session_manager still work
session_manager = session_manager_stateless
@ -267,10 +269,40 @@ if MCP_AVAILABLE:
_session_manager_cm = None
_session_manager_stateful_cm = None
_sse_session_manager_cm = None
_stateful_auth_context_cleanup_task: Optional[asyncio.Task] = None
async def _purge_expired_stateful_session_auth_contexts(
now: Optional[float] = None,
) -> None:
"""Terminate expired stateful sessions and drop their auth contexts."""
now = now or time.monotonic()
server_instances = getattr(session_manager_stateful, "_server_instances", {})
expired_session_ids = [
session_id
for session_id, last_seen in _stateful_session_auth_context_last_seen.items()
if now - last_seen >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
or session_id not in server_instances
]
for session_id in expired_session_ids:
_stateful_session_auth_contexts.pop(session_id, None)
_stateful_session_auth_context_last_seen.pop(session_id, None)
transport = server_instances.pop(session_id, None)
if transport is not None:
await transport.terminate()
for session_id in list(_stateful_session_auth_context_last_seen):
if session_id not in _stateful_session_auth_contexts:
_stateful_session_auth_context_last_seen.pop(session_id, None)
async def _cleanup_expired_stateful_session_auth_contexts() -> None:
while True:
await asyncio.sleep(_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS)
await _purge_expired_stateful_session_auth_contexts()
async def initialize_session_managers():
"""Initialize the session managers. Can be called from main app lifespan."""
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm, _stateful_auth_context_cleanup_task
# Use async lock to prevent concurrent initialization
async with _INITIALIZATION_LOCK:
@ -288,6 +320,9 @@ if MCP_AVAILABLE:
await _session_manager_cm.__aenter__()
await _session_manager_stateful_cm.__aenter__()
await _sse_session_manager_cm.__aenter__()
_stateful_auth_context_cleanup_task = asyncio.create_task(
_cleanup_expired_stateful_session_auth_contexts()
)
_SESSION_MANAGERS_INITIALIZED = True
verbose_logger.info(
@ -296,12 +331,16 @@ if MCP_AVAILABLE:
async def shutdown_session_managers():
"""Shutdown the session managers."""
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm, _stateful_auth_context_cleanup_task
if _SESSION_MANAGERS_INITIALIZED:
verbose_logger.info("Shutting down MCP session managers...")
try:
if _stateful_auth_context_cleanup_task:
_stateful_auth_context_cleanup_task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await _stateful_auth_context_cleanup_task
if _session_manager_cm:
await _session_manager_cm.__aexit__(None, None, None)
if _session_manager_stateful_cm:
@ -314,6 +353,7 @@ if MCP_AVAILABLE:
_session_manager_cm = None
_session_manager_stateful_cm = None
_sse_session_manager_cm = None
_stateful_auth_context_cleanup_task = None
_SESSION_MANAGERS_INITIALIZED = False
@contextlib.asynccontextmanager
@ -2958,6 +2998,7 @@ if MCP_AVAILABLE:
finally:
if use_stateful and session_id and scope.get("method") == "DELETE":
_stateful_session_auth_contexts.pop(session_id, None)
_stateful_session_auth_context_last_seen.pop(session_id, None)
except HTTPException:
# Re-raise HTTP exceptions to preserve status codes and details
raise
@ -3133,7 +3174,8 @@ if MCP_AVAILABLE:
auth_user = (
_stateful_session_auth_contexts.get(session_id) if session_id else None
)
if auth_user is not None:
if auth_user is not None and session_id is not None:
_stateful_session_auth_context_last_seen[session_id] = time.monotonic()
_update_auth_context(
auth_user=auth_user,
user_api_key_auth=user_api_key_auth,
@ -3164,7 +3206,11 @@ if MCP_AVAILABLE:
if message.get("type") == "http.response.start":
for key, value in message.get("headers", []):
if key.lower() == b"mcp-session-id":
_stateful_session_auth_contexts[value.decode()] = auth_user
session_id = value.decode()
_stateful_session_auth_contexts[session_id] = auth_user
_stateful_session_auth_context_last_seen[session_id] = (
time.monotonic()
)
break
await send(message)

View file

@ -1154,31 +1154,39 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless():
async def stateful_handle(s, r, se):
stateful_called.append(1)
with patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
), patch(
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
), patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
), patch.object(
session_manager_stateless,
"handle_request",
side_effect=stateless_handle,
), patch.object(
session_manager_stateful,
"handle_request",
side_effect=stateful_handle,
), patch.object(
session_manager_stateless,
"_server_instances",
{},
), patch.object(
session_manager_stateful,
"_server_instances",
{},
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
),
patch(
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch.object(
session_manager_stateless,
"handle_request",
side_effect=stateless_handle,
),
patch.object(
session_manager_stateful,
"handle_request",
side_effect=stateful_handle,
),
patch.object(
session_manager_stateless,
"_server_instances",
{},
),
patch.object(
session_manager_stateful,
"_server_instances",
{},
),
):
await handle_streamable_http_mcp(scope, receive, send)
@ -1187,16 +1195,16 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless():
# initialize → stateful
init_body = b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2024-11-05"}}'
stateless_called, stateful_called = await make_request(init_body)
assert stateful_called and not stateless_called, (
"initialize (no session) should route to stateful, not stateless"
)
assert (
stateful_called and not stateless_called
), "initialize (no session) should route to stateful, not stateless"
# tools/list → stateless
tools_body = b'{"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}'
stateless_called, stateful_called = await make_request(tools_body)
assert stateless_called and not stateful_called, (
"tools/list (no session) should route to stateless, not stateful"
)
assert (
stateless_called and not stateful_called
), "tools/list (no session) should route to stateless, not stateful"
@pytest.mark.asyncio
@ -1379,6 +1387,36 @@ async def test_stateful_mcp_requests_refresh_session_auth_context():
mcp_server._stateful_session_auth_contexts.pop(session_id, None)
@pytest.mark.asyncio
async def test_stateful_mcp_auth_contexts_expire_with_idle_sessions():
"""Expired session auth contexts should not remain in memory indefinitely."""
try:
from litellm.proxy._experimental.mcp_server import server as mcp_server
except ImportError:
pytest.skip("MCP server not available")
session_id = "expired-stateful-session"
auth_user = UserAPIKeyAuth(api_key="expired-key", user_id="expired-user")
transport = MagicMock()
now = 1000.0
mcp_server._stateful_session_auth_contexts[session_id] = auth_user
mcp_server._stateful_session_auth_context_last_seen[session_id] = (
now - mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
)
with patch.object(
mcp_server.session_manager_stateful,
"_server_instances",
{session_id: transport},
):
await mcp_server._purge_expired_stateful_session_auth_contexts(now=now)
assert session_id not in mcp_server._stateful_session_auth_contexts
assert session_id not in mcp_server._stateful_session_auth_context_last_seen
transport.terminate.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.no_parallel
async def test_mcp_routing_with_conflicting_alias_and_group_name():
@ -3405,4 +3443,3 @@ async def test_call_tool_empty_extra_headers_returns_none():
"P2 API consistency issue: expected None for empty extra_headers, got: "
+ str(captured_extra_headers)
)