mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Fix MCP stateful session cleanup
This commit is contained in:
parent
17c34f66f5
commit
6ea4046135
2 changed files with 119 additions and 36 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue