mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(mcp): drop minted JWT headers from catalog cache fingerprint
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
edf100913c
commit
32f1acadb4
3 changed files with 34 additions and 17 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue