From d03f1d14c9b86a42856a6ec4a7d521e8abbf78be Mon Sep 17 00:00:00 2001 From: joshua Date: Sat, 19 Sep 2026 01:37:46 +0000 Subject: [PATCH] fix(mcp): apply client allowlist and tools-style upstream headers to catalog routes The prompt and resource REST routes now call reject_disallowed_mcp_client before resolving the acting user or the server, matching /mcp-rest/tools/list. Prompt, resource and resource template listing share the tools listing header preparation, so ${ENV} static headers are interpolated and the MCPJWTSigner token is injected when nothing else carries an Authorization Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 145 +++++++++--------- .../mcp_server/rest_endpoints.py | 1 + .../mcp_server/test_mcp_server_manager.py | 63 ++++++++ .../mcp_server/test_rest_endpoints.py | 48 ++++++ 4 files changed, 183 insertions(+), 74 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 110062b081b..2505ea3dca4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4389,56 +4389,13 @@ class MCPServerManager: client = None try: - # Tool *listing* must not be blocked by missing per-user env vars — - # the server's tools should still appear so the client connects. The - # friendly "missing vars" error is raised only on the tool-*call* - # path (see _call_regular_mcp_tool). - resolved_static_headers: Final = await self._resolve_static_headers_with_env_vars( - server, user_api_key_auth, raise_on_missing=False + list_headers: Final = await self._resolve_list_headers( + server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + raw_headers=raw_headers, ) - if resolved_static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(resolved_static_headers) - - # MCPJWTSigner: inject signed JWT for tools/list (list path skips pre_call_hook). - # Skip entirely when the signer is not configured (avoid an unnecessary - # dict copy on every list call), when the server has its own static - # Authorization header, when a per-user mcp_auth_header has already - # been resolved, or when the caller already supplied an Authorization - # entry in extra_headers (e.g. a per-user OAuth token resolved - # upstream) — admin-configured static auth and per-user OAuth must - # take precedence so the signer doesn't silently overwrite e.g. an - # upstream API key or a user's OAuth token (MCPClient._get_auth_headers - # applies extra_headers after writing Authorization from auth_value, so - # an injected JWT would otherwise clobber the per-user token). - if user_api_key_auth is not None and not server.spec_path: - from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( - get_mcp_jwt_signer, - inject_mcp_jwt_headers_for_upstream, - ) - - static_headers: Final = server.static_headers or {} - has_static_authorization: Final = any( - isinstance(k, str) and k.lower() == "authorization" for k in static_headers - ) - has_extra_authorization: Final = bool(extra_headers) and any( - isinstance(k, str) and k.lower() == "authorization" for k in (extra_headers or {}) - ) - - if ( - get_mcp_jwt_signer() is not None - and not has_static_authorization - and not mcp_auth_header - and not has_extra_authorization - ): - extra_headers = await inject_mcp_jwt_headers_for_upstream( - user_api_key_dict=user_api_key_auth, - extra_headers=extra_headers, - raw_headers=raw_headers, - for_list_tools=True, - ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) # token_exchange (OBO) discovery needs the caller's token too: list it with the user's own @@ -4453,7 +4410,7 @@ class MCPServerManager: client = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, + extra_headers=list_headers, stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, @@ -4492,6 +4449,52 @@ class MCPServerManager: except Exception as e: _raise_single_server_list_failure(e, server, "tools") + async def _resolve_list_headers( + self, + server: MCPServer, + *, + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | dict[str, str] | None, + extra_headers: dict[str, str] | None, + raw_headers: dict[str, str] | None, + ) -> dict[str, str] | None: + """Listing stays best-effort on missing per-user env vars, and the JWT signer never overrides an + Authorization already supplied by static headers, a per-user auth header, or extra_headers.""" + resolved_static_headers: Final = await self._resolve_static_headers_with_env_vars( + server, user_api_key_auth, raise_on_missing=False + ) + headers: Final = ( + dict( + chain( + extra_headers.items() if extra_headers else (), + resolved_static_headers.items() if resolved_static_headers else (), + ) + ) + or None + ) + if user_api_key_auth is None or server.spec_path: + return headers + + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( + get_mcp_jwt_signer, + inject_mcp_jwt_headers_for_upstream, + ) + + has_static_authorization: Final = any( + isinstance(k, str) and k.lower() == "authorization" for k in (server.static_headers or {}) + ) + has_extra_authorization: Final = any( + isinstance(k, str) and k.lower() == "authorization" for k in (extra_headers or {}) + ) + if get_mcp_jwt_signer() is None or has_static_authorization or mcp_auth_header or has_extra_authorization: + return headers + return await inject_mcp_jwt_headers_for_upstream( + user_api_key_dict=user_api_key_auth, + extra_headers=headers, + raw_headers=raw_headers, + for_list_tools=True, + ) + def _invalidate_discovery_lists(self, server_id: str) -> None: self._prompt_discovery_cache.invalidate(server_id) self._resource_discovery_cache.invalidate(server_id) @@ -4538,14 +4541,12 @@ class MCPServerManager: raise_on_error: bool = False, ) -> list[Prompt]: try: - headers: Final = ( - dict( - chain( - extra_headers.items() if extra_headers else (), - server.static_headers.items() if server.static_headers else (), - ) - ) - or None + headers: Final = await self._resolve_list_headers( + server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + raw_headers=raw_headers, ) stdio_env: Final = self._build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) @@ -4584,14 +4585,12 @@ class MCPServerManager: raise_on_error: bool = False, ) -> list[Resource]: try: - headers: Final = ( - dict( - chain( - extra_headers.items() if extra_headers else (), - server.static_headers.items() if server.static_headers else (), - ) - ) - or None + headers: Final = await self._resolve_list_headers( + server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + raw_headers=raw_headers, ) stdio_env: Final = self._build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) @@ -4630,14 +4629,12 @@ class MCPServerManager: raise_on_error: bool = False, ) -> list[ResourceTemplate]: try: - headers: Final = ( - dict( - chain( - extra_headers.items() if extra_headers else (), - server.static_headers.items() if server.static_headers else (), - ) - ) - or None + headers: Final = await self._resolve_list_headers( + server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + raw_headers=raw_headers, ) stdio_env: Final = self._build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 2565089b051..ff66f8816aa 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -1094,6 +1094,7 @@ if MCP_AVAILABLE: server_id: str, user_api_key_dict: UserAPIKeyAuth, ) -> _CatalogServerContext: + reject_disallowed_mcp_client(request.headers, user_api_key_dict) acting_auth: Final = await acting_user_auth(user_api_key_dict) _, canonical_server_id = await _resolve_allowed_mcp_servers_with_ip_filter(request, acting_auth, server_id) server: Final = global_mcp_server_manager.get_mcp_server_by_id(canonical_server_id) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 4e315b605d6..9e6563acce0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -3997,6 +3997,69 @@ class TestMCPServerManager: assert exc_info.value.status_code == 401 assert exc_info.value.www_authenticate == challenge + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("manager_method", "client_method"), + [ + ("get_prompts_from_server", "list_prompts"), + ("get_resources_from_server", "list_resources"), + ("get_resource_templates_from_server", "list_resource_templates"), + ], + ) + async def test_catalog_fetch_prepares_upstream_headers_like_tools(self, manager_method, client_method, monkeypatch): + """Prompt and resource listings must reach the upstream with the same credentials the tools + listing sends: ``${NAME}`` static headers interpolated from the server's env vars and the + MCPJWTSigner token injected when nothing else carries an Authorization.""" + import jwt + + import litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer as jwt_signer_module + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import MCPJWTSigner + + monkeypatch.setattr(jwt_signer_module, "_mcp_jwt_signer_instance", None) + MCPJWTSigner( + guardrail_name="catalog-jwt-signer", + event_hook="pre_mcp_call", + default_on=True, + issuer="https://litellm.example.com", + audience="mcp", + ) + manager = MCPServerManager() + server = MCPServer( + server_id="server-1", + name="alias-server", + alias="alias-server", + server_name="alias-server", + url="https://example.com", + transport=MCPTransport.http, + static_headers={"X-Tenant": "${TENANT}"}, + env_vars=[{"name": "TENANT", "value": "acme", "scope": "global"}], + ) + mock_client = AsyncMock() + setattr(mock_client, client_method, AsyncMock(return_value=[])) + mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash") + upstream_headers: list[dict[str, str] | None] = [] + + async def create_client(**kwargs): + upstream_headers.append(kwargs["extra_headers"]) + return mock_client + + with patch.object(manager, "_create_mcp_client", create_client): + await getattr(manager, manager_method)( + server, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-test", user_id="alice"), + extra_headers={"X-Caller": "dashboard"}, + ) + + assert len(upstream_headers) == 1 + sent = upstream_headers[0] + assert sent is not None + assert sent["X-Caller"] == "dashboard" + assert sent["X-Tenant"] == "acme" + claims = jwt.decode(sent["Authorization"].removeprefix("Bearer "), options={"verify_signature": False}) + assert claims["sub"] == "alice" + assert claims["scope"] == "mcp:tools/list" + @pytest.mark.asyncio async def test_read_resource_from_server_success(self): manager = MCPServerManager() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 15d2ecc9124..f2c1667e6d3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -4653,3 +4653,51 @@ class TestClientAllowlistOnRestRoutes: assert denied.value.detail["error"] == "Forbidden" assert "'claude-code'" in denied.value.detail["details"] acting.assert_not_awaited() + + @pytest.mark.parametrize("route_name", ("list_prompts_rest_api", "list_resources_rest_api")) + @pytest.mark.parametrize( + ("caller", "headers", "expected_fragment"), + ( + (UserAPIKeyAuth(jwt_claims={"azp": "claude-code"}), {"x-mcp-client": "antigravity-cli"}, "'claude-code'"), + (UserAPIKeyAuth(), {}, "no 'x-mcp-client' header"), + ), + ) + async def test_catalog_routes_reject_unlisted_clients_before_resolving_servers( + self, + monkeypatch: pytest.MonkeyPatch, + route_name: str, + caller: UserAPIKeyAuth, + headers: dict[str, str], + expected_fragment: str, + ) -> None: + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", _CLIENT_ALLOWLIST_SETTINGS, raising=False) + acting: Final = AsyncMock() + monkeypatch.setattr(rest_endpoints, "acting_user_auth", acting, raising=False) + request: Final = _build_request(headers, path=f"/mcp-rest/{route_name}", method="GET") + + with pytest.raises(HTTPException) as denied: + await getattr(rest_endpoints, route_name)(request, server_id="server-1", user_api_key_dict=caller) + + assert denied.value.status_code == 403 + assert denied.value.detail["error"] == "Forbidden" + assert expected_fragment in denied.value.detail["details"] + assert "mcp_allowed_clients" in denied.value.detail["details"] + acting.assert_not_awaited() + + @pytest.mark.parametrize("route_name", ("list_prompts_rest_api", "list_resources_rest_api")) + async def test_catalog_routes_admit_listed_clients(self, monkeypatch: pytest.MonkeyPatch, route_name: str) -> None: + catalog_suite: Final = TestListPromptsAndResourcesRestAPI() + server: Final = catalog_suite._stub_server() + catalog_suite._grant(monkeypatch, server, allowed=[server.server_id]) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", _CLIENT_ALLOWLIST_SETTINGS, raising=False) + for method in ("get_prompts_from_server", "get_resources_from_server", "get_resource_templates_from_server"): + monkeypatch.setattr(rest_endpoints.global_mcp_server_manager, method, AsyncMock(return_value=[])) + request: Final = _build_request( + {"x-mcp-client": "antigravity-cli"}, path=f"/mcp-rest/{route_name}", method="GET" + ) + + result: Final = await getattr(rest_endpoints, route_name)( + request, server_id=server.server_id, user_api_key_dict=UserAPIKeyAuth() + ) + + assert result.model_dump() in ({"prompts": []}, {"resources": [], "resource_templates": []})