From fcfd7b4f1409ae70bc0a4b8d1e76b79cbddfd2cc Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 02:17:03 +0000 Subject: [PATCH] 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> --- .../mcp_server/discoverable_endpoints.py | 5 +- .../outbound_credentials/adapter.py | 7 ++- .../proxy/_experimental/mcp_server/server.py | 4 +- .../mcp/test_mcp_caller_sign_in.py | 43 +++++++++++++++ .../mcp_server/test_discoverable_endpoints.py | 50 +++++++++++++++++ .../mcp_server/test_mcp_server.py | 55 +++++++++++++++++++ 6 files changed, 159 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 6940e8d4f02..849c1feb524 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index 42947e39530..a4b723b5d67 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -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 = ( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b9dcb79dafe..7570ba36905 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index f80a0c086af..1d41be0d78c 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -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() == () 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 7349db67d5d..cb24b3a32b7 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 @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 48d43519b6a..0c4c6ce1ab8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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}"')