mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Fix stateful MCP auth context refresh
This commit is contained in:
parent
cf986ff343
commit
17c34f66f5
2 changed files with 192 additions and 14 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue