Fix MCP reinitialize session tracking

Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-13 19:14:28 +00:00
parent 225c11c2d9
commit 897aaa5dc1
No known key found for this signature in database
2 changed files with 120 additions and 22 deletions

View file

@ -72,8 +72,8 @@ _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
# Maximum bytes to peek when sniffing the JSON-RPC method on a no-session-id
# POST. An `initialize` envelope is a few hundred bytes; capping the peek
# Maximum bytes to peek when sniffing the JSON-RPC method on a POST.
# An `initialize` envelope is a few hundred bytes; capping the peek
# prevents an authenticated client from forcing the proxy to buffer an
# arbitrarily large body just to make a routing decision.
_MCP_ROUTING_PEEK_MAX_BYTES = 4096
@ -3083,7 +3083,7 @@ if MCP_AVAILABLE:
return
session_id = _get_session_id_from_scope(scope)
if scope.get("method") == "POST" and not session_id:
if scope.get("method") == "POST":
consumed_messages, body = await _read_request_body_for_routing(receive)
is_initialize = _is_initialize_request(body)
@ -3148,30 +3148,24 @@ if MCP_AVAILABLE:
session_id, asyncio.Lock()
)
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
)
active_request_session_ids: List[str] = []
def _increment_active_request_session(session_id_to_track: str) -> None:
if session_id_to_track in active_request_session_ids:
return
active_request_session_ids.append(session_id_to_track)
_stateful_session_active_request_counts[session_id_to_track] = (
_stateful_session_active_request_counts.get(session_id_to_track, 0)
+ 1
)
if use_stateful and session_id:
_increment_active_request_session(session_id)
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
)
_increment_active_request_session(initialized_session_id)
async def _dispatch() -> None:
auth_user = _set_or_update_auth_context(
@ -3212,7 +3206,7 @@ if MCP_AVAILABLE:
else:
await _dispatch()
finally:
if active_request_session_id:
for active_request_session_id in active_request_session_ids:
active_request_count = (
_stateful_session_active_request_counts.get(
active_request_session_id, 0

View file

@ -1639,6 +1639,110 @@ async def test_initialize_request_tracks_active_session_after_response_header():
mcp_server._remove_stateful_session_tracking(session_id)
@pytest.mark.asyncio
async def test_initialize_request_with_existing_session_tracks_new_session():
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")
existing_session_id = "existing-initialize-session"
new_session_id = "reinitialized-session"
owner_auth = UserAPIKeyAuth(api_key="initialize-key", user_id="user-a")
owner_fingerprint = mcp_server._owner_fingerprint_for(owner_auth)
initialize_body = b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}'
scope = {
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [
(b"content-type", b"application/json"),
(b"authorization", b"Bearer initialize-key"),
(b"mcp-session-id", existing_session_id.encode()),
],
}
receive = AsyncMock(
return_value={
"type": "http.request",
"body": initialize_body,
"more_body": False,
}
)
stateful_called = []
async def stateful_handle(s, r, se):
stateful_called.append(1)
message = await r()
assert message["body"] == initialize_body
await se(
{
"type": "http.response.start",
"headers": [(b"mcp-session-id", new_session_id.encode())],
}
)
assert mcp_server._stateful_session_auth_contexts[new_session_id]
assert mcp_server._stateful_session_owners[new_session_id] == owner_fingerprint
assert mcp_server._stateful_session_active_request_counts[new_session_id] == 1
now = (
mcp_server._stateful_session_auth_context_last_seen[new_session_id]
+ mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
)
await mcp_server._purge_expired_stateful_session_auth_contexts(now=now)
assert new_session_id in mcp_server._stateful_session_auth_contexts
async def stateless_handle(s, r, se):
raise AssertionError(
"initialize request with session should use stateful manager"
)
try:
mcp_server._stateful_session_auth_contexts[existing_session_id] = (
mcp_server.MCPAuthenticatedUser(user_api_key_auth=owner_auth)
)
mcp_server._stateful_session_auth_context_last_seen[existing_session_id] = 1.0
mcp_server._stateful_session_owners[existing_session_id] = owner_fingerprint
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",
{existing_session_id: MagicMock()},
),
):
await handle_streamable_http_mcp(scope, receive, AsyncMock())
assert stateful_called
assert new_session_id not in mcp_server._stateful_session_active_request_counts
assert new_session_id in mcp_server._stateful_session_auth_contexts
finally:
mcp_server._remove_stateful_session_tracking(existing_session_id)
mcp_server._remove_stateful_session_tracking(new_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."""