fix(mcp): match per-server OAuth metadata issuers

This commit is contained in:
Joshua Valluru 2026-09-11 17:05:35 -07:00
parent e4706fa409
commit 17863fa5cf
2 changed files with 71 additions and 9 deletions

View file

@ -2490,7 +2490,11 @@ async def _build_oauth_protected_resource_response(
return {
"authorization_servers": [
(f"{request_base_url}/{mcp_server_name}" if mcp_server_name else f"{request_base_url}")
resource_url
if explicitly_named
else f"{request_base_url}/{mcp_server_name}"
if mcp_server_name
else request_base_url
],
"resource": resource_url,
"scopes_supported": (mcp_server.scopes if mcp_server and mcp_server.scopes else []),
@ -2666,6 +2670,8 @@ async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str |
def _build_oauth_authorization_server_response(
request: Request,
mcp_server_name: str | None,
*,
issuer_path: str | None = None,
) -> dict:
"""Build OAuth authorization server metadata response (gateway-as-AS shape).
@ -2694,7 +2700,13 @@ def _build_oauth_authorization_server_response(
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth authorization server")
issuer: Final = f"{request_base_url}/{mcp_server_name}" if explicitly_named else request_base_url
issuer: Final = (
f"{request_base_url}/{issuer_path}"
if issuer_path is not None
else f"{request_base_url}/{mcp_server_name}"
if explicitly_named
else request_base_url
)
return {
"issuer": issuer,
@ -2724,6 +2736,7 @@ async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_n
return _build_oauth_authorization_server_response(
request=request,
mcp_server_name=mcp_server_name,
issuer_path=f"mcp/{mcp_server_name}",
)
@ -2802,7 +2815,7 @@ async def jwks_json(request: Request):
# Additional legacy pattern support
@router.get("/.well-known/oauth-authorization-server/{mcp_server_name}/mcp")
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
@ -2810,6 +2823,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
return _build_oauth_authorization_server_response(
request=request,
mcp_server_name=mcp_server_name,
issuer_path=f"{mcp_server_name}/mcp",
)

View file

@ -3419,16 +3419,16 @@ async def test_oauth_protected_resource_gateway_managed_oauth2_advertises_gatewa
relay_response = await _build_oauth_protected_resource_response(
request=mock_request, mcp_server_name="relay_mcp", use_standard_pattern=True
)
assert relay_response["authorization_servers"] == ["https://litellm.example.com/relay_mcp"]
assert relay_response["authorization_servers"] == ["https://litellm.example.com/mcp/relay_mcp"]
relay_legacy_response = await _build_oauth_protected_resource_response(
request=mock_request, mcp_server_name="relay_mcp", use_standard_pattern=False
)
assert relay_legacy_response["authorization_servers"] == ["https://litellm.example.com/relay_mcp"]
assert relay_legacy_response["authorization_servers"] == ["https://litellm.example.com/relay_mcp/mcp"]
delegated_response = await _build_oauth_protected_resource_response(
request=mock_request, mcp_server_name="delegated_mcp", use_standard_pattern=True
)
assert delegated_response["authorization_servers"] == ["https://litellm.example.com/delegated_mcp"]
assert delegated_response["authorization_servers"] == ["https://litellm.example.com/mcp/delegated_mcp"]
finally:
global_mcp_server_manager.registry.clear()
@ -9120,8 +9120,8 @@ async def test_named_discovery_issuer_matches_protected_resource_authorization_s
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_build_oauth_authorization_server_response,
_build_oauth_protected_resource_response,
oauth_authorization_server_mcp_standard,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
@ -9139,10 +9139,10 @@ async def test_named_discovery_issuer_matches_protected_resource_authorization_s
resource_response = await _build_oauth_protected_resource_response(
request=mock_request, mcp_server_name="test_oauth", use_standard_pattern=True
)
authorization_response = _build_oauth_authorization_server_response(
authorization_response = await oauth_authorization_server_mcp_standard(
request=mock_request, mcp_server_name="test_oauth"
)
assert resource_response["authorization_servers"] == ["https://llm.example.com/test_oauth"]
assert resource_response["authorization_servers"] == ["https://llm.example.com/mcp/test_oauth"]
assert authorization_response["issuer"] == resource_response["authorization_servers"][0]
finally:
global_mcp_server_manager.registry.clear()
@ -11186,3 +11186,51 @@ async def test_dcr_refusal_is_actionable_without_upstream_body(
assert f"HTTP {upstream_status}" in str(exc.value.detail)
assert "pre-registered OAuth client" in str(exc.value.detail)
assert "private upstream details" not in str(exc.value.detail)
@pytest.mark.parametrize("prefix", ["", "/tenant-a", "/tenant-b"])
@pytest.mark.parametrize(
("server_name", "pattern"),
[
("issuer_test", "mcp/{server}"),
("issuer_test", "{server}/mcp"),
("issuer_test", "{server}"),
("mcp", "mcp/{server}"),
("mcp", "{server}/mcp"),
],
)
def test_per_server_authorization_metadata_issuer_matches_discovery_path(
_no_proxy_base_url, _isolated_mcp_registry, prefix, server_name, pattern
):
server = _create_oauth2_server(server_id=server_name, name=server_name, server_name=server_name, alias=server_name)
_isolated_mcp_registry[server.server_id] = server
client = _prefixed_discovery_client(["/tenant-a", "/tenant-b"])
path = pattern.format(server=server_name)
response = client.get(f"{prefix}/.well-known/oauth-authorization-server/{path}")
assert response.status_code == 200
metadata = response.json()
assert metadata["issuer"] == f"http://testserver{prefix}/{path}"
assert metadata["authorization_endpoint"] == f"http://testserver{prefix}/{server_name}/authorize"
assert metadata["token_endpoint"] == f"http://testserver{prefix}/{server_name}/token"
assert metadata["registration_endpoint"] == f"http://testserver{prefix}/{server_name}/register"
@pytest.mark.parametrize("prefix", ["", "/tenant-a"])
@pytest.mark.parametrize("relay", [False, True])
@pytest.mark.parametrize("pattern", ["mcp/{server}", "{server}/mcp"])
def test_named_resource_discovery_follows_matching_authorization_issuer(
_no_proxy_base_url, _isolated_mcp_registry, prefix, relay, pattern
):
server = _create_oauth2_server().model_copy(update={"per_server_oauth_discovery": relay})
_isolated_mcp_registry[server.server_id] = server
client = _prefixed_discovery_client(["/tenant-a"])
path = pattern.format(server=server.server_name)
response = client.get(f"{prefix}/.well-known/oauth-protected-resource/{path}")
assert response.status_code == 200
resource = response.json()
issuer_path = path if relay else "mcp"
assert resource["resource"] == f"http://testserver{prefix}/{path}"
assert resource["authorization_servers"] == [f"http://testserver{prefix}/{issuer_path}"]
authorization = client.get(f"{prefix}/.well-known/oauth-authorization-server/{issuer_path}")
assert authorization.status_code == 200
assert authorization.json()["issuer"] == resource["authorization_servers"][0]