mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
48421a4723
commit
d03f1d14c9
4 changed files with 183 additions and 74 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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": []})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue