From aa4f585e4c0155d21ccf492cde1a1063e63eb1ee Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 10 Jul 2026 00:26:15 -0700 Subject: [PATCH] feat(mcp): dcr_bridge authorize and token relay redirect handling with mandatory S256 --- .../mcp_server/discoverable_endpoints.py | 88 ++++++- .../mcp_server/test_discoverable_endpoints.py | 241 ++++++++++++++++++ 2 files changed, 328 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index f746a42080a..86c2bd9cea5 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -516,6 +516,67 @@ def _raise_unless_oauth2_discovery_server( ) +def _dcr_bridge_relays_client_registration(mcp_server: MCPServer) -> bool: + """True when a DCR-bridge server relays client registration to the upstream authorization + server instead of short-circuiting to an admin-configured OAuth client. In the relay arm the + upstream holds each client's own registration, so the authorize and token relays pass the + client's ``client_id`` and ``redirect_uri`` through verbatim and the authorization code + returns directly to the client's redirect URI without transiting the gateway. Gateway-side + redirect trust and the ``/callback`` state relay therefore only apply to the short-circuit + arm, where the upstream only knows the gateway's own callback.""" + return mcp_server.is_dcr_bridge and bool(mcp_server.registration_url) and not mcp_server.client_id + + +def _require_s256_pkce( + code_challenge: Optional[str], + code_challenge_method: Optional[str], +) -> Tuple[str, str]: + """DCR-bridge servers serve unauthenticated public OAuth clients, so the PKCE downgrade + paths (no challenge, or a non-S256 method; RFC 7636 defaults a missing method to ``plain``) + are rejected at the gateway instead of relying on upstream enforcement. Returns the + validated pair so callers get non-optional values.""" + if code_challenge and code_challenge_method == "S256": + return code_challenge, code_challenge_method + raise HTTPException( + status_code=400, + detail={ + "error": ( + "This server requires PKCE: send code_challenge with " + "code_challenge_method=S256 on the authorization request" + ) + }, + ) + + +def _redirect_to_upstream_authorize( + *, + mcp_server: MCPServer, + client_id: str, + redirect_uri: str, + state: str, + code_challenge: str, + code_challenge_method: str, + response_type: Optional[str], + scope: Optional[str], +) -> RedirectResponse: + """The bridge relay arm's authorize redirect: every client-supplied parameter passes through + to the upstream authorize endpoint verbatim, no relay state cookie is set, and the upstream + enforces its own registered redirect binding for the client.""" + scope_value = scope or (" ".join(mcp_server.scopes) if mcp_server.scopes else None) + passthrough_params = { + "client_id": client_id, + "redirect_uri": redirect_uri, + "state": state, + "response_type": response_type or "code", + "code_challenge": code_challenge, + "code_challenge_method": code_challenge_method, + **({"scope": scope_value} if scope_value else {}), + } + parsed_auth_url = urlparse(mcp_server.authorization_url or "") + merged_params = {**dict(parse_qsl(parsed_auth_url.query)), **passthrough_params} + return RedirectResponse(urlunparse(parsed_auth_url._replace(query=urlencode(merged_params)))) + + async def authorize_with_server( request: Request, mcp_server: MCPServer, @@ -531,6 +592,20 @@ async def authorize_with_server( if mcp_server.authorization_url is None: raise HTTPException(status_code=400, detail="MCP server authorization url is not set") + if mcp_server.is_dcr_bridge: + bridge_challenge, bridge_method = _require_s256_pkce(code_challenge, code_challenge_method) + if _dcr_bridge_relays_client_registration(mcp_server): + return _redirect_to_upstream_authorize( + mcp_server=mcp_server, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=bridge_challenge, + code_challenge_method=bridge_method, + response_type=response_type, + scope=scope, + ) + # Trusted redirect_uri: same-origin, loopback, or ops-allowlisted. # The URI is encrypted into the OAuth state and decoded on # /callback to redirect the user back; a non-trusted URI would be @@ -626,11 +701,22 @@ async def exchange_token_with_server( status_code=400, detail="code is required for authorization_code grant", ) + if _dcr_bridge_relays_client_registration(mcp_server) and not redirect_uri: + raise HTTPException( + status_code=400, + detail=( + "redirect_uri is required for the authorization_code grant on this server; " + "send the same redirect_uri used on the authorization request" + ), + ) proxy_base_url = get_request_base_url(request) + resolved_redirect_uri = ( + redirect_uri if _dcr_bridge_relays_client_registration(mcp_server) else f"{proxy_base_url}/callback" + ) token_data = { "grant_type": "authorization_code", "code": code, - "redirect_uri": f"{proxy_base_url}/callback", + "redirect_uri": resolved_redirect_uri, **client_auth.body, } if code_verifier: 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 c8e871b3bc1..5b43f39a7a4 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 @@ -3577,6 +3577,247 @@ async def test_token_exchange_passes_through_upstream_expires_in(): assert body["expires_in"] == 43200 +_BRIDGE_CLIENT_REDIRECT = "https://claude.ai/api/mcp/auth_callback" + + +def _bridge_server(**overrides): + from litellm.proxy._types import MCPTransport + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + fields = { + "server_id": "bridge_srv", + "name": "bridge_srv", + "server_name": "bridge_srv", + "alias": "bridge_srv", + "transport": MCPTransport.http, + "auth_type": MCPAuth.true_passthrough, + "dcr_bridge": True, + "authorization_url": "https://provider.com/oauth/authorize", + "token_url": "https://provider.com/oauth/token", + "registration_url": "https://provider.com/oauth/register", + **overrides, + } + return MCPServer(**fields) + + +def _bridge_mock_request(): + from fastapi import Request + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + return mock_request + + +@pytest.mark.asyncio +@pytest.mark.parametrize("auth_type_value", ["true_passthrough", "oauth_delegate"]) +async def test_authorize_bridge_relay_passes_client_params_verbatim(auth_type_value): + """The bridge relay arm (registration relayed upstream, no admin-configured client) passes the + client's client_id, redirect_uri, state, and PKCE through verbatim: the code returns straight + to the client's own redirect URI, so the gateway sets no state cookie, injects no /callback, + and applies no gateway-side redirect trust (the upstream enforces its registered binding).""" + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize_with_server, + ) + from litellm.types.mcp import MCPAuth + + response = await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=_bridge_server(auth_type=MCPAuth(auth_type_value)), + client_id="dcr-client-123", + redirect_uri=_BRIDGE_CLIENT_REDIRECT, + state="client-state", + code_challenge="chal", + code_challenge_method="S256", + ) + + assert response.status_code == 307 + location = response.headers["location"] + assert location.startswith("https://provider.com/oauth/authorize") + query = parse_qs(urlparse(location).query) + assert query["client_id"] == ["dcr-client-123"] + assert query["redirect_uri"] == [_BRIDGE_CLIENT_REDIRECT] + assert query["state"] == ["client-state"] + assert query["code_challenge"] == ["chal"] + assert query["code_challenge_method"] == ["S256"] + assert "litellm.example.com" not in location + assert "set-cookie" not in {key.lower() for key in response.headers.keys()} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "code_challenge,code_challenge_method", + [(None, None), ("chal", None), ("chal", "plain"), (None, "S256")], +) +async def test_authorize_bridge_requires_s256_pkce(code_challenge, code_challenge_method): + """Bridge servers serve unauthenticated public clients, so the PKCE downgrade paths (missing + challenge, or a method that is not S256; RFC 7636 defaults a missing method to plain) are + rejected at the gateway on both bridge arms.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize_with_server, + ) + + with pytest.raises(HTTPException) as exc: + await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=_bridge_server(), + client_id="dcr-client-123", + redirect_uri=_BRIDGE_CLIENT_REDIRECT, + state="s", + code_challenge=code_challenge, + code_challenge_method=code_challenge_method, + ) + + assert exc.value.status_code == 400 + assert "S256" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_authorize_bridge_short_circuit_keeps_callback_and_redirect_trust(): + """The bridge short-circuit arm (admin-configured OAuth client, upstream only knows the + gateway callback) keeps the /callback state relay and the gateway redirect trust: a public + client redirect target is rejected unless ops allowlist it, and a trusted target still routes + through the gateway callback with the state cookie.""" + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize_with_server, + ) + + short_circuit_server = _bridge_server(client_id="admin-client", registration_url=None) + + with pytest.raises(HTTPException) as exc: + await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=short_circuit_server, + client_id="ignored", + redirect_uri=_BRIDGE_CLIENT_REDIRECT, + state="s", + code_challenge="chal", + code_challenge_method="S256", + ) + assert exc.value.status_code in (400, 403) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper", + return_value="mocked_encrypted_state", + ): + response = await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=short_circuit_server, + client_id="ignored", + redirect_uri="http://127.0.0.1:60108/callback", + state="s", + code_challenge="chal", + code_challenge_method="S256", + ) + + query = parse_qs(urlparse(response.headers["location"]).query) + assert query["redirect_uri"] == ["https://litellm.example.com/callback"] + assert query["client_id"] == ["admin-client"] + + +@pytest.mark.asyncio +async def test_authorize_non_bridge_client_forwarded_keeps_pre_bridge_contract(): + """A client-forwarded server without dcr_bridge keeps the pre-bridge behavior: no PKCE + requirement and the gateway /callback relay (this is the browser-only Authorize path).""" + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize_with_server, + ) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper", + return_value="mocked_encrypted_state", + ): + response = await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=_bridge_server(dcr_bridge=None), + client_id="cid", + redirect_uri="http://127.0.0.1:60108/callback", + state="s", + ) + + assert response.status_code == 307 + query = parse_qs(urlparse(response.headers["location"]).query) + assert query["redirect_uri"] == ["https://litellm.example.com/callback"] + + +async def _bridge_token_post_data(server, redirect_uri): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + + fake_http_response = MagicMock() + fake_http_response.json.return_value = {"access_token": "tok", "token_type": "Bearer"} + fake_http_response.raise_for_status = MagicMock() + fake_http_client = MagicMock() + fake_http_client.post = AsyncMock(return_value=fake_http_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=fake_http_client, + ): + await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri=redirect_uri, + client_id="dcr-client-123", + client_secret=None, + code_verifier="verifier", + ) + return fake_http_client.post.call_args.kwargs["data"] + + +@pytest.mark.asyncio +async def test_token_bridge_relay_posts_client_redirect_uri(): + """The bridge relay arm's token exchange sends the client's own redirect_uri upstream (it must + match the authorize leg) with the caller's public client_id and PKCE verifier.""" + data = await _bridge_token_post_data(_bridge_server(), redirect_uri=_BRIDGE_CLIENT_REDIRECT) + + assert data["redirect_uri"] == _BRIDGE_CLIENT_REDIRECT + assert data["client_id"] == "dcr-client-123" + assert data["code_verifier"] == "verifier" + assert "client_secret" not in data + + +@pytest.mark.asyncio +async def test_token_bridge_relay_requires_redirect_uri(): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + + with pytest.raises(HTTPException) as exc: + await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=_bridge_server(), + grant_type="authorization_code", + code="auth-code", + redirect_uri=None, + client_id="dcr-client-123", + client_secret=None, + code_verifier="verifier", + ) + + assert exc.value.status_code == 400 + assert "redirect_uri" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_token_non_bridge_keeps_gateway_callback(): + """Without dcr_bridge the token exchange keeps posting the gateway callback as redirect_uri, + pinning the pre-bridge contract for the browser-only Authorize path.""" + data = await _bridge_token_post_data(_bridge_server(dcr_bridge=None), redirect_uri=_BRIDGE_CLIENT_REDIRECT) + + assert data["redirect_uri"] == "https://litellm.example.com/callback" + + async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool: """Run exchange_token_with_server for a server of ``auth_type`` and report whether it attempted to persist the exchanged token server-side. The client-forwarded token modes must not persist: