diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 1794cd14381..246e0d397d3 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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 / 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" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 85d5d6ba466..ef23c258154 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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 = "upstream gateway error" + + 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 "upstream gateway error" in body["error_description"]