From 94b350be8f517fd4f51345dc93f6c0eff46fffae Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Thu, 5 Mar 2026 18:28:50 -0800 Subject: [PATCH] feat(mcp): add OAuth2 authorization-code flow for OpenAPI MCPs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds three new endpoints so users can authorize their own accounts through a provider's OAuth2 consent screen (GitHub, Spotify, Linear, etc.) instead of pasting static API keys for BYOK OpenAPI MCP servers. - GET /v1/mcp/server/{server_id}/oauth2/connect — initiates the flow, returns an authorization_url the UI opens as a popup - GET /v1/mcp/oauth2/callback — receives code+state from provider, exchanges for access token, stores in LiteLLM_MCPUserCredentials - GET /v1/mcp/server/{server_id}/oauth2/status — returns connected:true/false State tokens are HMAC-SHA256 signed with master_key and expire after 10 min. The callback shows a success HTML page (with auto-close for popups) instead of redirecting, avoiding 404s in environments without a full UI deploy. --- .../mcp_server/openapi_oauth2_endpoints.py | 452 ++++++++++++++++++ 1 file changed, 452 insertions(+) create mode 100644 litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py diff --git a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py new file mode 100644 index 00000000000..f8841cab25d --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -0,0 +1,452 @@ +""" +OAuth2 authorization code flow endpoints for OpenAPI-backed MCP servers. + +These endpoints let users authorize LiteLLM to call an OAuth2-protected API +(e.g. GitHub, Spotify) on their behalf, rather than entering a static key. + +Endpoints: + GET /v1/mcp/server/{server_id}/oauth2/connect — initiate the OAuth2 flow + GET /v1/mcp/oauth2/callback — receive the code from the provider + GET /v1/mcp/server/{server_id}/oauth2/status — check if user is connected +""" + +import base64 +import hashlib +import hmac +import html as _html_module +import time +from typing import Dict, Optional +from urllib.parse import parse_qs, urlencode + +import httpx +from fastapi import APIRouter, Depends, HTTPException, Query, Request +from fastapi.responses import HTMLResponse, JSONResponse, Response + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._experimental.mcp_server.db import ( + has_user_credential, + store_user_credential, +) +from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, +) +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +# --------------------------------------------------------------------------- +# In-memory state store for pending OAuth2 flows. +# Each entry: {state: {server_id, user_id, timestamp, expires_at}} +# --------------------------------------------------------------------------- +_pending_oauth2_states: Dict[str, dict] = {} + +_STATE_TTL_SECONDS = 600 # 10 minutes +_STATES_MAX_SIZE = 1000 + +router = APIRouter(tags=["mcp"]) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _purge_expired_states() -> None: + now = time.time() + expired = [k for k, v in _pending_oauth2_states.items() if v["expires_at"] < now] + for k in expired: + del _pending_oauth2_states[k] + + +def _make_state_token(server_id: str, user_id: str, timestamp: float, master_key: str) -> str: + """HMAC-SHA256 of '{server_id}:{user_id}:{timestamp}' base64url-encoded.""" + message = f"{server_id}:{user_id}:{timestamp}".encode() + digest = hmac.new(master_key.encode(), message, hashlib.sha256).digest() + return base64.urlsafe_b64encode(digest).rstrip(b"=").decode() + + +def _build_success_html(server_name: str) -> str: + e = _html_module.escape + return f""" + + + +Connected — LiteLLM + + + + +
+
✓
+

Connected to {e(server_name)}

+

Authorization successful. You can close this window.

+
+ +""" + + +def _build_error_html(title: str, message: str) -> str: + e = _html_module.escape + return f""" + + + +{e(title)} — LiteLLM + + + +
+

{e(title)}

+

{e(message)}

