Fix MCP initialize session active tracking

Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-13 17:32:23 +00:00
parent dd4af19c49
commit eaf99edfe3
No known key found for this signature in database
2 changed files with 125 additions and 16 deletions

View file

@ -3148,10 +3148,29 @@ if MCP_AVAILABLE:
session_id, asyncio.Lock()
)
track_active_stateful_request = bool(use_stateful and session_id)
if track_active_stateful_request and session_id:
_stateful_session_active_request_counts[session_id] = (
_stateful_session_active_request_counts.get(session_id, 0) + 1
active_request_session_id = (
session_id if use_stateful and session_id else None
)
if active_request_session_id:
_stateful_session_active_request_counts[active_request_session_id] = (
_stateful_session_active_request_counts.get(
active_request_session_id, 0
)
+ 1
)
def _track_initialized_stateful_session(
initialized_session_id: str,
) -> None:
nonlocal active_request_session_id
if active_request_session_id is not None:
return
active_request_session_id = initialized_session_id
_stateful_session_active_request_counts[initialized_session_id] = (
_stateful_session_active_request_counts.get(
initialized_session_id, 0
)
+ 1
)
async def _dispatch() -> None:
@ -3174,6 +3193,7 @@ if MCP_AVAILABLE:
_owner_fingerprint_for(
user_api_key_auth, oauth2_headers, _client_ip
),
_track_initialized_stateful_session,
)
async with _gateway_initialize_instructions_request_scope(
@ -3192,32 +3212,38 @@ if MCP_AVAILABLE:
else:
await _dispatch()
finally:
if track_active_stateful_request and session_id:
if active_request_session_id:
active_request_count = (
_stateful_session_active_request_counts.get(session_id, 0) - 1
_stateful_session_active_request_counts.get(
active_request_session_id, 0
)
- 1
)
if active_request_count > 0:
_stateful_session_active_request_counts[session_id] = (
active_request_count
)
_stateful_session_active_request_counts[
active_request_session_id
] = active_request_count
else:
_stateful_session_active_request_counts.pop(session_id, None)
_stateful_session_active_request_counts.pop(
active_request_session_id, None
)
if (
scope.get("method") != "DELETE"
and session_id in _stateful_session_auth_contexts
and active_request_session_id in _stateful_session_auth_contexts
):
_stateful_session_auth_context_last_seen[session_id] = (
time.monotonic()
)
_stateful_session_auth_context_last_seen[
active_request_session_id
] = time.monotonic()
# Periodic cleanup iterates _stateful_session_auth_context_last_seen,
# so locks for untracked sessions must be dropped here.
if (
active_request_count <= 0
and session_id not in _stateful_session_auth_contexts
and active_request_session_id
not in _stateful_session_auth_contexts
):
_stateful_session_locks.pop(session_id, None)
_stateful_session_locks.pop(active_request_session_id, None)
except HTTPException:
# Re-raise HTTP exceptions to preserve status codes and details
raise
@ -3422,6 +3448,7 @@ if MCP_AVAILABLE:
send: Send,
auth_user: MCPAuthenticatedUser,
owner_fingerprint: str,
on_session_registered: Optional[Callable[[str], None]] = None,
) -> Send:
async def wrapped_send(message: Message) -> None:
if message.get("type") == "http.response.start":
@ -3431,6 +3458,8 @@ if MCP_AVAILABLE:
session_id = (
value.decode() if isinstance(value, bytes) else str(value)
)
if on_session_registered is not None:
on_session_registered(session_id)
auth_context_var.set(auth_user)
_stateful_session_auth_contexts[session_id] = auth_user
_stateful_session_auth_context_last_seen[session_id] = (

View file

@ -1559,6 +1559,86 @@ async def test_initialize_response_capture_accepts_str_headers_and_sets_auth_con
mcp_server._remove_stateful_session_tracking(session_id)
@pytest.mark.asyncio
async def test_initialize_request_tracks_active_session_after_response_header():
try:
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager_stateful,
session_manager_stateless,
)
except ImportError:
pytest.skip("MCP server not available")
session_id = "initialize-active-session-1"
owner_auth = UserAPIKeyAuth(api_key="initialize-key", user_id="user-a")
scope = {
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [
(b"content-type", b"application/json"),
(b"authorization", b"Bearer initialize-key"),
],
}
receive = AsyncMock(
return_value={
"type": "http.request",
"body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}',
"more_body": False,
}
)
async def stateful_handle(s, r, se):
await se(
{
"type": "http.response.start",
"headers": [(b"mcp-session-id", session_id.encode())],
}
)
assert mcp_server._stateful_session_active_request_counts[session_id] == 1
now = (
mcp_server._stateful_session_auth_context_last_seen[session_id]
+ mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
)
await mcp_server._purge_expired_stateful_session_auth_contexts(now=now)
assert session_id in mcp_server._stateful_session_auth_contexts
async def stateless_handle(s, r, se):
raise AssertionError("initialize request should use stateful manager")
try:
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(owner_auth, None, None, None, None, None),
),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch.object(
session_manager_stateful,
"handle_request",
side_effect=stateful_handle,
),
patch.object(
session_manager_stateless,
"handle_request",
side_effect=stateless_handle,
),
patch.object(session_manager_stateful, "_server_instances", {}),
):
await handle_streamable_http_mcp(scope, receive, AsyncMock())
assert session_id not in mcp_server._stateful_session_active_request_counts
assert session_id in mcp_server._stateful_session_auth_contexts
finally:
mcp_server._remove_stateful_session_tracking(session_id)
@pytest.mark.asyncio
async def test_stateful_mcp_auth_contexts_expire_with_idle_sessions():
"""Expired session auth contexts should not remain in memory indefinitely."""