mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(mcp): support public-client / PKCE in OAuth metadata, /register, and /token
This commit is contained in:
parent
144279eb57
commit
44f6da9693
2 changed files with 502 additions and 43 deletions
|
|
@ -2,6 +2,7 @@ import json
|
|||
from typing import Any, Dict, Optional
|
||||
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
|
||||
|
||||
|
|
@ -378,6 +379,50 @@ async def authorize_with_server(
|
|||
return RedirectResponse(final_url)
|
||||
|
||||
|
||||
async def _post_to_upstream_token_endpoint(token_url: str, token_data: Dict[str, Any]):
|
||||
"""POST to the upstream IdP /token endpoint.
|
||||
|
||||
Uses a raw httpx client (instead of get_async_httpx_client) so we can inspect
|
||||
4xx/5xx responses and surface them as proper OAuth2 error JSON per RFC 6749
|
||||
§5.2. The wrapped client raises MaskedHTTPStatusError on raise_for_status()
|
||||
before downstream code can read the response body, so the actual
|
||||
AADSTS<code> / error_description from Entra/Okta is lost and the client sees
|
||||
a generic 500.
|
||||
|
||||
Returns the parsed JSON dict on success, or a JSONResponse with the upstream
|
||||
error body on 4xx/5xx (or a transport failure).
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30.0) as raw_client:
|
||||
response = await raw_client.post(
|
||||
token_url,
|
||||
headers={"Accept": "application/json"},
|
||||
data=token_data,
|
||||
)
|
||||
except httpx.HTTPError as exc:
|
||||
return JSONResponse(
|
||||
{"error": "server_error", "error_description": str(exc)},
|
||||
status_code=502,
|
||||
headers=TOKEN_NO_CACHE_HEADERS,
|
||||
)
|
||||
|
||||
if response.status_code >= 400:
|
||||
try:
|
||||
err_body = response.json()
|
||||
except ValueError:
|
||||
err_body = {
|
||||
"error": "invalid_request",
|
||||
"error_description": response.text or "upstream token endpoint error",
|
||||
}
|
||||
return JSONResponse(
|
||||
err_body,
|
||||
status_code=response.status_code,
|
||||
headers=TOKEN_NO_CACHE_HEADERS,
|
||||
)
|
||||
|
||||
return response.json()
|
||||
|
||||
|
||||
async def exchange_token_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -400,6 +445,12 @@ async def exchange_token_with_server(
|
|||
resolved_client_secret = (
|
||||
mcp_server.client_secret if mcp_server.client_secret else client_secret
|
||||
)
|
||||
# Drop placeholder / empty secrets that come from LiteLLM's dummy /register
|
||||
# response (or from in-flight clients that already cached "dummy" before
|
||||
# the public-client branch deployed). Forwarding these to a real IdP yields
|
||||
# a 401.
|
||||
if resolved_client_secret in (None, "", "dummy"):
|
||||
resolved_client_secret = None
|
||||
|
||||
if grant_type == "refresh_token":
|
||||
if not refresh_token:
|
||||
|
|
@ -434,15 +485,13 @@ async def exchange_token_with_server(
|
|||
if code_verifier:
|
||||
token_data["code_verifier"] = code_verifier
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
headers={"Accept": "application/json"},
|
||||
data=token_data,
|
||||
upstream_result = await _post_to_upstream_token_endpoint(
|
||||
mcp_server.token_url, token_data
|
||||
)
|
||||
if isinstance(upstream_result, JSONResponse):
|
||||
return upstream_result
|
||||
|
||||
response.raise_for_status()
|
||||
token_response = response.json()
|
||||
token_response = upstream_result
|
||||
access_token = token_response["access_token"]
|
||||
|
||||
# Validate token response against server-configured rules before any storage.
|
||||
|
|
@ -514,6 +563,16 @@ async def register_client_with_server(
|
|||
"redirect_uris": [f"{request_base_url}/callback"],
|
||||
}
|
||||
|
||||
# Public-client / PKCE-broker mode: configured client_id but no client_secret.
|
||||
# Return the real client_id and OMIT client_secret so the downstream MCP
|
||||
# client doesn't cache a placeholder and forward it to the upstream IdP.
|
||||
if mcp_server.client_id and not mcp_server.client_secret:
|
||||
return {
|
||||
"client_id": mcp_server.client_id,
|
||||
"redirect_uris": [f"{request_base_url}/callback"],
|
||||
"token_endpoint_auth_method": "none",
|
||||
}
|
||||
|
||||
if mcp_server.client_id and mcp_server.client_secret:
|
||||
return dummy_return
|
||||
|
||||
|
|
@ -871,6 +930,17 @@ def _build_oauth_authorization_server_response(
|
|||
mcp_server_name, client_ip=client_ip
|
||||
)
|
||||
|
||||
# Per RFC 8414 §2: token_endpoint_auth_methods_supported declares which
|
||||
# client-authentication methods this auth server accepts at /token.
|
||||
# Public clients (PKCE-broker, no client_secret stored) must advertise
|
||||
# "none". Fall back to "client_secret_post" when the server can't be
|
||||
# resolved (root /.well-known or unknown server name) — preserves legacy
|
||||
# behavior for that case.
|
||||
if mcp_server and mcp_server.client_id and not mcp_server.client_secret:
|
||||
token_endpoint_auth_methods_supported = ["none"]
|
||||
else:
|
||||
token_endpoint_auth_methods_supported = ["client_secret_post"]
|
||||
|
||||
return {
|
||||
"issuer": request_base_url, # point to your proxy
|
||||
"authorization_endpoint": authorization_endpoint,
|
||||
|
|
@ -881,7 +951,7 @@ def _build_oauth_authorization_server_response(
|
|||
),
|
||||
"grant_types_supported": ["authorization_code", "refresh_token"],
|
||||
"code_challenge_methods_supported": ["S256"],
|
||||
"token_endpoint_auth_methods_supported": ["client_secret_post"],
|
||||
"token_endpoint_auth_methods_supported": token_endpoint_auth_methods_supported,
|
||||
# Claude expects a registration endpoint, even if we just fake it
|
||||
"registration_endpoint": (
|
||||
f"{request_base_url}/{mcp_server_name}/register"
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Tests for MCP OAuth discoverable endpoints"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -285,25 +286,24 @@ async def test_token_endpoint_forwards_code_verifier():
|
|||
|
||||
# Mock httpx client response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "ya29.test_access_token",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3599,
|
||||
"scope": "openid email https://www.googleapis.com/auth/drive",
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
# Mock the async httpx client with AsyncMock for async methods
|
||||
from unittest.mock import AsyncMock
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
mock_async_client = MagicMock()
|
||||
# Use AsyncMock for the async post method
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
mock_get_client.return_value = mock_async_client
|
||||
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
# Call token endpoint with code_verifier
|
||||
response = await token_endpoint(
|
||||
request=mock_request,
|
||||
|
|
@ -632,22 +632,23 @@ async def test_token_endpoint_respects_x_forwarded_proto():
|
|||
|
||||
# Mock httpx client response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "test_token",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3599,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
# Mock the async httpx client
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
mock_get_client.return_value = mock_async_client
|
||||
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
await token_endpoint(
|
||||
request=mock_request,
|
||||
grant_type="authorization_code",
|
||||
|
|
@ -937,22 +938,23 @@ async def test_token_endpoint_respects_x_forwarded_host():
|
|||
|
||||
# Mock httpx client response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "test_token",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3599,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
# Mock the async httpx client
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
mock_get_client.return_value = mock_async_client
|
||||
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
await token_endpoint(
|
||||
request=mock_request,
|
||||
grant_type="authorization_code",
|
||||
|
|
@ -1567,22 +1569,24 @@ async def test_token_root_resolves_single_oauth2_server():
|
|||
mock_request.headers = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "ya29.test_token",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3599,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
mock_get_client.return_value = mock_async_client
|
||||
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
# Call /token WITHOUT mcp_server_name
|
||||
response = await token_endpoint(
|
||||
request=mock_request,
|
||||
|
|
@ -2087,22 +2091,24 @@ async def test_token_endpoint_refresh_token_grant():
|
|||
|
||||
# Mock httpx client response with new tokens
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "new_access_token",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3599,
|
||||
"refresh_token": "new_refresh_token",
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
mock_get_client.return_value = mock_async_client
|
||||
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
response = await token_endpoint(
|
||||
request=mock_request,
|
||||
grant_type="refresh_token",
|
||||
|
|
@ -2384,18 +2390,21 @@ async def test_token_endpoint_sets_no_store_cache_control():
|
|||
mock_request.headers = {}
|
||||
|
||||
fake_http_response = MagicMock()
|
||||
fake_http_response.status_code = 200
|
||||
fake_http_response.json.return_value = {
|
||||
"access_token": "tok",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
fake_http_response.raise_for_status = MagicMock()
|
||||
fake_http_client = MagicMock()
|
||||
fake_http_client.post = AsyncMock(return_value=fake_http_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=fake_http_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
|
||||
return_value=fake_http_client,
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
response = await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
|
|
@ -2410,3 +2419,383 @@ async def test_token_endpoint_sets_no_store_cache_control():
|
|||
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
assert response.headers["pragma"] == "no-cache"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_client_public_client_returns_real_client_id_and_no_secret():
|
||||
"""Bug 1 fix: when client_id is set without client_secret, /register returns
|
||||
the real client_id, no client_secret, and token_endpoint_auth_method=none."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
register_client,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
public_server = MCPServer(
|
||||
server_id="public_server",
|
||||
name="public_server",
|
||||
server_name="public_server",
|
||||
alias="public_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-entra-client-id",
|
||||
client_secret=None,
|
||||
authorization_url="https://login.microsoftonline.com/tid/oauth2/v2.0/authorize",
|
||||
token_url="https://login.microsoftonline.com/tid/oauth2/v2.0/token",
|
||||
)
|
||||
global_mcp_server_manager.registry[public_server.server_id] = public_server
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={}),
|
||||
):
|
||||
result = await register_client(
|
||||
request=mock_request, mcp_server_name="public_server"
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert result == {
|
||||
"client_id": "real-entra-client-id",
|
||||
"redirect_uris": ["https://proxy.litellm.example/callback"],
|
||||
"token_endpoint_auth_method": "none",
|
||||
}
|
||||
assert "client_secret" not in result
|
||||
|
||||
|
||||
def test_oauth_authorization_server_metadata_advertises_none_for_public_client():
|
||||
"""Bug 2 fix: token_endpoint_auth_methods_supported includes 'none' when the
|
||||
server has client_id but no client_secret."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_build_oauth_authorization_server_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
public_server = MCPServer(
|
||||
server_id="public_server",
|
||||
name="public_server",
|
||||
server_name="public_server",
|
||||
alias="public_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-entra-client-id",
|
||||
client_secret=None,
|
||||
authorization_url="https://login.microsoftonline.com/tid/oauth2/v2.0/authorize",
|
||||
token_url="https://login.microsoftonline.com/tid/oauth2/v2.0/token",
|
||||
)
|
||||
global_mcp_server_manager.registry[public_server.server_id] = public_server
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
result = _build_oauth_authorization_server_response(
|
||||
request=mock_request, mcp_server_name="public_server"
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert result["token_endpoint_auth_methods_supported"] == ["none"]
|
||||
|
||||
|
||||
def test_oauth_authorization_server_metadata_keeps_client_secret_post_for_confidential_client():
|
||||
"""Bug 2 fix regression guard: confidential client (both client_id +
|
||||
client_secret stored) keeps advertising client_secret_post."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_build_oauth_authorization_server_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
confidential_server = MCPServer(
|
||||
server_id="confidential_server",
|
||||
name="confidential_server",
|
||||
server_name="confidential_server",
|
||||
alias="confidential_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-client-id",
|
||||
client_secret="real-client-secret",
|
||||
authorization_url="https://provider.example/authorize",
|
||||
token_url="https://provider.example/token",
|
||||
)
|
||||
global_mcp_server_manager.registry[confidential_server.server_id] = (
|
||||
confidential_server
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
result = _build_oauth_authorization_server_response(
|
||||
request=mock_request, mcp_server_name="confidential_server"
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert result["token_endpoint_auth_methods_supported"] == ["client_secret_post"]
|
||||
|
||||
|
||||
def test_oauth_authorization_server_metadata_default_for_unresolved_server():
|
||||
"""Bug 2 fix regression guard: unknown server name and empty registry fall
|
||||
back to legacy ['client_secret_post']."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_build_oauth_authorization_server_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
result = _build_oauth_authorization_server_response(
|
||||
request=mock_request, mcp_server_name=None
|
||||
)
|
||||
|
||||
assert result["token_endpoint_auth_methods_supported"] == ["client_secret_post"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("placeholder", [None, "", "dummy"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_token_omits_placeholder_client_secret(placeholder):
|
||||
"""Cyrus Patch 2: client_secret in (None, '', 'dummy') is dropped from the
|
||||
upstream /token form body."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
server = MCPServer(
|
||||
server_id="public_server",
|
||||
name="public_server",
|
||||
server_name="public_server",
|
||||
alias="public_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-client-id",
|
||||
client_secret=None,
|
||||
authorization_url="https://provider.example/authorize",
|
||||
token_url="https://provider.example/token",
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "tok",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type="authorization_code",
|
||||
code="c",
|
||||
redirect_uri="http://127.0.0.1:3000/cb",
|
||||
client_id="real-client-id",
|
||||
client_secret=placeholder,
|
||||
code_verifier="cv",
|
||||
)
|
||||
|
||||
call_args = mock_async_client.post.call_args
|
||||
assert "client_secret" not in call_args.kwargs["data"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_token_surfaces_upstream_4xx_as_oauth_error_json():
|
||||
"""Cyrus Patch 3: when upstream /token returns 400 with AADSTS body, we
|
||||
return that body verbatim with the same status — not a generic 500."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
server = MCPServer(
|
||||
server_id="entra_server",
|
||||
name="entra_server",
|
||||
server_name="entra_server",
|
||||
alias="entra_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-client-id",
|
||||
client_secret=None,
|
||||
authorization_url="https://login.microsoftonline.com/tid/oauth2/v2.0/authorize",
|
||||
token_url="https://login.microsoftonline.com/tid/oauth2/v2.0/token",
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
upstream_body = {
|
||||
"error": "invalid_grant",
|
||||
"error_description": "AADSTS70008: The provided authorization code or refresh token has expired",
|
||||
}
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 400
|
||||
mock_response.json.return_value = upstream_body
|
||||
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
response = await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type="authorization_code",
|
||||
code="bad-code",
|
||||
redirect_uri="http://127.0.0.1:3000/cb",
|
||||
client_id="real-client-id",
|
||||
client_secret=None,
|
||||
code_verifier="cv",
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body) == upstream_body
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_token_handles_upstream_non_json_4xx():
|
||||
"""Cyrus Patch 3 edge: upstream returns 4xx with non-JSON body — we wrap it
|
||||
as RFC 6749 invalid_request with the raw text in error_description."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
server = MCPServer(
|
||||
server_id="entra_server",
|
||||
name="entra_server",
|
||||
server_name="entra_server",
|
||||
alias="entra_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-client-id",
|
||||
client_secret=None,
|
||||
authorization_url="https://login.microsoftonline.com/tid/oauth2/v2.0/authorize",
|
||||
token_url="https://login.microsoftonline.com/tid/oauth2/v2.0/token",
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 502
|
||||
mock_response.json.side_effect = ValueError("not json")
|
||||
mock_response.text = "<html>upstream gateway error</html>"
|
||||
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
fake_cm = MagicMock()
|
||||
fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client)
|
||||
fake_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient",
|
||||
return_value=fake_cm,
|
||||
):
|
||||
response = await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type="authorization_code",
|
||||
code="c",
|
||||
redirect_uri="http://127.0.0.1:3000/cb",
|
||||
client_id="real-client-id",
|
||||
client_secret=None,
|
||||
code_verifier="cv",
|
||||
)
|
||||
|
||||
assert response.status_code == 502
|
||||
body = json.loads(response.body)
|
||||
assert body["error"] == "invalid_request"
|
||||
assert "<html>upstream gateway error</html>" in body["error_description"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue