fix(mcp-oauth2): address greptile review - random state nonce, provider error body, polling fixes, GET Content-Type, tests

This commit is contained in:
Ishaan Jaffer 2026-03-06 18:04:28 -08:00
parent 2c8d9d0f32
commit 0b6fa2b4b1
4 changed files with 295 additions and 10 deletions

View file

@ -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",

View file

@ -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"

View file

@ -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<OAuth2ConnectButtonProps> = ({
server,
@ -22,12 +23,14 @@ export const OAuth2ConnectButton: React.FC<OAuth2ConnectButtonProps> = ({
const [error, setError] = useState<string | null>(null);
const popupRef = useRef<Window | null>(null);
const pollTimerRef = useRef<ReturnType<typeof setInterval> | null>(null);
const pollStartRef = useRef<number | null>(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<OAuth2ConnectButtonProps> = ({
};
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<OAuth2ConnectButtonProps> = ({
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);
};

View file

@ -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",
},
});