fix(mcp): name the connected segment in sign-in challenge resource metadata

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-27 02:17:03 +00:00
parent 6a71a3b1a0
commit fcfd7b4f14
6 changed files with 159 additions and 5 deletions

View file

@ -531,7 +531,10 @@ def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCP
by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip)
if by_name is not None:
return by_name
return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
by_id: Final = global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
if by_id is not None:
return by_id
return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip)
def _resolve_oauth2_server_for_root_endpoints(

View file

@ -298,7 +298,7 @@ def raise_public(error: CredError) -> NoReturn:
assert_never(error.tag)
def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str:
def oauth_protected_resource_path(root_path: str, server: MCPServer, *, connected_as: str | None = None) -> str:
"""The server's RFC 9728 Protected Resource Metadata path, the shared anchor of both challenges.
``root_path`` is the prefix the request was routed under, resolved by the caller (the imperative
@ -320,7 +320,7 @@ def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str:
challenge would then disagree on where the resource metadata lives.
"""
prefix: Final = "" if root_path == "/" else root_path
name: Final = server.alias or server.server_name or server.name or server.server_id
name: Final = connected_as or server.alias or server.server_name or server.name or server.server_id
scalar_env: Final = os.getenv("SERVER_ROOT_PATH", "").rstrip("/")
if not prefix or (scalar_env and prefix == scalar_env):
return f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}"
@ -348,6 +348,7 @@ def raise_token_exchange_challenge(
*,
root_path: str,
claims: str | None = None,
connected_as: str | None = None,
) -> NoReturn:
"""Raise the RFC 9728 / RFC 6750 challenge an OBO (``token_exchange``) server returns when the
caller's subject token is missing or the IdP rejected it.
@ -366,7 +367,7 @@ def raise_token_exchange_challenge(
two literals) and the base64 claims draw from a fixed alphabet, so nothing from the IdP body
reaches the header unescaped.
"""
resource_metadata: Final = oauth_protected_resource_path(root_path, server)
resource_metadata: Final = oauth_protected_resource_path(root_path, server, connected_as=connected_as)
encoded_claims: Final = base64.b64encode(claims.encode()).decode() if claims else None
error: Final = "insufficient_claims" if encoded_claims else "invalid_token"
error_description: Final = (

View file

@ -1737,7 +1737,9 @@ if MCP_AVAILABLE:
get_request_root_path,
)
raise_token_exchange_challenge(server, root_path=get_request_root_path())
raise_token_exchange_challenge(
server, root_path=get_request_root_path(), connected_as=server_name
)
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728

View file

@ -216,3 +216,46 @@ def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gatew
assert refused.status_code == 403, refused.text
assert "www-authenticate" not in refused.headers
assert api.drain() == ()
def test_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, tmp_path: Path) -> None:
def nothing(request: Request) -> Reply:
return Reply(status=500)
with wire_server(nothing) as api:
config: Final = _sign_in_config(
{
"guardrail": "agent_365",
"mode": "pre_mcp_call",
"default_on": True,
"tenant_id": "00000000-0000-0000-0000-000000000000",
"client_id": "22222222-2222-2222-2222-222222222222",
"client_secret": "secret",
"api_base": api.url,
},
tmp_path / "agent365-case.yaml",
)
with (
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
mcp_peer() as peer,
candidate.scenario() as scenario,
):
alias: Final = "a365" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
connected_as: Final = alias.upper()
challenged: Final = _rpc(candidate, f"/mcp/{connected_as}", granted, {})
assert challenged.status_code == 401, challenged.text
authenticate: Final = challenged.headers.get("www-authenticate", "")
assert f'resource_metadata="/.well-known/oauth-protected-resource/mcp/{connected_as}"' in authenticate
assert 'error="invalid_token"' in authenticate
discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{connected_as}")
assert discovery.status_code == 200, discovery.text
document: Final = discovery.json()
assert document["authorization_servers"] == [
"https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0"
]
assert document["scopes_supported"] == ["api://22222222-2222-2222-2222-222222222222/access_as_user"]
assert api.drain() == ()

View file

@ -3743,6 +3743,56 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku
by_id.assert_called_once_with(server.server_id, client_ip=None)
@pytest.mark.asyncio
async def test_protected_resource_metadata_resolves_the_connected_case_variant():
"""The challenge points clients at the segment they connected with (``/mcp/CATALOG``), so the PRM
route must resolve that same segment through the answering-to fallback."""
from fastapi import Request
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="catalog-server-id-001",
name="catalog",
alias="catalog",
server_name="catalog",
url="https://catalog.test/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
mcp_info={"server_name": "catalog"},
)
sign_in: Final = CallerSignIn(
issuers=("https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0",),
scopes=("api://22222222-2222-2222-2222-222222222222/access_as_user",),
)
request = MagicMock(spec=Request)
request.base_url = "https://llm.example.com/"
request.headers = {}
with (
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None), # test-quality-ok: resolver seam
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=None), # test-quality-ok: resolver seam
patch.object(
global_mcp_server_manager, "get_filtered_registry", return_value={server.server_id: server}
), # test-quality-ok: resolver seam
patch.object(discoverable_endpoints, "caller_sign_in_for", return_value=sign_in), # test-quality-ok: provider seam
):
result = await discoverable_endpoints._build_oauth_protected_resource_response(
request=request,
mcp_server_name="CATALOG",
use_standard_pattern=True,
)
assert result["authorization_servers"] == [
"https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0"
]
assert result["resource"] == "https://llm.example.com/mcp/CATALOG"
def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails():
from fastapi import Request

View file

@ -11063,11 +11063,22 @@ def _catalog_server() -> MCPServer:
class _CallerSignInGuardrail(CustomGuardrail):
def __init__(self, *args, preflight_result=None, **kwargs):
super().__init__(*args, **kwargs)
self._preflight_result = preflight_result
self.preflight_calls = [] # mutable-ok: call recorder
def caller_sign_in(self, server, user_api_key_auth):
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",))
async def preflight_caller_sign_in(self, server, user_api_key_auth, subject_token):
from litellm.proxy._experimental.mcp_server.caller_sign_in import SignedIn
self.preflight_calls.append(subject_token)
return self._preflight_result if self._preflight_result is not None else SignedIn()
class TestConnectChallengeResolver:
"""The connect-time sign-in challenge must resolve the server the same way the router resolves
@ -11195,3 +11206,47 @@ class TestConnectChallengeResolver:
'error="invalid_token", '
'error_description="Missing or invalid subject token; authenticate with the IdP and retry"'
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route_name",
["catalog-server-id-001", "CATALOG"],
ids=["server_id", "uppercase_name"],
)
async def test_challenge_resource_metadata_names_the_connected_segment(self, route_name):
"""The PRM path in the challenge must name the segment the client connected with, or the
client's follow-up metadata fetch 404s against the route it was pointed at."""
from litellm.proxy._experimental.mcp_server import server as server_module
server = _catalog_server()
guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub")
litellm.logging_callback_manager.add_litellm_callback(guardrail)
try:
with (
patch.object(
mcp_operations.global_mcp_server_manager,
"get_filtered_registry",
return_value={server.server_id: server},
),
patch.object(
mcp_operations,
"_get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
),
pytest.raises(HTTPException) as exc,
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []},
mcp_servers=[route_name],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, guardrail, require_self=False
)
authenticate: Final = (exc.value.headers or {}).get("WWW-Authenticate") or ""
assert authenticate.startswith(f'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/{route_name}"')