From 0b6fa2b4b1a2908608fb94bfdf56317ef3b963da Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Fri, 6 Mar 2026 18:04:28 -0800 Subject: [PATCH] fix(mcp-oauth2): address greptile review - random state nonce, provider error body, polling fixes, GET Content-Type, tests --- .../mcp_server/openapi_oauth2_endpoints.py | 61 ++++- .../test_openapi_oauth2_endpoints.py | 224 ++++++++++++++++++ .../mcp_tools/OAuth2ConnectButton.tsx | 18 +- .../src/components/networking.tsx | 2 - 4 files changed, 295 insertions(+), 10 deletions(-) create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_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 index f8841cab25d..426b072b473 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -14,6 +14,7 @@ import base64 import hashlib import hmac import html as _html_module +import os import time from typing import Dict, Optional from urllib.parse import parse_qs, urlencode @@ -39,6 +40,13 @@ 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}} +# +# NOTE: This is an in-process dict. In horizontally-scaled deployments +# (multiple uvicorn workers, Kubernetes pods, etc.) the /connect request may +# be handled by one instance while the provider's redirect hits /callback on +# a different instance — that second instance won't find the state and the +# flow will fail. A follow-up should persist state in a shared store +# (e.g. LiteLLM_MCPUserCredentials with a sentinel user_id, or Redis cache). # --------------------------------------------------------------------------- _pending_oauth2_states: Dict[str, dict] = {} @@ -61,8 +69,14 @@ def _purge_expired_states() -> None: 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() + """HMAC-SHA256 of '{server_id}:{user_id}:{timestamp}:{nonce}' base64url-encoded. + + A 16-byte random nonce is appended so that concurrent flows for the same + (server_id, user_id) pair — e.g. from multiple open browser tabs — each + produce a unique state token instead of colliding in _pending_oauth2_states. + """ + nonce = base64.urlsafe_b64encode(os.urandom(16)).rstrip(b"=").decode() + message = f"{server_id}:{user_id}:{timestamp}:{nonce}".encode() digest = hmac.new(master_key.encode(), message, hashlib.sha256).digest() return base64.urlsafe_b64encode(digest).rstrip(b"=").decode() @@ -158,6 +172,11 @@ async def openapi_oauth2_connect( status_code=400, detail=f"Server '{server_id}' has no client_id configured", ) + if not server.client_secret: + raise HTTPException( + status_code=400, + detail=f"Server '{server_id}' has no client_secret configured", + ) if master_key is None: raise HTTPException(status_code=500, detail="Master key not configured") @@ -331,24 +350,52 @@ async def openapi_oauth2_callback( ) # Parse response: try JSON first, fall back to URL-encoded form (GitHub can return either) + # Some providers return HTTP 200 with an error body, so check for error fields explicitly. access_token: Optional[str] = None + provider_error: 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") + if token_data.get("error"): + err = token_data["error"] + err_desc = token_data.get("error_description", "") + provider_error = f"{err}: {err_desc}" if err_desc else err + else: + access_token = token_data.get("access_token") except Exception: pass - if access_token is None: + if access_token is None and provider_error 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] + if form_data.get("error"): + err = form_data["error"][0] + err_desc_list = form_data.get("error_description", []) + err_desc = err_desc_list[0] if err_desc_list else "" + provider_error = f"{err}: {err_desc}" if err_desc else err + else: + tokens = form_data.get("access_token", []) + if tokens: + access_token = tokens[0] except Exception: pass + if provider_error: + verbose_proxy_logger.error( + "openapi_oauth2_callback: provider error user=%s server=%s error=%s", + user_id, + server_id, + provider_error, + ) + return HTMLResponse( + content=_build_error_html( + "Token exchange failed", + f"The provider returned an error: {provider_error}", + ), + status_code=502, + ) + if not access_token: verbose_proxy_logger.error( "openapi_oauth2_callback: no access_token in response user=%s server=%s", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py new file mode 100644 index 00000000000..7994fe8a5cc --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py @@ -0,0 +1,224 @@ +"""Unit tests for openapi_oauth2_endpoints.py""" + +import sys +import time +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, "../../../../../") + +pytest.importorskip("mcp", reason="mcp package not installed; skipping MCP OAuth2 tests") + +from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + _make_state_token, + _pending_oauth2_states, + _purge_expired_states, +) + +# --------------------------------------------------------------------------- +# _make_state_token +# --------------------------------------------------------------------------- + + +def test_make_state_token_returns_string(): + token = _make_state_token("server1", "user1", time.time(), "master-key") + assert isinstance(token, str) + assert len(token) > 0 + + +def test_make_state_token_is_unique_for_same_inputs(): + """Two calls with identical inputs must produce different tokens (random nonce).""" + ts = time.time() + t1 = _make_state_token("server1", "user1", ts, "master-key") + t2 = _make_state_token("server1", "user1", ts, "master-key") + assert t1 != t2, "Tokens should differ due to random nonce" + + +def test_make_state_token_differs_by_server(): + ts = time.time() + t1 = _make_state_token("server1", "user1", ts, "master-key") + t2 = _make_state_token("server2", "user1", ts, "master-key") + assert t1 != t2 + + +# --------------------------------------------------------------------------- +# _purge_expired_states +# --------------------------------------------------------------------------- + + +def test_purge_expired_states_removes_expired(): + _pending_oauth2_states.clear() + now = time.time() + _pending_oauth2_states["expired"] = {"expires_at": now - 1} + _pending_oauth2_states["valid"] = {"expires_at": now + 600} + _purge_expired_states() + assert "expired" not in _pending_oauth2_states + assert "valid" in _pending_oauth2_states + _pending_oauth2_states.clear() + + +# --------------------------------------------------------------------------- +# openapi_oauth2_connect — validation +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_missing_server_raises_404(): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + openapi_oauth2_connect, + ) + + mock_request = MagicMock() + mock_user = MagicMock() + mock_user.user_id = "user1" + mock_user.api_key = None + + with patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.global_mcp_server_manager" + ) as mock_mgr: + mock_mgr.get_mcp_server_by_id.return_value = None + with pytest.raises(HTTPException) as exc_info: + await openapi_oauth2_connect("nonexistent", mock_request, mock_user) + assert exc_info.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_connect_missing_client_secret_raises_400(): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + openapi_oauth2_connect, + ) + + mock_request = MagicMock() + mock_user = MagicMock() + mock_user.user_id = "user1" + mock_user.api_key = None + + mock_server = MagicMock() + mock_server.authorization_url = "https://provider.example/auth" + mock_server.token_url = "https://provider.example/token" + mock_server.client_id = "my-client-id" + mock_server.client_secret = None # missing + + with patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.global_mcp_server_manager" + ) as mock_mgr, patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.master_key", + "sk-test", + create=True, + ): + mock_mgr.get_mcp_server_by_id.return_value = mock_server + with pytest.raises(HTTPException) as exc_info: + await openapi_oauth2_connect("server1", mock_request, mock_user) + assert exc_info.value.status_code == 400 + assert "client_secret" in exc_info.value.detail + + +# --------------------------------------------------------------------------- +# openapi_oauth2_callback — token parsing (provider error body) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_callback_provider_error_in_json_body(): + """HTTP 200 with JSON error body should render error page, not succeed.""" + from fastapi.responses import HTMLResponse + + from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + openapi_oauth2_callback, + ) + + state = "test-state-123" + now = time.time() + _pending_oauth2_states[state] = { + "server_id": "server1", + "user_id": "user1", + "timestamp": now, + "expires_at": now + 600, + } + + mock_server = MagicMock() + mock_server.token_url = "https://provider.example/token" + mock_server.client_id = "cid" + mock_server.client_secret = "csecret" + + # Simulate provider returning HTTP 200 with an error body + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = { + "error": "invalid_grant", + "error_description": "Code already used", + } + mock_response.raise_for_status = MagicMock() # does not raise + + mock_request = MagicMock() + mock_request.base_url = "http://localhost:4000" + + with patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.global_mcp_server_manager" + ) as mock_mgr, patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.get_request_base_url", + return_value="http://localhost:4000", + ), patch( + "httpx.AsyncClient" + ) as mock_client_cls: + mock_mgr.get_mcp_server_by_id.return_value = mock_server + mock_async_client = AsyncMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + mock_client_cls.return_value.__aenter__ = AsyncMock(return_value=mock_async_client) + mock_client_cls.return_value.__aexit__ = AsyncMock(return_value=None) + + result = await openapi_oauth2_callback( + request=mock_request, + code="auth-code", + state=state, + error=None, + error_description=None, + ) + + assert isinstance(result, HTMLResponse) + assert result.status_code == 502 + body_bytes = bytes(result.body) if not isinstance(result.body, bytes) else result.body + assert b"invalid_grant" in body_bytes or b"provider" in body_bytes.lower() + + +# --------------------------------------------------------------------------- +# openapi_oauth2_status — basic happy path +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_status_no_prisma_returns_not_connected(): + from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + openapi_oauth2_status, + ) + + mock_user = MagicMock() + mock_user.user_id = "user1" + mock_user.api_key = None + + mock_server = MagicMock() + mock_server.server_name = "GitHub" + mock_server.name = "github" + + with patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.global_mcp_server_manager" + ) as mock_mgr, patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.prisma_client", + None, + create=True, + ): + mock_mgr.get_mcp_server_by_id.return_value = mock_server + result = await openapi_oauth2_status("server1", mock_user) + + import json + + raw = result.body + body = json.loads(raw.decode() if isinstance(raw, (bytes, bytearray)) else str(raw)) + assert body["connected"] is False + assert body["server_id"] == "server1" diff --git a/ui/litellm-dashboard/src/components/mcp_tools/OAuth2ConnectButton.tsx b/ui/litellm-dashboard/src/components/mcp_tools/OAuth2ConnectButton.tsx index 853a757b871..d1e6be825b2 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/OAuth2ConnectButton.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/OAuth2ConnectButton.tsx @@ -12,6 +12,7 @@ interface OAuth2ConnectButtonProps { } const POLL_INTERVAL_MS = 2000; +const MAX_POLL_MS = 10 * 60 * 1000; // 10 minutes export const OAuth2ConnectButton: React.FC = ({ server, @@ -22,12 +23,14 @@ export const OAuth2ConnectButton: React.FC = ({ const [error, setError] = useState(null); const popupRef = useRef(null); const pollTimerRef = useRef | null>(null); + const pollStartRef = useRef(null); const stopPolling = () => { if (pollTimerRef.current !== null) { clearInterval(pollTimerRef.current); pollTimerRef.current = null; } + pollStartRef.current = null; }; const handleConnected = () => { @@ -42,7 +45,20 @@ export const OAuth2ConnectButton: React.FC = ({ }; const startPolling = () => { + // Always clear any existing interval before starting a new one to avoid + // double-polling if the button is clicked while a previous flow is still active. + stopPolling(); + pollStartRef.current = Date.now(); pollTimerRef.current = setInterval(async () => { + // Enforce a maximum polling duration to avoid indefinite requests when + // the popup is left open but the OAuth flow never completes. + if (pollStartRef.current !== null && Date.now() - pollStartRef.current > MAX_POLL_MS) { + stopPolling(); + setLoading(false); + setError("OAuth2 connection timed out. Please try again."); + return; + } + // Stop if popup was closed by user if (popupRef.current && popupRef.current.closed) { stopPolling(); @@ -56,7 +72,7 @@ export const OAuth2ConnectButton: React.FC = ({ handleConnected(); } } catch { - // Ignore polling errors; keep trying until popup is closed + // Ignore polling errors; keep trying until popup is closed or timeout } }, POLL_INTERVAL_MS); }; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index a7c1a418aba..4f63cbe2d92 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -6408,7 +6408,6 @@ export const getMcpOAuth2ConnectUrl = async ( method: HTTP_REQUEST.GET, headers: { [globalLitellmHeaderName]: `Bearer ${accessToken}`, - "Content-Type": "application/json", }, }); @@ -6439,7 +6438,6 @@ export const getMcpOAuth2Status = async ( method: HTTP_REQUEST.GET, headers: { [globalLitellmHeaderName]: `Bearer ${accessToken}`, - "Content-Type": "application/json", }, });