mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(mcp-oauth2): address greptile review - random state nonce, provider error body, polling fixes, GET Content-Type, tests
This commit is contained in:
parent
2c8d9d0f32
commit
0b6fa2b4b1
4 changed files with 295 additions and 10 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
},
|
||||
});
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue