fix(mcp): drop minted JWT headers from catalog cache fingerprint
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua 2026-09-19 20:38:40 +00:00
parent edf100913c
commit 32f1acadb4
3 changed files with 34 additions and 17 deletions

View file

@ -331,8 +331,8 @@ class MCPClient:
if auth_value:
self.update_auth_value(auth_value)
async def discovery_auth_fingerprint(self) -> str:
return self._hash_discovery_auth(await self.prepare_request_auth())
async def discovery_auth_fingerprint(self, *, ignore_headers: frozenset[str] = frozenset()) -> str:
return self._hash_discovery_auth(await self.prepare_request_auth(), ignore_headers)
async def prepare_request_auth(self) -> httpx2.Request:
"""Preview the authenticated request without sending it, closing the auth flow afterwards."""
@ -349,8 +349,11 @@ class MCPClient:
await flow.aclose()
@staticmethod
def _hash_discovery_auth(request: httpx2.Request) -> str:
material: Final = json.dumps((str(request.url), tuple(sorted(request.headers.multi_items()))))
def _hash_discovery_auth(request: httpx2.Request, ignore_headers: frozenset[str] = frozenset()) -> str:
kept: Final = tuple(
sorted(item for item in request.headers.multi_items() if item[0].lower() not in ignore_headers)
)
material: Final = json.dumps((str(request.url), kept))
return hashlib.sha256(material.encode()).hexdigest()
def _create_transport_context(

View file

@ -323,6 +323,7 @@ _OAuthDiscoveryOutcome: TypeAlias = _OAuthDiscoveryResolved | _OAuthDiscoveryFai
class _ListHeaders:
upstream: dict[str, str] | None
signed_for_user: bool
minted: frozenset[str] = frozenset()
@dataclass(frozen=True, slots=True)
@ -4519,14 +4520,17 @@ class MCPServerManager:
)
if get_mcp_jwt_signer() is None or has_static_authorization or mcp_auth_header or has_extra_authorization:
return _ListHeaders(headers, signed_for_user=False)
signed: Final = 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,
)
unsigned_items: Final = frozenset((headers or {}).items())
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,
signed_for_user=True,
minted=frozenset(name.lower() for name, value in signed.items() if (name, value) not in unsigned_items),
)
def _invalidate_discovery_lists(self, server_id: str) -> None:
@ -4595,7 +4599,7 @@ class MCPServerManager:
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
credential_fingerprint: Final = await client.discovery_auth_fingerprint(ignore_headers=headers.minted)
key: Final = self._discovery_key(
server,
user_api_key_auth,
@ -4646,7 +4650,7 @@ class MCPServerManager:
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
credential_fingerprint: Final = await client.discovery_auth_fingerprint(ignore_headers=headers.minted)
key: Final = self._discovery_key(
server,
user_api_key_auth,
@ -4697,7 +4701,7 @@ class MCPServerManager:
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
credential_fingerprint: Final = await client.discovery_auth_fingerprint(ignore_headers=headers.minted)
key: Final = self._discovery_key(
server,
user_api_key_auth,

View file

@ -4098,6 +4098,7 @@ class TestMCPServerManager:
from types import SimpleNamespace
import litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer as jwt_signer_module
from litellm.experimental_mcp_client.client import MCPClient
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import MCPJWTSigner
@ -4120,19 +4121,28 @@ class TestMCPServerManager:
url="https://example.com",
transport=MCPTransport.http,
)
mock_client = AsyncMock()
setattr(mock_client, client_method, AsyncMock(side_effect=[[make_item("first")], [make_item("second")]]))
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
upstream_call = AsyncMock(side_effect=[[make_item("first")], [make_item("second")]])
seen_authorization: list[str] = []
async def create_client(**kwargs: object) -> MCPClient:
extra_headers = kwargs["extra_headers"]
assert isinstance(extra_headers, dict)
seen_authorization.append(extra_headers["Authorization"])
client = MCPClient(server_url=server.url, transport_type=MCPTransport.http, extra_headers=extra_headers)
setattr(client, client_method, upstream_call)
return client
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)):
with patch.object(manager, "_create_mcp_client", AsyncMock(side_effect=create_client)):
alice_first = await getattr(manager, manager_method)(server, user_api_key_auth=alice, add_prefix=False)
alice_second = await getattr(manager, manager_method)(server, user_api_key_auth=alice, add_prefix=False)
bob_first = await getattr(manager, manager_method)(server, user_api_key_auth=bob, add_prefix=False)
assert [item.name for item in alice_first] == ["first"]
assert alice_second == alice_first
assert [item.name for item in bob_first] == ["second"]
assert len(seen_authorization) == 3 and len(set(seen_authorization)) == 3
@pytest.mark.asyncio
async def test_read_resource_from_server_success(self):