diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 4b456710057..1c1ce9a3bbd 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index cd679e3a03f..5a4d3fca5fd 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 1dba4bb7c6c..90b7f353aa7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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):