diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index b9aa3aded80..12e7fd7fe8b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -499,17 +499,21 @@ def _raise_unless_oauth2_discovery_server( mcp_server_name: Optional[str], description: str, ) -> None: - """404 a NAMED discovery request unless it resolves to an oauth2 server. + """404 a NAMED discovery request unless it resolves to an oauth2 or DCR-bridge server. A named server that is unknown (or hidden from the caller) and one that exists but is non-oauth2 both return the same 404, so the well-known discovery paths cannot be used to enumerate non-OAuth server names. Root discovery (no name) is unaffected, and pass-through servers are resolved by the caller before this runs. + DCR-bridge servers are admitted because they serve the gateway's own authorization + server metadata (the register, authorize, and token relays). """ if mcp_server_name is None: return if mcp_server is not None and mcp_server.auth_type == MCPAuth.oauth2: return + if mcp_server is not None and mcp_server.is_dcr_bridge: + return raise HTTPException( status_code=404, detail=f"MCP server '{mcp_server_name}' is {description}", @@ -735,7 +739,17 @@ async def exchange_token_with_server( detail="MCP upstream token endpoint returned no response", ) - response.raise_for_status() + try: + response.raise_for_status() + except httpx.HTTPStatusError as exc: + if "invalid_target" in exc.response.text: + verbose_logger.warning( + "MCP server %s: the upstream authorization server rejected the token request with " + "invalid_target; it may require RFC 8707 resource indicators, which the gateway " + "does not send yet (tracked as LIT-4339)", + mcp_server.server_id, + ) + raise token_response = response.json() access_token = token_response["access_token"] @@ -984,6 +998,7 @@ async def register_client_with_server( token_endpoint_auth_method: Optional[str], fallback_client_id: Optional[str] = None, persist_credentials: bool = False, + client_redirect_uris: Optional[list] = None, ): _raise_if_not_oauth2(mcp_server) request_base_url = get_request_base_url(request) @@ -1005,12 +1020,19 @@ async def register_client_with_server( if mcp_server.registration_url is None: return dummy_return + bridge_relay = _dcr_bridge_relays_client_registration(mcp_server) + if bridge_relay and not client_redirect_uris: + raise HTTPException( + status_code=400, + detail={"error": "redirect_uris is required to register a client with this server"}, + ) + register_data = { "client_name": client_name, - "redirect_uris": [f"{request_base_url}/callback"], - "grant_types": grant_types or [], - "response_types": response_types or [], - "token_endpoint_auth_method": token_endpoint_auth_method or "", + "redirect_uris": client_redirect_uris if bridge_relay else [f"{request_base_url}/callback"], + "grant_types": grant_types or (["authorization_code", "refresh_token"] if bridge_relay else []), + "response_types": response_types or (["code"] if bridge_relay else []), + "token_endpoint_auth_method": token_endpoint_auth_method or ("none" if bridge_relay else ""), } headers = { "Content-Type": "application/json", @@ -1032,7 +1054,7 @@ async def register_client_with_server( token_response = response.json() - if persist_credentials: + if persist_credentials and not bridge_relay: persistence_result = await _persist_dcr_client_registration(mcp_server, token_response) if persistence_result == "reused": return dummy_return @@ -1456,6 +1478,13 @@ async def _build_oauth_protected_resource_response( else: resource_url = f"{request_base_url}/mcp" + if mcp_server is not None and mcp_server_name and mcp_server.is_dcr_bridge: + return { + "authorization_servers": [f"{request_base_url}/{mcp_server_name}"], + "resource": resource_url, + "scopes_supported": (mcp_server.scopes if mcp_server.scopes else []), + } + # Pass-through branch: proxy the upstream's own metadata so discovery # directs the client at the real IdP (Okta, Keycloak, …) instead of us. if mcp_server is not None and ( @@ -1799,4 +1828,5 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non response_types=data.get("response_types", []), token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""), fallback_client_id=mcp_server_name, + client_redirect_uris=data.get("redirect_uris"), ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a3da237ca09..1d681b43b9e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2630,10 +2630,16 @@ class MCPServerManager: return prefixed_or_original_tools - except MCPUpstreamAuthError: + except MCPUpstreamAuthError as upstream_auth_error: # Pass-through 401 must surface to single-server routes so the # client triggers the upstream OAuth flow. The multi-server # aggregator catches this explicitly to keep absorbing. + if server.is_dcr_bridge and upstream_auth_error.www_authenticate is not None: + raise MCPUpstreamAuthError( + status_code=upstream_auth_error.status_code, + www_authenticate=None, + server_name=upstream_auth_error.server_name, + ) from upstream_auth_error raise except HTTPException as e: # A v2 resolver auth challenge (token_exchange's RFC 9728 401, authorization_code's @@ -2643,9 +2649,10 @@ class MCPServerManager: # Non-auth HTTP errors stay absorbed so one misconfigured server can't blank the listing. if e.status_code in (401, 403): headers = e.headers or {} + challenge_header = headers.get("WWW-Authenticate") or headers.get("www-authenticate") raise MCPUpstreamAuthError( status_code=e.status_code, - www_authenticate=headers.get("WWW-Authenticate") or headers.get("www-authenticate"), + www_authenticate=None if server.is_dcr_bridge else challenge_header, server_name=server.name, ) from e verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}") diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 238cb51bb64..5090aa7d7d5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3700,6 +3700,17 @@ if MCP_AVAILABLE: and not _scope_has_authorization_header(scope) and not _client_has_per_server_auth_header(server, mcp_server_auth_headers) ): + if server.is_dcr_bridge: + raise HTTPException( + status_code=401, + detail="Unauthorized", + headers={ + "www-authenticate": _get_passthrough_www_authenticate( + scope=scope, + server_name=server_name, + ) + }, + ) upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "") if upstream_status == 401 and upstream_www_authenticate: raise HTTPException( 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 5b43f39a7a4..5871925dc93 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 @@ -3818,6 +3818,135 @@ async def test_token_non_bridge_keeps_gateway_callback(): assert data["redirect_uri"] == "https://litellm.example.com/callback" +def _named_as_metadata_response(server): + 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, + ) + + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry[server.server_id] = server + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip", + return_value=None, + ): + return _build_oauth_authorization_server_response( + request=_bridge_mock_request(), + mcp_server_name=server.server_name, + ) + finally: + global_mcp_server_manager.registry.clear() + + +def test_oauth_authorization_server_metadata_served_for_bridge_server(): + """Bridge servers get the gateway's AS metadata (the register, authorize, and token relays), + which is what makes the DCR front door discoverable to standard MCP clients.""" + result = _named_as_metadata_response(_bridge_server()) + + assert result["authorization_endpoint"] == "https://litellm.example.com/bridge_srv/authorize" + assert result["token_endpoint"] == "https://litellm.example.com/bridge_srv/token" + assert result["registration_endpoint"] == "https://litellm.example.com/bridge_srv/register" + + +def test_oauth_authorization_server_404_for_non_bridge_client_forwarded_server(): + """Without dcr_bridge a client-forwarded server keeps 404ing AS-metadata discovery: verbatim + upstream discovery is the contract and the gateway must not advertise itself as its AS.""" + with pytest.raises(HTTPException) as exc: + _named_as_metadata_response(_bridge_server(dcr_bridge=None)) + + assert exc.value.status_code == 404 + + +async def _bridge_register_response(server, request_payload, persist_credentials=False): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client_with_server, + ) + + mock_response = MagicMock() + mock_response.json.return_value = { + "client_id": "upstream-issued-client", + "redirect_uris": request_payload.get("redirect_uris", []), + "token_endpoint_auth_method": "none", + } + mock_response.raise_for_status = MagicMock() + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._persist_dcr_client_registration", + new_callable=AsyncMock, + ) as mock_persist, + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._reuse_persisted_dcr_client_if_available", + new_callable=AsyncMock, + return_value=False, + ), + ): + response = await register_client_with_server( + request=_bridge_mock_request(), + mcp_server=server, + client_name=request_payload.get("client_name", ""), + grant_types=request_payload.get("grant_types"), + response_types=request_payload.get("response_types"), + token_endpoint_auth_method=request_payload.get("token_endpoint_auth_method"), + persist_credentials=persist_credentials, + client_redirect_uris=request_payload.get("redirect_uris"), + ) + return response, mock_async_client, mock_persist + + +@pytest.mark.asyncio +async def test_register_bridge_relay_forwards_client_redirect_uris(): + """The bridge relay arm registers the client's own redirect_uris upstream with public-client + defaults and relays the upstream response verbatim, so the upstream AS enforces the redirect + binding for that client and the auth code never transits the gateway.""" + import json + + response, mock_async_client, _ = await _bridge_register_response( + _bridge_server(), + {"client_name": "Claude", "redirect_uris": [_BRIDGE_CLIENT_REDIRECT]}, + ) + + posted = mock_async_client.post.call_args.kwargs["json"] + assert posted["redirect_uris"] == [_BRIDGE_CLIENT_REDIRECT] + assert posted["grant_types"] == ["authorization_code", "refresh_token"] + assert posted["response_types"] == ["code"] + assert posted["token_endpoint_auth_method"] == "none" + + payload = json.loads(response.body.decode("utf-8")) + assert payload["client_id"] == "upstream-issued-client" + + +@pytest.mark.asyncio +async def test_register_bridge_relay_requires_redirect_uris(): + with pytest.raises(HTTPException) as exc: + await _bridge_register_response(_bridge_server(), {"client_name": "Claude"}) + + assert exc.value.status_code == 400 + assert "redirect_uris" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_register_bridge_relay_never_persists(): + """Relayed registrations belong to individual clients; persisting one as the server's own DCR + client would hand every future caller the first client's identity.""" + _, _, mock_persist = await _bridge_register_response( + _bridge_server(), + {"client_name": "Claude", "redirect_uris": [_BRIDGE_CLIENT_REDIRECT]}, + persist_credentials=True, + ) + + mock_persist.assert_not_called() + + 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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py index 6de35ebc524..ec285f8eba0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py @@ -571,3 +571,53 @@ async def test_oauth_protected_resource_true_passthrough_returns_upstream_metada assert result["resource"] == "https://upstream.example.com/mcp" finally: global_mcp_server_manager.registry.clear() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate]) +@pytest.mark.parametrize("use_standard_pattern", [True, False]) +async def test_oauth_protected_resource_dcr_bridge_returns_gateway_facade(auth_type, use_standard_pattern): + """With dcr_bridge on, discovery flips from the upstream-verbatim contract to the gateway + facade: resource is the gateway URL the client dialed and authorization_servers names the + gateway's per-server AS, so DCR-only clients (which enforce the RFC 9728 resource match) + can register and sign in through the gateway. No upstream metadata fetch happens.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + bridge_server = MCPServer( + server_id="bridge-1", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + dcr_bridge=True, + scopes=["read"], + registration_url="https://okta.example.com/register", + ) + global_mcp_server_manager.registry[bridge_server.server_id] = bridge_server + + try: + with patch.object(discoverable_endpoints, "get_async_httpx_client") as mock_client_factory: + result = await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=use_standard_pattern, + ) + finally: + global_mcp_server_manager.registry.clear() + + expected_resource = ( + "https://gateway.example.com/mcp/sample_docs" + if use_standard_pattern + else "https://gateway.example.com/sample_docs/mcp" + ) + assert result == { + "authorization_servers": ["https://gateway.example.com/sample_docs"], + "resource": expected_resource, + "scopes_supported": ["read"], + } + mock_client_factory.assert_not_called() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f81e460a8be..7c55bd4560f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6678,6 +6678,67 @@ class TestMCPToolsListAuthSurfacing: assert await manager._get_tools_from_server(server) == [] + @pytest.mark.asyncio + async def test_get_tools_from_server_suppresses_upstream_challenge_for_dcr_bridge(self): + """A dcr_bridge server must never relay the upstream's own WWW-Authenticate: it points + clients at the upstream protected-resource metadata, which fails the RFC 9728 resource + match against the gateway URL they dialed. Stripping it makes the single-server route + fabricate the gateway well-known challenge, whose content is the bridge facade.""" + from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPUpstreamAuthError, + ) + from litellm.types.mcp import MCPAuth + + manager = MCPServerManager() + bridge_server = MCPServer( + server_id="bridge-srv", + name="bridge-srv", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + dcr_bridge=True, + ) + upstream_challenge = 'Bearer resource_metadata="https://upstream.example/.well-known/oauth-protected-resource"' + client = MagicMock() + client.list_tools = AsyncMock(side_effect=_upstream_status_error(401, upstream_challenge)) + manager._create_mcp_client = AsyncMock(return_value=client) + + with pytest.raises(MCPUpstreamAuthError) as exc_info: + await manager._get_tools_from_server(bridge_server) + + assert exc_info.value.status_code == 401 + assert exc_info.value.www_authenticate is None + assert exc_info.value.server_name == "bridge-srv" + + @pytest.mark.asyncio + async def test_get_tools_from_server_suppresses_resolver_challenge_for_dcr_bridge(self): + """The client-build-time HTTPException conversion path applies the same suppression.""" + from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPUpstreamAuthError, + ) + from litellm.types.mcp import MCPAuth + + manager = MCPServerManager() + bridge_server = MCPServer( + server_id="bridge-srv", + name="bridge-srv", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth_delegate, + dcr_bridge=True, + ) + manager._create_mcp_client = AsyncMock( + side_effect=HTTPException( + status_code=401, + detail="Unauthorized", + headers={"WWW-Authenticate": 'Bearer resource_metadata="https://upstream.example/prm"'}, + ) + ) + + with pytest.raises(MCPUpstreamAuthError) as exc_info: + await manager._get_tools_from_server(bridge_server) + + assert exc_info.value.status_code == 401 + assert exc_info.value.www_authenticate is None + @pytest.mark.asyncio async def test_aggregate_list_tools_absorbs_unauthenticated_server(self): from litellm.proxy._experimental.mcp_server.exceptions import ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index da183e8d02a..be95b3f3f73 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -1400,6 +1400,76 @@ async def test_handle_streamable_http_mcp_true_passthrough_without_token_surface probe_client.post.assert_awaited_once() +@pytest.mark.asyncio +async def test_handle_streamable_http_mcp_true_passthrough_dcr_bridge_challenges_with_gateway_metadata(): + """With dcr_bridge on, the missing-token challenge names the GATEWAY's well-known instead of + relaying the upstream's: the gateway is the authorization server for bridge clients, and the + upstream's own challenge would point them at metadata that fails the RFC 9728 resource match. + The upstream probe is skipped entirely; the gateway can answer authoritatively.""" + from fastapi import HTTPException + + try: + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + probe_client = MagicMock() + probe_client.post = AsyncMock() + + scope = _passthrough_mode_scope("tp_bridge_server") + receive = AsyncMock( + return_value={ + "type": "http.request", + "body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}', + "more_body": False, + } + ) + send = AsyncMock() + user_auth = MagicMock() + user_auth.user_id = None + bridge_server = _build_passthrough_mode_server("tp_bridge_server", MCPAuth.true_passthrough).model_copy( + update={"dcr_bridge": True} + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, ["tp_bridge_server"], None, None, None), + ), + patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.get_async_httpx_client", + return_value=probe_client, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=bridge_server, + ), + patch.object( + session_manager_stateful, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): + with pytest.raises(HTTPException) as exc_info: + await handle_streamable_http_mcp(scope, receive, send) + + assert mock_handle_request.await_count == 0 + assert exc_info.value.status_code == 401 + challenge = exc_info.value.headers["www-authenticate"] + assert "/.well-known/oauth-protected-resource/tp_bridge_server/mcp" in challenge + assert "upstream.example.com" not in challenge + probe_client.post.assert_not_awaited() + + @pytest.mark.asyncio async def test_handle_streamable_http_mcp_true_passthrough_with_token_skips_probe_and_challenge(): """When the true_passthrough caller already carries an Authorization the