Merge pull request #33174 from BerriAI/litellm_lit3637_aggregate_dcr

feat(mcp): always-on aggregate gateway DCR discovery front door
This commit is contained in:
tin-berri 2026-07-18 19:25:26 -07:00 • committed by GitHub
commit ac1b4b35ce
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 489 additions and 56 deletions

View file

@ -10,6 +10,10 @@ from typing_extensions import assert_never
import litellm
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_request_base_url,
well_known_root_suffix,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import (
BridgeEnvelopeAdmitted,
BridgeEnvelopeInvalid,
@ -120,6 +124,96 @@ def _has_client_supplied_mcp_auth(
return bool(mcp_auth_header) or bool(mcp_server_auth_headers)
def _is_aggregate_gateway_dcr_challenge_scope(
route: str,
mcp_servers: list[str] | None,
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
exc: Exception,
) -> bool:
"""True when an unauthenticated request to the aggregate ``/mcp`` endpoint
should receive the RFC 9728 401 challenge that advertises the gateway as
the authorization server.
Fires only for a genuine 401 on the aggregate scope: any named target
(path or ``x-mcp-servers``) belongs to the per-server challenge paths, and
client-supplied MCP auth headers mean the caller is not a cold-start DCR
client. Fails closed to the original admission error otherwise."""
if not _is_litellm_auth_admission_error(exc):
return False
if mcp_servers:
return False
if _has_client_supplied_mcp_auth(mcp_auth_header, mcp_server_auth_headers):
return False
return len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0
def _aggregate_gateway_dcr_challenge(request: Request, invalid_token: bool) -> HTTPException:
"""The RFC 9728 challenge for the aggregate endpoint: points the client at
the gateway's own protected-resource metadata so a DCR client discovers
the gateway as its authorization server and starts the sign-in flow.
``invalid_token`` adds the RFC 6750 error code for a request that DID
present a bearer that failed admission (expired or revoked), telling
spec-compliant clients to re-authorize rather than retry; a request with
no credentials at all gets the bare challenge per RFC 6750 section 3.1."""
error_attr = 'error="invalid_token", ' if invalid_token else ""
resource_metadata_url = (
f"{get_request_base_url(request)}/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp"
)
return HTTPException(
status_code=401,
detail={
"error": "authentication_required",
"message": "Authenticate with the gateway to use the MCP endpoint.",
},
headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"'},
)
def _admission_failure_fallback(
request: Request,
request_route: str,
mcp_servers: list[str] | None,
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
exc: Exception,
bearer_presented: bool,
) -> UserAPIKeyAuth:
"""Map a failed LiteLLM admission to its anonymous fallback or challenge.
Two fallbacks exist, both gated on a genuine 401 with no client-supplied
MCP auth headers. The pass-through cold start (RFC 9728 / MCP
Authorization spec discovery return) admits anonymously so the route's
401 emitter can produce the per-server challenge. The aggregate
gateway-DCR scope converts the failure into the gateway's own
resource_metadata challenge, with the RFC 6750 ``invalid_token`` error
code when the caller DID present a bearer (an expired gateway session
must re-authorize, not retry a dead token). Anything else re-raises the
original admission error unchanged."""
mcp_servers_from_path = _parse_mcp_server_names_from_path(request_route, mcp_servers)
if (
mcp_servers_from_path is not None
and not _has_client_supplied_mcp_auth(mcp_auth_header, mcp_server_auth_headers)
and _is_litellm_auth_admission_error(exc)
and _is_mcp_passthrough_cold_start(
mcp_servers_from_path,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
):
verbose_logger.debug("MCP pass-through cold start: deferring admission to route 401 emitter")
return UserAPIKeyAuth()
if _is_aggregate_gateway_dcr_challenge_scope(
route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=exc,
):
raise _aggregate_gateway_dcr_challenge(request, invalid_token=bearer_presented) from exc
raise exc
class MCPRequestHandler:
"""
Class to handle MCP request processing, including:
@ -271,56 +365,32 @@ class MCPRequestHandler:
elif oauth2_headers:
# Authorization on a non-delegated server: the bearer must be a real
# LiteLLM credential, so a failed validation is a genuine 401/403 and
# propagates. The sole anonymous fallback is the auth_type=none
# pass-through cold-start (RFC 9728 discovery return), gated on a 401
# so a recognized-but-forbidden key still fails closed.
client_ip = IPAddressUtils.get_mcp_client_ip(request)
# propagates unless a fallback in _admission_failure_fallback applies.
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as e:
# ProxyException.code is normalized to str (possibly "None"), so
# compare both int and str forms rather than coercing.
status = e.status_code if isinstance(e, HTTPException) else e.code
is_unauthenticated = status in (401, "401")
mcp_servers_from_path = _parse_mcp_server_names_from_path(request_route, mcp_servers)
if (
is_unauthenticated
and mcp_servers_from_path is not None
and not _has_client_supplied_mcp_auth(
mcp_auth_header,
mcp_server_auth_headers,
)
and _is_mcp_passthrough_cold_start(mcp_servers_from_path, client_ip=client_ip)
):
verbose_logger.debug(
"MCP pass-through return: forwarding Authorization as upstream OAuth token for delegated auth"
)
validated_user_api_key_auth = UserAPIKeyAuth()
else:
raise
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=e,
bearer_presented=True,
)
else:
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as exc:
# Cold-start MCP OAuth discovery: RFC 9728 / MCP Authorization spec
# require unauthenticated requests to protected resources to receive
# 401 + WWW-Authenticate. Defer to _raise_preemptive_401_for_unauthenticated_servers
# for pass-through servers instead of surfacing a generic admission error.
mcp_servers_from_path = _parse_mcp_server_names_from_path(request_route, mcp_servers)
client_ip = IPAddressUtils.get_mcp_client_ip(request)
if (
mcp_servers_from_path is not None
and not _has_client_supplied_mcp_auth(
mcp_auth_header,
mcp_server_auth_headers,
)
and _is_litellm_auth_admission_error(exc)
and _is_mcp_passthrough_cold_start(mcp_servers_from_path, client_ip=client_ip)
):
verbose_logger.debug("MCP pass-through cold start: deferring admission to route 401 emitter")
validated_user_api_key_auth = UserAPIKeyAuth()
else:
raise
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=exc,
bearer_presented=False,
)
return (
validated_user_api_key_auth,

View file

@ -43,6 +43,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
get_request_base_url,
validate_trusted_redirect_uri,
well_known_root_suffix,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
@ -50,7 +51,6 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.utils import get_server_root_path
from litellm.types.mcp import MCPAuth, MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -1838,11 +1838,88 @@ def _jwt_auth_issuers() -> list:
return issuers
def _build_aggregate_protected_resource_response(request: Request) -> dict:
"""RFC 9728 metadata for the aggregate /mcp resource: the gateway itself is
the authorization server. No per-server names or scopes leak here; access
is resolved after sign-in from the authenticated user's grants.
The advertised authorization server is ``{base}/mcp`` (not the bare
origin) so RFC 8414 path-insertion resolves its metadata at
``/.well-known/oauth-authorization-server/mcp``, a route this module
owns. The bare-origin well-known is registered first by the BYOK OAuth
feature and describes the BYOK flow, so it must not be the aggregate
discovery entry point (same pattern as the per-server documents, which
advertise ``{base}/{server_name}``)."""
request_base_url = get_request_base_url(request)
return {
"authorization_servers": [f"{request_base_url}/mcp"],
"resource": f"{request_base_url}/mcp",
"scopes_supported": [],
}
def _build_aggregate_authorization_server_response(request: Request) -> dict:
"""RFC 8414 metadata for the gateway as the aggregate authorization server.
The issuer is ``{base}/mcp`` and must stay equal to the value the
aggregate protected-resource document advertises: spec clients verify the
issuer in the metadata matches the one that derived the well-known URL.
Advertises the root /authorize, /token, and /register endpoints and
``token_endpoint_auth_methods_supported: ["none", ...]`` because DCR
clients (Claude Desktop, MCP Inspector) register as public clients; PKCE
S256 is mandatory in the gateway's authorize flow."""
request_base_url = get_request_base_url(request)
return {
"issuer": f"{request_base_url}/mcp",
"authorization_endpoint": f"{request_base_url}/authorize",
"token_endpoint": f"{request_base_url}/token",
"registration_endpoint": f"{request_base_url}/register",
"response_types_supported": ["code"],
"scopes_supported": [],
"grant_types_supported": ["authorization_code", "refresh_token"],
"code_challenge_methods_supported": ["S256"],
"token_endpoint_auth_methods_supported": ["none", "client_secret_post"],
}
# RFC 9728 path-appended discovery for the aggregate /mcp endpoint. A client
# pointed at {base}/mcp inserts the well-known segment before the resource
# path, so this exact route must exist for aggregate discovery to work at all.
# Declared before the parameterized well-known routes below: Starlette matches
# in registration order, and /.well-known/oauth-authorization-server/{name}
# would otherwise capture the "/mcp" suffix as a server name.
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp")
async def oauth_protected_resource_aggregate(request: Request):
"""
OAuth protected resource discovery for the aggregate /mcp endpoint.
The single-segment ``/mcp`` path does not collide with any per-server PRM pattern
(those are two-segment: ``/mcp/{server}`` or ``/{server}/mcp``), so this unambiguously
describes the aggregate resource.
"""
return _build_aggregate_protected_resource_response(request)
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/mcp")
async def oauth_authorization_server_aggregate(request: Request):
"""
OAuth authorization server discovery for the aggregate /mcp endpoint, the RFC 8414
path-inserted form for a client that treats {base}/mcp as its authorization base URL.
The single-segment /mcp is reserved for the aggregate so the discovery chain stays
consistent: the aggregate protected-resource document advertises {base}/mcp as its
authorization server, so the document served here must have issuer {base}/mcp. A server
literally named ``mcp`` therefore does not take this route; it keeps its standard
two-segment discovery at /.well-known/oauth-authorization-server/mcp/mcp. Letting the
per-server row win here instead would serve an issuer of {base} against a resource that
advertised {base}/mcp, which fails the RFC 8414 issuer check and breaks the front door.
"""
return _build_aggregate_authorization_server_response(request)
# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name}
# This is the pattern expected by standard MCP clients (mcp-inspector, VSCode Copilot)
@router.get(
f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/mcp/{{mcp_server_name}}"
)
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth protected resource discovery endpoint using standard MCP URL pattern.
@ -1862,9 +1939,7 @@ async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_nam
# LiteLLM legacy pattern: /.well-known/oauth-protected-resource/{server_name}/mcp
# Kept for backward compatibility with existing deployments
@router.get(
f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp"
)
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
@router.get("/.well-known/oauth-protected-resource")
async def oauth_protected_resource_mcp(request: Request, mcp_server_name: Optional[str] = None):
"""
@ -1934,9 +2009,7 @@ def _build_oauth_authorization_server_response(
# Standard MCP pattern: /.well-known/oauth-authorization-server/mcp/{server_name}
@router.get(
f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/mcp/{{mcp_server_name}}"
)
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery endpoint using standard MCP URL pattern.
@ -1951,9 +2024,7 @@ async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_n
# LiteLLM legacy pattern and root endpoint
@router.get(
f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}"
)
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}")
@router.get("/.well-known/oauth-authorization-server")
async def oauth_authorization_server_mcp(request: Request, mcp_server_name: Optional[str] = None):
"""

View file

@ -132,6 +132,18 @@ def get_request_base_url(request: Request) -> str:
return urlunparse((scheme, _strip_default_port(scheme, netloc), parsed.path, "", "", ""))
def well_known_root_suffix() -> str:
"""The ``SERVER_ROOT_PATH`` segment inserted into a ``.well-known`` path (RFC 8414 / 9728
path insertion), empty for a root-mounted proxy or an explicit ``/``.
The discovery route registrations and the 401 challenges that advertise those routes both
derive their path from this one function, so the ``resource_metadata`` URL a client is told
to fetch cannot drift from the route that actually serves it.
"""
root = os.getenv("SERVER_ROOT_PATH", "")
return "" if root == "/" else root
def validate_loopback_redirect_uri(redirect_uri: str) -> None:
"""Require a loopback ``redirect_uri`` (OAuth 2.1 §4.1.2.1 + RFC 8252
§7.3 native-app pattern). MCP clients are native apps that listen on

View file

@ -6131,3 +6131,133 @@ class TestMCPDcrBridgeDelegateAdmission:
route="/mcp/bridge_delegate_server",
)
assert exc_info.value.status_code == 500
@pytest.mark.asyncio
class TestAggregateGatewayDcrChallenge:
"""The mcp_gateway_dcr front door: a 401 on the aggregate /mcp scope must
carry the RFC 9728 resource_metadata challenge pointing at the gateway's
own protected-resource metadata, and must NOT fire for named-server
targets, explicit litellm keys, or non-401 failures."""
_AUTH_PATCH_TARGET = "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth"
_EXPECTED_RESOURCE_METADATA = 'resource_metadata="http://testserver/.well-known/oauth-protected-resource/mcp"'
def _scope(self, path="/mcp", extra_headers=()):
return {
"type": "http",
"method": "POST",
"path": path,
"headers": [(b"host", b"testserver"), *extra_headers],
}
def _auth_401(self):
async def _raise(api_key, request):
raise ProxyException(
message="Authentication Error: Invalid API key",
type="auth_error",
param="api_key",
code=401,
)
return _raise
async def test_challenge_on_anonymous_aggregate_mcp(self):
"""Anonymous request to the aggregate /mcp: 401 plus
the bare bearer challenge (no error attribute, RFC 6750 section 3.1)."""
with (
patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()),
):
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(self._scope())
assert exc_info.value.status_code == 401
www_authenticate = (exc_info.value.headers or {})["WWW-Authenticate"]
assert www_authenticate == f"Bearer {self._EXPECTED_RESOURCE_METADATA}"
async def test_challenge_invalid_token_on_failed_bearer(self):
"""A bearer that fails LiteLLM admission at aggregate scope (an expired
gateway session, a revoked key) re-challenges with error=invalid_token
so a spec client re-authorizes instead of retrying the dead token."""
with (
patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()),
):
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(
self._scope(extra_headers=((b"authorization", b"Bearer expired-session-token"),))
)
assert exc_info.value.status_code == 401
www_authenticate = (exc_info.value.headers or {})["WWW-Authenticate"]
assert www_authenticate == f'Bearer error="invalid_token", {self._EXPECTED_RESOURCE_METADATA}'
async def test_challenge_inserts_server_root_path(self):
"""With SERVER_ROOT_PATH set the resource_metadata URL must carry the same path-inserted
root segment the aggregate PRM route is registered with (both derive it from
well_known_root_suffix), so a DCR client behind a sub-path is pointed at a route that
exists instead of a 404. Regression: the challenge used to hard-code /mcp and omit the
root path the route inserts."""
import os
with (
patch.dict(os.environ, {"SERVER_ROOT_PATH": "/litellm"}),
patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()),
):
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(self._scope())
www_authenticate = (exc_info.value.headers or {})["WWW-Authenticate"]
assert 'resource_metadata="http://testserver/.well-known/oauth-protected-resource/litellm/mcp"' in www_authenticate
async def test_no_challenge_for_explicit_litellm_key(self):
"""An explicit x-litellm-api-key declares a litellm-key client; a typo
there must surface the real auth error, never a DCR challenge that
would send SDKs into a sign-in flow."""
with (
patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()),
):
with pytest.raises(ProxyException):
await MCPRequestHandler.process_mcp_request(
self._scope(extra_headers=((b"x-litellm-api-key", b"sk-typo"),))
)
async def test_no_challenge_for_named_servers_header(self):
"""x-mcp-servers names explicit targets; the per-server challenge paths
own those, so the aggregate challenge must not fire."""
with (
patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()),
):
with pytest.raises(ProxyException):
await MCPRequestHandler.process_mcp_request(
self._scope(extra_headers=((b"x-mcp-servers", b"github"),))
)
async def test_no_challenge_for_path_named_server(self):
"""/mcp/{server} targets one server; the aggregate challenge must not
fire even when that server does not resolve."""
with (
patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()),
):
with pytest.raises(ProxyException):
await MCPRequestHandler.process_mcp_request(self._scope(path="/mcp/github"))
async def test_no_challenge_for_client_supplied_mcp_auth(self):
"""Per-server x-mcp-{alias}-authorization headers mean the caller is
not a cold-start DCR client; keep the original error."""
with (
patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()),
):
with pytest.raises(ProxyException):
await MCPRequestHandler.process_mcp_request(
self._scope(extra_headers=((b"x-mcp-github-authorization", b"Bearer upstream"),))
)
async def test_no_challenge_for_non_401_failure(self):
"""Only genuine 401s convert to a challenge; a 500 stays a 500."""
async def _raise_500(api_key, request):
raise ProxyException(message="boom", type="server_error", param=None, code=500)
with (
patch(self._AUTH_PATCH_TARGET, side_effect=_raise_500),
):
with pytest.raises(ProxyException) as exc_info:
await MCPRequestHandler.process_mcp_request(self._scope())
assert str(exc_info.value.code) == "500"

View file

@ -0,0 +1,22 @@
import os
import pytest
@pytest.fixture(autouse=True)
def _hermetic_server_root_path():
"""Isolate MCP discovery tests from a leaked ``SERVER_ROOT_PATH``.
``tests/test_litellm/proxy/test_custom_proxy.py`` sets ``SERVER_ROOT_PATH`` at import time
(its app mounts under a custom path) and never restores it, so in a shared shard the value
leaks into this process. The discovery routes and the 401 challenges read it, so a leaked
value would silently rewrite every ``resource_metadata`` URL and make these tests depend on
shard ordering. Clearing it here pins the default (root-mounted) deployment; a test that
exercises a sub-path deployment sets the value explicitly within its own body.
"""
saved = os.environ.pop("SERVER_ROOT_PATH", None)
try:
yield
finally:
if saved is not None:
os.environ["SERVER_ROOT_PATH"] = saved

View file

@ -7500,3 +7500,131 @@ async def test_reload_servers_from_database_hydrates_dcr_clients():
await global_mcp_server_manager.reload_servers_from_database()
hydrate_spy.assert_awaited_once()
def test_aggregate_wellknown_routes_serve_gateway_metadata():
"""Both path-appended aggregate routes serve the gateway documents. Exercises real
routing, so this also pins registration order: the parameterized
/.well-known/oauth-authorization-server/{name} route would otherwise capture the /mcp
suffix as a server name."""
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.registry.clear()
app = FastAPI()
app.include_router(router)
client = TestClient(app)
prm = client.get("/.well-known/oauth-protected-resource/mcp")
asm = client.get("/.well-known/oauth-authorization-server/mcp")
assert prm.status_code == 200
assert prm.json()["resource"] == "http://testserver/mcp"
assert prm.json()["authorization_servers"] == ["http://testserver/mcp"]
assert asm.status_code == 200
assert asm.json()["issuer"] == "http://testserver/mcp"
assert asm.json()["authorization_endpoint"] == "http://testserver/authorize"
assert "none" in asm.json()["token_endpoint_auth_methods_supported"]
def test_as_aggregate_route_reserves_mcp_for_the_aggregate():
"""The single-segment /.well-known/oauth-authorization-server/mcp is reserved for the
aggregate even when a server is literally named ``mcp``. The aggregate protected-resource
document advertises {base}/mcp as its authorization server, so the document served here
must carry issuer {base}/mcp for the RFC 8414 issuer check to pass. Letting the per-server
row win (issuer {base}) breaks that chain, so the aggregate wins and the mcp-named server
keeps its standard two-segment discovery at /.well-known/oauth-authorization-server/mcp/mcp."""
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.registry.clear()
server_named_mcp = _create_oauth2_server(server_id="mcp_srv", name="mcp", server_name="mcp", alias="mcp")
global_mcp_server_manager.registry[server_named_mcp.server_id] = server_named_mcp
app = FastAPI()
app.include_router(router)
client = TestClient(app)
try:
asm = client.get("/.well-known/oauth-authorization-server/mcp")
assert asm.status_code == 200
# the aggregate document, whose issuer matches what the aggregate PRM advertises
assert asm.json()["issuer"] == "http://testserver/mcp"
prm = client.get("/.well-known/oauth-protected-resource/mcp")
assert prm.status_code == 200
assert prm.json()["authorization_servers"] == [asm.json()["issuer"]]
# the mcp-named server keeps its own document on the standard two-segment route
per_server = client.get("/.well-known/oauth-authorization-server/mcp/mcp")
assert per_server.status_code == 200
assert "/mcp/authorize" in per_server.json()["authorization_endpoint"]
finally:
global_mcp_server_manager.registry.clear()
def test_well_known_root_suffix_reflects_server_root_path():
"""The single path segment both the discovery routes and the 401 challenges insert for RFC
8414/9728 path insertion: empty for a root-mounted proxy or an explicit ``/``, the configured
path otherwise. Sharing this one function is what keeps the advertised resource_metadata URL
equal to the route that serves it."""
import os
from unittest.mock import patch
from litellm.proxy._experimental.mcp_server.oauth_utils import well_known_root_suffix
with patch.dict(os.environ, {"SERVER_ROOT_PATH": ""}):
assert well_known_root_suffix() == ""
with patch.dict(os.environ, {"SERVER_ROOT_PATH": "/"}):
assert well_known_root_suffix() == ""
with patch.dict(os.environ, {"SERVER_ROOT_PATH": "/litellm"}):
assert well_known_root_suffix() == "/litellm"
@pytest.mark.asyncio
async def test_bare_origin_discovery_resolves_single_server_not_aggregate():
"""The always-on aggregate front door must not change bare-origin discovery: with one
oauth2 server configured, the no-suffix /.well-known/oauth-{authorization-server,
protected-resource} still resolves THAT server, so an existing single-server deployment's
discovery is unchanged. The aggregate document lives only at the /mcp-suffixed routes."""
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_build_oauth_authorization_server_response,
_build_oauth_protected_resource_response,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server()
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
authorization_response = _build_oauth_authorization_server_response(
request=mock_request, mcp_server_name=None
)
resource_response = await _build_oauth_protected_resource_response(
request=mock_request, mcp_server_name=None, use_standard_pattern=True
)
# per-server, not aggregate: the single server's name is in the endpoints
assert "/test_oauth/authorize" in authorization_response["authorization_endpoint"]
assert authorization_response["issuer"] == "https://llm.example.com"
assert resource_response["authorization_servers"] == ["https://llm.example.com/test_oauth"]
finally:
global_mcp_server_manager.registry.clear()