diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9e579500f69..5fb84c64258 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1625,21 +1625,19 @@ if MCP_AVAILABLE: registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to( server_name, client_ip=client_ip ) - obo_without_subject = ( - registry_pick is not None - and registry_pick.auth_type == MCPAuth.oauth2_token_exchange - and not oauth2_headers - ) allowed_single = ( await operations._get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip ) - if registry_pick and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1 + if registry_pick and mcp_servers is not None and len(mcp_servers) == 1 else () ) granted = next(iter(allowed_single), None) server = granted if granted is not None else registry_pick granted_single = granted is not None + obo_without_subject = ( + server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers + ) if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids: # Caller's narrowed scope excludes this server — skip the # preemptive challenge and let downstream authorization diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index e73168a4448..7db98279e3a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -8985,6 +8985,100 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: ) assert guardrail.asked_about == ["d-id"] + @staticmethod + def _register_obo_alias_collision() -> tuple[MCPServer, MCPServer]: + obo: Final = MCPServer( + server_id="o-id", + name="obo_server", + server_name="obo_server", + alias="obo", + url="https://obo.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + token_exchange_endpoint="https://idp.test/token", + client_id="cid", + client_secret="csecret", + ) + plain: Final = MCPServer( + server_id="p-id", + name="obo", + server_name="obo", + url="https://plain.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + mcp_operations.global_mcp_server_manager.registry.update({"o-id": obo, "p-id": plain}) + return obo, plain + + @staticmethod + async def _connect_to(route_name: str) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + 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-collision", user_id="u-1"), + client_ip=None, + raw_headers={"x-litellm-api-key": "sk-collision"}, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("route_name", ["obo", "OBO"], ids=["exact_alias", "case_variant"]) + async def test_granted_plain_server_connects_past_an_ungranted_obo_alias_holder(self, route_name): + _, plain = self._register_obo_alias_collision() + preflight: Final = AsyncMock() + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["p-id"]) + ), + patch.object( # test-quality-ok: a real exchanger would call an IdP; this one records which server ran + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight + ), + ): + await self._connect_to(route_name) + + assert preflight.await_args is not None, "the granted plain server must reach the preflight, not a 401" + assert preflight.await_args.kwargs["server"] is plain + + @pytest.mark.asyncio + @pytest.mark.parametrize("route_name", ["obo", "obo_server"], ids=["alias", "server_name"]) + async def test_granted_obo_server_still_challenges_without_a_subject(self, route_name, monkeypatch): + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + self._register_obo_alias_collision() + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["o-id"]) + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to(route_name) + + assert exc.value.status_code == 401 + assert ((exc.value.headers or {}).get("WWW-Authenticate") or "").startswith( + 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/obo"' + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("route_name", ["obo", "OBO"], ids=["exact_alias", "case_variant"]) + async def test_key_granting_neither_server_is_not_preflighted_for_the_plain_one(self, route_name): + self._register_obo_alias_collision() + preflight: Final = AsyncMock() + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=[]) + ), + patch.object( # test-quality-ok: a real exchanger would call an IdP; this one records which server ran + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to(route_name) + + assert exc.value.status_code == 401 + assert preflight.await_count == 0 + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown():