fix(mcp): keep minted JWT out of the catalog discovery cache key

The MCPJWTSigner mints a token with a fresh iat on every prompt, resource and template listing, and that Authorization header was hashed into the discovery cache key, so the 60s cache missed on every request that crossed a second boundary. The listing now keys on the user the token was signed for instead of the token itself, keeping per-user isolation while reusing cached results across refreshed JWTs. Tools listing was unaffected and only adapts to the new return type.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua 2026-09-19 20:12:47 +00:00
parent 34ec9da012
commit 4ae476995f
2 changed files with 102 additions and 16 deletions

View file

@ -319,6 +319,12 @@ class _OAuthDiscoveryStale:
_OAuthDiscoveryOutcome: TypeAlias = _OAuthDiscoveryResolved | _OAuthDiscoveryFailed | _OAuthDiscoveryStale
@dataclass(frozen=True, slots=True)
class _ListHeaders:
upstream: dict[str, str] | None
signed_for_user: bool
@dataclass(frozen=True, slots=True)
class _OAuthDiscoverySlot:
server_id: str
@ -4403,7 +4409,7 @@ class MCPServerManager:
client = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=list_headers,
extra_headers=list_headers.upstream,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
@ -4482,7 +4488,7 @@ class MCPServerManager:
mcp_auth_header: str | dict[str, str] | None,
extra_headers: dict[str, str] | None,
raw_headers: dict[str, str] | None,
) -> dict[str, str] | None:
) -> _ListHeaders:
"""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(
@ -4498,7 +4504,7 @@ class MCPServerManager:
or None
)
if user_api_key_auth is None or server.spec_path:
return headers
return _ListHeaders(headers, signed_for_user=False)
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
get_mcp_jwt_signer,
@ -4512,12 +4518,15 @@ class MCPServerManager:
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,
return _ListHeaders(headers, signed_for_user=False)
return _ListHeaders(
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,
),
signed_for_user=True,
)
def _invalidate_discovery_lists(self, server_id: str) -> None:
@ -4534,9 +4543,12 @@ class MCPServerManager:
stdio_env: dict[str, str] | None,
subject_token: str | None,
credential_fingerprint: str | None = None,
*,
signed_for_user: bool = False,
) -> _DiscoveryKey:
per_user: Final = (
server.requires_per_user_auth
signed_for_user
or server.requires_per_user_auth
or self._references_per_user_env_var(server)
or server.delegate_auth_to_upstream
or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)
@ -4578,14 +4590,21 @@ class MCPServerManager:
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
extra_headers=headers.upstream,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
server,
user_api_key_auth,
mcp_auth_header,
extra_headers,
stdio_env,
subject_token,
credential_fingerprint,
signed_for_user=headers.signed_for_user,
)
async def fetch() -> list[Prompt]:
@ -4622,14 +4641,21 @@ class MCPServerManager:
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
extra_headers=headers.upstream,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
server,
user_api_key_auth,
mcp_auth_header,
extra_headers,
stdio_env,
subject_token,
credential_fingerprint,
signed_for_user=headers.signed_for_user,
)
async def fetch() -> list[Resource]:
@ -4666,14 +4692,21 @@ class MCPServerManager:
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
extra_headers=headers.upstream,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
server,
user_api_key_auth,
mcp_auth_header,
extra_headers,
stdio_env,
subject_token,
credential_fingerprint,
signed_for_user=headers.signed_for_user,
)
async def fetch() -> list[ResourceTemplate]:

View file

@ -4072,6 +4072,59 @@ class TestMCPServerManager:
assert claims["sub"] == "alice"
assert claims["scope"] == "mcp:tools/list"
@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_discovery_cache_survives_jwt_signer_and_isolates_users(
self, manager_method: str, client_method: str, monkeypatch: pytest.MonkeyPatch
) -> None:
"""The MCPJWTSigner mints a fresh token on every listing, so the token itself must not be part of
the discovery cache key; the user it was signed for must be, since the upstream sees that identity."""
from types import SimpleNamespace
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)
signer_clock = iter(range(1_700_000_000, 1_700_000_100))
monkeypatch.setattr(jwt_signer_module, "time", SimpleNamespace(time=lambda: next(signer_clock)))
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,
)
mock_client = AsyncMock()
upstream_call = AsyncMock(return_value=[])
setattr(mock_client, client_method, upstream_call)
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice")
bob = UserAPIKeyAuth(api_key="sk-bob", user_id="bob")
with patch.object(manager, "_create_mcp_client", AsyncMock(return_value=mock_client)):
await getattr(manager, manager_method)(server, user_api_key_auth=alice)
await getattr(manager, manager_method)(server, user_api_key_auth=alice)
assert upstream_call.await_count == 1
await getattr(manager, manager_method)(server, user_api_key_auth=bob)
assert upstream_call.await_count == 2
@pytest.mark.asyncio
async def test_read_resource_from_server_success(self):
manager = MCPServerManager()