mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
34ec9da012
commit
4ae476995f
2 changed files with 102 additions and 16 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue