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>
This commit is contained in:
joshua 2026-09-19 01:37:46 +00:00
parent 48421a4723
commit d03f1d14c9
4 changed files with 183 additions and 74 deletions

View file

@ -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)

View file

@ -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)

View file

@ -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()

View file

@ -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": []})