+
+ +""" + + +# --------------------------------------------------------------------------- +# Endpoints +# --------------------------------------------------------------------------- + + +@router.get("/v1/mcp/server/{server_id}/oauth2/connect", include_in_schema=False) +async def openapi_oauth2_connect( + server_id: str, + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> JSONResponse: + from litellm.proxy.proxy_server import master_key + + server = global_mcp_server_manager.get_mcp_server_by_id(server_id) + if server is None: + raise HTTPException(status_code=404, detail=f"MCP server '{server_id}' not found") + + if not server.authorization_url: + raise HTTPException( + status_code=400, + detail=f"Server '{server_id}' has no authorization_url configured", + ) + if not server.token_url: + raise HTTPException( + status_code=400, + detail=f"Server '{server_id}' has no token_url configured", + ) + if not server.client_id: + raise HTTPException( + status_code=400, + detail=f"Server '{server_id}' has no client_id configured", + ) + if master_key is None: + raise HTTPException(status_code=500, detail="Master key not configured") + + user_id = user_api_key_dict.user_id or user_api_key_dict.api_key or "" + if not user_id: + raise HTTPException(status_code=400, detail="Cannot determine user identity from token") + + _purge_expired_states() + if len(_pending_oauth2_states) >= _STATES_MAX_SIZE: + raise HTTPException(status_code=503, detail="Too many pending OAuth2 flows") + + timestamp = time.time() + state = _make_state_token(server_id, user_id, timestamp, master_key) + _pending_oauth2_states[state] = { + "server_id": server_id, + "user_id": user_id, + "timestamp": timestamp, + "expires_at": timestamp + _STATE_TTL_SECONDS, + } + + base_url = get_request_base_url(request) + callback_url = f"{base_url}/v1/mcp/oauth2/callback" + + params: dict = { + "client_id": server.client_id, + "redirect_uri": callback_url, + "response_type": "code", + "state": state, + } + if server.scopes: + params["scope"] = " ".join(server.scopes) + + authorization_url = f"{server.authorization_url}?{urlencode(params)}" + server_name = server.server_name or server.name or server_id + + verbose_proxy_logger.debug( + "openapi_oauth2_connect: user=%s server=%s", + user_id, + server_id, + ) + + return JSONResponse( + { + "authorization_url": authorization_url, + "server_id": server_id, + "server_name": server_name, + } + ) + + +@router.get("/v1/mcp/oauth2/callback", include_in_schema=False, response_model=None) +async def openapi_oauth2_callback( + request: Request, + code: Optional[str] = Query(default=None), + state: Optional[str] = Query(default=None), + error: Optional[str] = Query(default=None), + error_description: Optional[str] = Query(default=None), +) -> Response: + from litellm.proxy.proxy_server import prisma_client + + if error: + msg = error_description or error + return HTMLResponse( + content=_build_error_html( + "Authorization failed", + f"The provider returned an error: {msg}", + ), + status_code=400, + ) + + if not code or not state: + return HTMLResponse( + content=_build_error_html( + "Missing parameters", + "Expected 'code' and 'state' query parameters.", + ), + status_code=400, + ) + + _purge_expired_states() + state_data = _pending_oauth2_states.get(state) + if state_data is None: + return HTMLResponse( + content=_build_error_html( + "Invalid state", + "The OAuth2 state token is invalid or has already been used.", + ), + status_code=400, + ) + if time.time() > state_data["expires_at"]: + del _pending_oauth2_states[state] + return HTMLResponse( + content=_build_error_html( + "State expired", + "The OAuth2 state token has expired. Please start the connection flow again.", + ), + status_code=400, + ) + + # Consume the state (one-time use) + del _pending_oauth2_states[state] + + server_id: str = state_data["server_id"] + user_id: str = state_data["user_id"] + + server = global_mcp_server_manager.get_mcp_server_by_id(server_id) + if server is None: + return HTMLResponse( + content=_build_error_html( + "Server not found", + f"MCP server '{server_id}' could not be found.", + ), + status_code=404, + ) + if not server.token_url: + return HTMLResponse( + content=_build_error_html( + "Configuration error", + f"Server '{server_id}' has no token_url configured.", + ), + status_code=500, + ) + + base_url = get_request_base_url(request) + callback_url = f"{base_url}/v1/mcp/oauth2/callback" + + token_request_data = { + "client_id": server.client_id or "", + "client_secret": server.client_secret or "", + "code": code, + "redirect_uri": callback_url, + "grant_type": "authorization_code", + } + + try: + async with httpx.AsyncClient() as client: + response = await client.post( + server.token_url, + data=token_request_data, + headers={"Accept": "application/json"}, + timeout=30.0, + ) + response.raise_for_status() + except httpx.HTTPStatusError as exc: + verbose_proxy_logger.error( + "openapi_oauth2_callback: token exchange HTTP error user=%s server=%s status=%s", + user_id, + server_id, + exc.response.status_code, + ) + return HTMLResponse( + content=_build_error_html( + "Token exchange failed", + f"The provider returned HTTP {exc.response.status_code} during token exchange.", + ), + status_code=502, + ) + except Exception as exc: + verbose_proxy_logger.error( + "openapi_oauth2_callback: token exchange error user=%s server=%s: %s", + user_id, + server_id, + exc, + ) + return HTMLResponse( + content=_build_error_html( + "Token exchange failed", + "An unexpected error occurred during token exchange.", + ), + status_code=502, + ) + + # Parse response: try JSON first, fall back to URL-encoded form (GitHub can return either) + access_token: Optional[str] = None + content_type = response.headers.get("content-type", "") + if "application/json" in content_type: + try: + token_data = response.json() + access_token = token_data.get("access_token") + except Exception: + pass + if access_token is None: + # URL-encoded form fallback (GitHub without Accept: application/json header) + try: + form_data = parse_qs(response.text) + tokens = form_data.get("access_token", []) + if tokens: + access_token = tokens[0] + except Exception: + pass + + if not access_token: + verbose_proxy_logger.error( + "openapi_oauth2_callback: no access_token in response user=%s server=%s", + user_id, + server_id, + ) + return HTMLResponse( + content=_build_error_html( + "Token exchange failed", + "The provider did not return an access token. Check client credentials and scopes.", + ), + status_code=502, + ) + + if prisma_client is None: + verbose_proxy_logger.warning( + "openapi_oauth2_callback: prisma_client is None — credential not persisted user=%s server=%s", + user_id, + server_id, + ) + return HTMLResponse( + content=_build_error_html( + "Database unavailable", + "Cannot persist credentials: database is not configured.", + ), + status_code=500, + ) + + try: + await store_user_credential( + prisma_client=prisma_client, + user_id=user_id, + server_id=server_id, + credential=access_token, + ) + from litellm.proxy._experimental.mcp_server.server import ( + _invalidate_byok_cred_cache, + ) + + _invalidate_byok_cred_cache(user_id, server_id) + except Exception as exc: + verbose_proxy_logger.error( + "openapi_oauth2_callback: failed to store credential user=%s server=%s: %s", + user_id, + server_id, + exc, + ) + return HTMLResponse( + content=_build_error_html( + "Storage error", + "Failed to store the access token. Please try again.", + ), + status_code=500, + ) + + verbose_proxy_logger.info( + "openapi_oauth2_callback: connected user=%s server=%s", + user_id, + server_id, + ) + server_name = server.server_name or server.name or server_id + return HTMLResponse(content=_build_success_html(server_name), status_code=200) + + +@router.get("/v1/mcp/server/{server_id}/oauth2/status", include_in_schema=False) +async def openapi_oauth2_status( + server_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> JSONResponse: + from litellm.proxy.proxy_server import prisma_client + + server = global_mcp_server_manager.get_mcp_server_by_id(server_id) + if server is None: + raise HTTPException(status_code=404, detail=f"MCP server '{server_id}' not found") + + user_id = user_api_key_dict.user_id or user_api_key_dict.api_key or "" + server_name = server.server_name or server.name or server_id + + if prisma_client is None or not user_id: + return JSONResponse( + {"connected": False, "server_id": server_id, "server_name": server_name} + ) + + try: + connected = await has_user_credential( + prisma_client=prisma_client, + user_id=user_id, + server_id=server_id, + ) + except Exception as exc: + verbose_proxy_logger.error( + "openapi_oauth2_status: error checking credential user=%s server=%s: %s", + user_id, + server_id, + exc, + ) + connected = False + + return JSONResponse( + {"connected": connected, "server_id": server_id, "server_name": server_name} + )