From f7a590901924f9419256fae5422f365e30c1a8e6 Mon Sep 17 00:00:00 2001 From: Yug Date: Fri, 1 May 2026 07:12:27 +0530 Subject: [PATCH] resolve --- .../mcp_server/mcp_server_manager.py | 1 + .../proxy/_experimental/mcp_server/server.py | 228 +++++++----------- tests/mcp_tests/test_mcp_session_binding.py | 37 ++- 3 files changed, 119 insertions(+), 147 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f3fd1360525..a80befb64f7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1127,6 +1127,7 @@ class MCPServerManager: tools = await self._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, + user_api_key_auth=user_api_key_auth, ) return tools except Exception as e: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index cdc3c2d7189..9e8eea04140 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -285,8 +285,19 @@ if MCP_AVAILABLE: from mcp.server.sse import SseServerTransport as _McpSseServerTransport sse = _McpSseServerTransport("/mcp/messages") - # Wrap the dict to capture session IDs race-free during connect_sse - sse._read_stream_writers = _SessionIdCapturingDict(sse._read_stream_writers) # type: ignore + # Wrap the dict to capture session IDs race-free during connect_sse. + # We use a hasattr guard to avoid module-load crashes if the private SDK + # attribute is renamed or removed in future versions. + if hasattr(sse, "_read_stream_writers"): + sse._read_stream_writers = _SessionIdCapturingDict( + getattr(sse, "_read_stream_writers") + ) # type: ignore + else: + verbose_logger.error( + "MCP SSE Transport: '_read_stream_writers' attribute not found on " + "SseServerTransport. Session ID capture for auth-binding will be " + "disabled. Consider upgrading LiteLLM or reporting an SDK compatibility issue." + ) # Create session managers (StreamableHTTP — stateless by default) session_manager = StreamableHTTPSessionManager( app=server, @@ -2547,6 +2558,33 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op}) + async def _render_mcp_error( + e: Exception, scope: Scope, receive: Receive, send: Send + ) -> None: + """Render an exception as a JSON response for ASGI handlers.""" + from starlette.responses import JSONResponse + + status_code = 500 + detail = str(e) + + if isinstance(e, HTTPException): + status_code = e.status_code + detail = e.detail + elif hasattr(e, "status_code"): + status_code = getattr(e, "status_code") + elif hasattr(e, "code"): + # LiteLLM ProxyException uses 'code' + try: + status_code = int(getattr(e, "code")) + except (ValueError, TypeError): + status_code = 500 + + error_response = JSONResponse( + status_code=status_code, + content={"error": "MCP request failed", "details": detail}, + ) + await error_response(scope, receive, send) + async def handle_streamable_http_mcp( scope: Scope, receive: Receive, send: Send ) -> None: @@ -2563,22 +2601,12 @@ if MCP_AVAILABLE: ) = await extract_mcp_auth_context(scope, path) # Extract client IP for MCP access control _client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) - verbose_logger.debug( - f"MCP request mcp_servers (header/path): {mcp_servers}" - ) - verbose_logger.debug( - f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" - ) - # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response + # ... (rest of authentication checks) for server_name in mcp_servers or []: server = global_mcp_server_manager.get_mcp_server_by_name( server_name, client_ip=_client_ip ) if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers: - # For per-user OAuth servers, only skip the pre-emptive 401 when - # a stored token actually exists for this user+server pair. - # If no stored token exists, fail fast with 401 so clients can - # kick off PKCE/interactive OAuth flow immediately. if server.needs_user_oauth_token: stored_oauth_headers = ( await _get_user_oauth_extra_headers_from_db( @@ -2602,8 +2630,7 @@ if MCP_AVAILABLE: for k, v in scope.get("headers", []) if k.lower() != b"x-mcp-toolset-id" ] - # Apply toolset scope if set server-side via ContextVar (set by - # /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py). + # Apply toolset scope if set server-side via ContextVar active_toolset_id = _mcp_active_toolset_id.get() if active_toolset_id and user_api_key_auth is not None: user_api_key_auth = await _apply_toolset_scope( @@ -2634,15 +2661,12 @@ if MCP_AVAILABLE: # Ensure session managers are initialized if not _SESSION_MANAGERS_INITIALIZED: await initialize_session_managers() - # Give it a moment to start up await asyncio.sleep(0.1) - # Handle stale session IDs - either strip them for reconnection - # or return success for idempotent DELETE operations + # Handle stale session IDs handled = await _handle_stale_mcp_session( scope, receive, send, session_manager ) if handled: - # Request was fully handled (e.g., DELETE on non-existent session) return async with _gateway_initialize_instructions_request_scope( user_api_key_auth, @@ -2650,27 +2674,8 @@ if MCP_AVAILABLE: _client_ip, ): await session_manager.handle_request(scope, receive, send) - except HTTPException: - # Re-raise HTTP exceptions to preserve status codes and details - raise except Exception as e: - verbose_logger.exception(f"Error handling MCP request: {e}") - # Try to send a graceful error response for non-HTTP exceptions - try: - from starlette.responses import JSONResponse - from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR - - error_response = JSONResponse( - status_code=HTTP_500_INTERNAL_SERVER_ERROR, - content={"error": "MCP request failed", "details": str(e)}, - ) - await error_response(scope, receive, send) - except Exception as response_error: - verbose_logger.exception( - f"Failed to send error response: {response_error}" - ) - # If we can't send a proper response, re-raise the original error - raise e + await _render_mcp_error(e, scope, receive, send) async def handle_sse_mcp_endpoint( scope: Scope, receive: Receive, send: Send @@ -2693,12 +2698,6 @@ if MCP_AVAILABLE: ) = await extract_mcp_auth_context(scope, path) # Extract client IP for MCP access control _sse_client_ip = IPAddressUtils.get_mcp_client_ip(request) - verbose_logger.debug( - f"MCP SSE request mcp_servers (header/path): {mcp_servers}" - ) - verbose_logger.debug( - f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" - ) set_auth_context( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -2711,13 +2710,9 @@ if MCP_AVAILABLE: if not _SESSION_MANAGERS_INITIALIZED: await initialize_session_managers() - await asyncio.sleep(0.1) async with _gateway_initialize_instructions_request_scope( - user_api_key_auth, - mcp_servers, - _sse_client_ip, + user_api_key_auth, mcp_servers, _sse_client_ip ): - verbose_logger.info("Initializing SSE session...") options = server.create_initialization_options() # Capture existing session IDs to find the newly created one # We use a mutable container because ContextVar writes are task-local @@ -2777,27 +2772,8 @@ if MCP_AVAILABLE: finally: _captured_session_id_container_var.reset(_capture_token) - except HTTPException: - raise except Exception as e: - verbose_logger.exception(f"Error handling MCP request: {e}") - # Instead of re-raising, try to send a graceful error response - try: - # Send a proper HTTP error response instead of letting the exception bubble up - from starlette.responses import JSONResponse - from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR - - error_response = JSONResponse( - status_code=HTTP_500_INTERNAL_SERVER_ERROR, - content={"error": "MCP request failed", "details": str(e)}, - ) - await error_response(scope, receive, send) - except Exception as response_error: - verbose_logger.exception( - f"Failed to send error response: {response_error}" - ) - # If we can't send a proper response, re-raise the original error - raise e + await _render_mcp_error(e, scope, receive, send) # No need to return Response for raw ASGI app. async def handle_sse_post_messages( @@ -2816,63 +2792,38 @@ if MCP_AVAILABLE: raw_headers, ) = await extract_mcp_auth_context(scope, path) _sse_client_ip = IPAddressUtils.get_mcp_client_ip(request) - # set_auth_context here is a no-op for actual tool execution since the SDK session_id_param = request.query_params.get( "sessionId" ) or request.query_params.get("session_id") if session_id_param: - import uuid - try: session_id = uuid.UUID(hex=session_id_param) - # Look up auth associated with this session ID session_auth = _session_id_auth_storage.get(session_id) - if session_auth is not None: - # P1 Security: Bind POST auth to the existing SSE session - session_api_key = getattr( + # P1 Security Hardening: Ensure POST auth matches session owner + session_key = getattr( session_auth.user_api_key_auth, "api_key", None ) - session_has_auth = ( - session_api_key is not None or session_auth.oauth2_headers + session_oauth = (session_auth.oauth2_headers or {}).get( + "Authorization" ) + post_key = ( + getattr(user_api_key_auth, "api_key", None) + if user_api_key_auth + else None + ) + post_oauth = (oauth2_headers or {}).get("Authorization") - if session_has_auth: - post_api_key = ( - getattr(user_api_key_auth, "api_key", None) - if user_api_key_auth - else None - ) - - # Match by API Key - if session_api_key is not None: - if session_api_key != post_api_key: - raise HTTPException( - status_code=403, - detail="Authentication mismatch: POST auth does not match the session owner.", - ) - # Match by OAuth2 Token if no API key - elif session_auth.oauth2_headers: - if ( - not oauth2_headers - or session_auth.oauth2_headers.get("Authorization") - != oauth2_headers.get("Authorization") - ): - raise HTTPException( - status_code=403, - detail="Authentication mismatch: POST OAuth2 token does not match the session owner.", - ) - # If session was authenticated but POST has no auth, deny. - elif user_api_key_auth is None and not oauth2_headers: + # If session is authenticated, POST must have matching credentials + if session_key or session_oauth: + if session_key != post_key or session_oauth != post_oauth: raise HTTPException( status_code=403, - detail="Authentication required: Target session is authenticated but POST lacks credentials.", + detail="Authentication mismatch: POST does not match session owner.", ) except ValueError: - # Invalid UUID format in session_id_param pass - # Set auth context for this POST request (important for SDK background tasks) set_auth_context( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -2883,28 +2834,8 @@ if MCP_AVAILABLE: client_ip=_sse_client_ip, ) await sse.handle_post_message(scope, receive, send) - except HTTPException: - raise except Exception as e: - verbose_logger.exception( - f"Failed to extract auth context in POST /messages: {e}" - ) - from starlette.responses import JSONResponse - - error_response = JSONResponse( - status_code=500, - content={ - "error": "Authentication processing failed", - "details": "An internal error occurred during authentication processing.", - }, - ) - await error_response(scope, receive, send) - # CRITICAL: Stop execution if auth extraction fails to prevent unauthenticated fall-through - return - - # The SDK's handler calls `send` directly. - # Only proceed if auth was successfully extracted - # No need to return NoOpResponse for raw ASGI app. + await _render_mcp_error(e, scope, receive, send) def get_active_mcp_session() -> Optional[_McpServerSession]: """Get the active downstream MCP session from the current context.""" @@ -3088,15 +3019,36 @@ if MCP_AVAILABLE: if user_api_key_auth is not None: session = _get_current_session() if session is not None and id(session) not in _session_obj_auth_storage: - _session_obj_auth_storage[id(session)] = MCPAuthenticatedUser( - 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, + # Reuse the *same* MCPAuthenticatedUser reference that + # handle_sse_mcp_endpoint stored in _session_auth_storage. + # This is critical: the cleanup ``finally`` block uses + # ``v is auth`` (identity check), so we must store the + # original object, not a fresh copy. + read_stream = getattr(session, "_read_stream", None) + existing_auth = ( + _session_auth_storage.get(read_stream) + if read_stream is not None + else None ) + if existing_auth is not None: + _session_obj_auth_storage[id(session)] = existing_auth + else: + # Shouldn't normally happen — _session_auth_storage is + # populated before server.run(). Log for diagnostics. + verbose_logger.debug( + "get_or_extract_auth_context: no existing auth in " + "_session_auth_storage; creating standalone entry " + "(cleanup may not remove it via identity check)" + ) + _session_obj_auth_storage[id(session)] = MCPAuthenticatedUser( + 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, + ) else: # Fallback: try session-object-identity lookup first (robust) stored: Optional[MCPAuthenticatedUser] = None diff --git a/tests/mcp_tests/test_mcp_session_binding.py b/tests/mcp_tests/test_mcp_session_binding.py index 1b0b3e253c3..96803e4a1f5 100644 --- a/tests/mcp_tests/test_mcp_session_binding.py +++ b/tests/mcp_tests/test_mcp_session_binding.py @@ -2,7 +2,7 @@ import pytest import uuid import asyncio from unittest.mock import AsyncMock, MagicMock, patch -from fastapi import HTTPException + from contextlib import asynccontextmanager from litellm.proxy._experimental.mcp_server.server import ( handle_sse_mcp_endpoint, @@ -160,10 +160,20 @@ async def test_session_id_capture_and_binding(): "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", AsyncMock(return_value=wrong_post_auth_context), ): - with pytest.raises(HTTPException) as exc: - await handle_sse_post_messages(post_scope, mock_receive, mock_send) - assert exc.value.status_code == 403 - assert "Authentication mismatch" in exc.value.detail + # Reset mock_send calls + mock_send.reset_mock() + await handle_sse_post_messages(post_scope, mock_receive, mock_send) + + # Verify 403 was sent + assert mock_send.call_count >= 1 + # Check the status code in the http.response.start message + status_sent = False + for call in mock_send.call_args_list: + msg = call[0][0] + if msg.get("type") == "http.response.start": + assert msg.get("status") == 403 + status_sent = True + assert status_sent, "Should have sent a 403 status" # 3b. Test POST Message Binding (No Auth - Security Fix) no_post_auth_context = ( @@ -179,10 +189,19 @@ async def test_session_id_capture_and_binding(): "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", AsyncMock(return_value=no_post_auth_context), ): - with pytest.raises(HTTPException) as exc: - await handle_sse_post_messages(post_scope, mock_receive, mock_send) - assert exc.value.status_code == 403 - assert "Authentication mismatch" in exc.value.detail + # Reset mock_send calls + mock_send.reset_mock() + await handle_sse_post_messages(post_scope, mock_receive, mock_send) + + # Verify 403 was sent + assert mock_send.call_count >= 1 + status_sent = False + for call in mock_send.call_args_list: + msg = call[0][0] + if msg.get("type") == "http.response.start": + assert msg.get("status") == 403 + status_sent = True + assert status_sent, "Should have sent a 403 status" # Finish the run finish_run.set()