Fix stateful MCP auth context refresh

This commit is contained in:
Cursor Agent 2026-05-05 18:49:22 +00:00
parent cf986ff343
commit 17c34f66f5
No known key found for this signature in database
2 changed files with 192 additions and 14 deletions

View file

@ -250,6 +250,7 @@ if MCP_AVAILABLE:
json_response=False, # enables SSE streaming
stateless=False,
)
_stateful_session_auth_contexts: Dict[str, MCPAuthenticatedUser] = {}
# Keep this alias so existing references to session_manager still work
session_manager = session_manager_stateless
@ -2882,17 +2883,6 @@ if MCP_AVAILABLE:
if _debug_headers:
send = MCPDebug.wrap_send_with_debug_headers(send, _debug_headers)
# Set the auth context variable for easy access in MCP functions
set_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
)
# Ensure session managers are initialized
if not _SESSION_MANAGERS_INITIALIZED:
await initialize_session_managers()
@ -2945,12 +2935,29 @@ if MCP_AVAILABLE:
receive = wrapped_receive
auth_user = _set_or_update_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
session_id=session_id if use_stateful else None,
)
if use_stateful and is_initialize:
send = _wrap_send_with_stateful_session_auth_context(send, auth_user)
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth,
mcp_servers,
_client_ip,
):
await target_manager.handle_request(scope, receive, send)
try:
await target_manager.handle_request(scope, receive, send)
finally:
if use_stateful and session_id and scope.get("method") == "DELETE":
_stateful_session_auth_contexts.pop(session_id, None)
except HTTPException:
# Re-raise HTTP exceptions to preserve status codes and details
raise
@ -3064,8 +3071,9 @@ if MCP_AVAILABLE:
############ Auth Context Functions ####################
########################################################
def set_auth_context(
user_api_key_auth: UserAPIKeyAuth,
def _update_auth_context(
auth_user: MCPAuthenticatedUser,
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
@ -3073,6 +3081,23 @@ if MCP_AVAILABLE:
raw_headers: Optional[Dict[str, str]] = None,
client_ip: Optional[str] = None,
) -> None:
auth_user.user_api_key_auth = user_api_key_auth
auth_user.mcp_auth_header = mcp_auth_header
auth_user.mcp_servers = mcp_servers
auth_user.mcp_server_auth_headers = mcp_server_auth_headers or {}
auth_user.oauth2_headers = oauth2_headers
auth_user.raw_headers = raw_headers
auth_user.client_ip = client_ip
def set_auth_context(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
client_ip: Optional[str] = None,
) -> MCPAuthenticatedUser:
"""
Set the UserAPIKeyAuth in the auth context variable.
@ -3093,6 +3118,57 @@ if MCP_AVAILABLE:
client_ip=client_ip,
)
auth_context_var.set(auth_user)
return auth_user
def _set_or_update_auth_context(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
client_ip: Optional[str] = None,
session_id: Optional[str] = None,
) -> MCPAuthenticatedUser:
auth_user = (
_stateful_session_auth_contexts.get(session_id) if session_id else None
)
if auth_user is not None:
_update_auth_context(
auth_user=auth_user,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
)
auth_context_var.set(auth_user)
return auth_user
return set_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
)
def _wrap_send_with_stateful_session_auth_context(
send: Send,
auth_user: MCPAuthenticatedUser,
) -> Send:
async def wrapped_send(message: Message) -> None:
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
break
await send(message)
return wrapped_send
def get_auth_context() -> Tuple[
Optional[UserAPIKeyAuth],

View file

@ -1,4 +1,5 @@
import asyncio
import contextvars
from datetime import datetime, timedelta
from unittest.mock import AsyncMock, MagicMock, patch
@ -1277,6 +1278,107 @@ async def test_mcp_routing_chunked_initialize_to_stateful():
)
@pytest.mark.asyncio
async def test_stateful_mcp_requests_refresh_session_auth_context():
"""
Stateful MCP sessions run callbacks in the initialize task's context; the
stored auth object must be refreshed for each mcp-session-id request.
"""
try:
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import (
get_auth_context,
handle_streamable_http_mcp,
session_manager_stateful,
)
except ImportError:
pytest.skip("MCP server not available")
session_id = "stateful-session-1"
initialize_auth = UserAPIKeyAuth(api_key="initialize-key", user_id="user-a")
current_auth = UserAPIKeyAuth(api_key="current-key", user_id="user-b")
callback_context = contextvars.copy_context()
callback_context.run(
mcp_server.set_auth_context,
initialize_auth,
None,
["old-server"],
None,
None,
None,
"1.1.1.1",
)
mcp_server._stateful_session_auth_contexts[session_id] = callback_context.run(
mcp_server.auth_context_var.get
)
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/current-server",
"headers": [
(b"content-type", b"application/json"),
(b"authorization", b"Bearer current-key"),
(b"mcp-session-id", session_id.encode()),
],
}
receive = AsyncMock(
return_value={
"type": "http.request",
"body": b'{"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}',
"more_body": False,
}
)
send = AsyncMock()
captured_context = None
async def stateful_handle(s, r, se):
nonlocal captured_context
captured_context = callback_context.run(get_auth_context)
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(
current_auth,
"current-mcp-auth",
["current-server"],
{"current-server": {"Authorization": "Bearer server-key"}},
{"Authorization": "Bearer oauth-key"},
{"mcp-session-id": session_id},
),
),
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_stateful,
"_server_instances",
{session_id: MagicMock()},
),
):
await handle_streamable_http_mcp(scope, receive, send)
assert captured_context == (
current_auth,
"current-mcp-auth",
["current-server"],
{"current-server": {"Authorization": "Bearer server-key"}},
{"Authorization": "Bearer oauth-key"},
{"mcp-session-id": session_id},
"",
)
mcp_server._stateful_session_auth_contexts.pop(session_id, None)
@pytest.mark.asyncio
@pytest.mark.no_parallel
async def test_mcp_routing_with_conflicting_alias_and_group_name():