mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates
Discovery-list cache identity now uses the hashed token instead of the raw api_key and treats MCPJWTSigner-signed servers as per caller. Server definition changes also drop the cached upstream OAuth metadata. OpenAPI listings look tools up under the normalized registry prefix with the separator, so an overlapping sibling prefix no longer leaks into the list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
9fd78ff6f4
commit
444894be2a
4 changed files with 180 additions and 28 deletions
|
|
@ -141,6 +141,17 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None:
|
|||
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
|
||||
|
||||
|
||||
def invalidate_oauth_metadata_cache(server_id: str) -> None:
|
||||
"""Drop cached upstream IdP metadata for a server whose definition changed."""
|
||||
for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]:
|
||||
del _OAUTH_METADATA_CACHE[cache_key]
|
||||
for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]:
|
||||
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
|
||||
if lock is None or lock.locked():
|
||||
continue
|
||||
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
|
||||
|
||||
|
||||
def encode_state_with_base_url(
|
||||
base_url: str,
|
||||
original_state: str,
|
||||
|
|
|
|||
|
|
@ -2664,7 +2664,7 @@ class MCPServerManager:
|
|||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_legacy_delegate_auth_if_applicable(new_server, source="config")
|
||||
_warn_config_id_jag_server_outruns_sso(new_server)
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._invalidate_server_definition_caches(server_id)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
self._set_oauth_discovery_deferred(
|
||||
server_id,
|
||||
|
|
@ -2868,7 +2868,7 @@ class MCPServerManager:
|
|||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server.server_id)
|
||||
self._invalidate_server_definition_caches(server.server_id)
|
||||
prefix_root: Final = normalize_server_name(get_server_prefix(server))
|
||||
if server.spec_path and prefix_root:
|
||||
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
|
||||
|
|
@ -3275,7 +3275,7 @@ class MCPServerManager:
|
|||
# env_vars_are_encrypted=False.
|
||||
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -3312,7 +3312,7 @@ class MCPServerManager:
|
|||
previous_server=self.registry[mcp_server.server_id],
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -4492,24 +4492,16 @@ class MCPServerManager:
|
|||
if server.spec_path:
|
||||
# OpenAPI tools were stored in the registry under the prefix
|
||||
# active at registration time — fetch by that same prefix.
|
||||
_tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server))
|
||||
registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR
|
||||
_tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix)
|
||||
tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(_tools)
|
||||
# OpenAPI tools are stored in the registry with their prefix already
|
||||
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
|
||||
# through _create_prefixed_tools — that would add the prefix a second
|
||||
# time producing "test_petstore-test_petstore-getinventory".
|
||||
if not add_prefix:
|
||||
prefix: Final = get_server_prefix(server)
|
||||
sep: Final = MCP_TOOL_PREFIX_SEPARATOR
|
||||
tools = [
|
||||
(
|
||||
t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]})
|
||||
if t.name.startswith(f"{prefix}{sep}")
|
||||
else t
|
||||
)
|
||||
for t in tools
|
||||
]
|
||||
return tools
|
||||
if add_prefix:
|
||||
return tools
|
||||
return [t.model_copy(update=MappingProxyType({"name": t.name[len(registry_prefix) :]})) for t in tools]
|
||||
else:
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
self._remember_upstream_initialize_instructions(server, client)
|
||||
|
|
@ -4558,6 +4550,35 @@ class MCPServerManager:
|
|||
self._resource_discovery_cache.invalidate(server_id)
|
||||
self._template_discovery_cache.invalidate(server_id)
|
||||
|
||||
def _invalidate_server_definition_caches(self, server_id: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton
|
||||
invalidate_oauth_metadata_cache,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _discovers_per_caller(self, server: MCPServer) -> bool:
|
||||
return (
|
||||
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)
|
||||
or self._signs_caller_identity_upstream(server)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _signs_caller_identity_upstream(server: MCPServer) -> bool:
|
||||
"""Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may
|
||||
tailor its catalog to the caller even though the server itself is configured as shared."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server
|
||||
get_mcp_jwt_signer,
|
||||
)
|
||||
|
||||
if get_mcp_jwt_signer() is None:
|
||||
return False
|
||||
return not any(k.lower() == "authorization" for k in (server.static_headers or {}))
|
||||
|
||||
def _discovery_key(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -4568,25 +4589,18 @@ class MCPServerManager:
|
|||
subject_token: str | None,
|
||||
credential_fingerprint: str | None = None,
|
||||
) -> _DiscoveryKey:
|
||||
per_user: Final = (
|
||||
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)
|
||||
)
|
||||
per_user: Final = self._discovers_per_caller(server)
|
||||
if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token):
|
||||
return server.server_id, None
|
||||
identity: Final = (
|
||||
(user_api_key_auth.user_id, user_api_key_auth.api_key)
|
||||
if per_user and user_api_key_auth is not None
|
||||
else None
|
||||
(user_api_key_auth.user_id, user_api_key_auth.token) if per_user and user_api_key_auth is not None else None
|
||||
)
|
||||
material: Final = json.dumps(
|
||||
(identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint),
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return server.server_id, hashlib.sha256(material.encode()).hexdigest()
|
||||
return server.server_id, hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest()
|
||||
|
||||
async def get_prompts_from_server(
|
||||
self,
|
||||
|
|
@ -6704,7 +6718,7 @@ class MCPServerManager:
|
|||
|
||||
for server_id in previous_registry.keys() | registered_registry.keys():
|
||||
if previous_registry.get(server_id) != registered_registry.get(server_id):
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._invalidate_server_definition_caches(server_id)
|
||||
self.registry = registered_registry
|
||||
_warn_on_shared_identifier_prefixes(registered_registry.values())
|
||||
# A discovery task may have published into ``previous_registry`` while
|
||||
|
|
|
|||
|
|
@ -12611,3 +12611,38 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
|
|||
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
|
||||
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
|
||||
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_server_drops_cached_upstream_oauth_metadata():
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="oauth-cache-server",
|
||||
name="oauth_cache_server",
|
||||
url="http://old-upstream/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
stale_key: Final = (server.server_id, server.url)
|
||||
other_key: Final = ("other-server", "http://other/mcp")
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"})
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"})
|
||||
try:
|
||||
await manager.update_server(
|
||||
LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.name,
|
||||
url="http://new-upstream/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
)
|
||||
assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
finally:
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None)
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None)
|
||||
|
|
|
|||
|
|
@ -14058,6 +14058,98 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) ->
|
|||
assert "first" not in str(first)
|
||||
assert "second" not in str(second)
|
||||
|
||||
from litellm.proxy._types import hash_token
|
||||
|
||||
same_user_other_token: Final = manager._discovery_key(
|
||||
server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-second")), None, None, None, None
|
||||
)
|
||||
same_token_no_key: Final = manager._discovery_key(
|
||||
server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-first")), None, None, None, None
|
||||
)
|
||||
with_key: Final = manager._discovery_key(
|
||||
server, UserAPIKeyAuth(user_id="first", api_key="sk-first"), None, None, None, None
|
||||
)
|
||||
assert same_user_other_token != with_key
|
||||
assert same_token_no_key == with_key
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("signer", "static_headers", "shared"),
|
||||
[
|
||||
pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"),
|
||||
pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"),
|
||||
pytest.param(None, None, True, id="no-signer-stays-shared"),
|
||||
],
|
||||
)
|
||||
def test_jwt_signer_makes_a_shared_server_discover_per_caller(signer, static_headers, shared) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = _discovery_server().model_copy(update={"static_headers": static_headers})
|
||||
alice: Final = UserAPIKeyAuth(user_id="alice", token="hashed-alice")
|
||||
bob: Final = UserAPIKeyAuth(user_id="bob", token="hashed-bob")
|
||||
|
||||
with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam
|
||||
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
|
||||
return_value=signer,
|
||||
):
|
||||
for_alice: Final = manager._discovery_key(server, alice, None, None, None, None)
|
||||
for_bob: Final = manager._discovery_key(server, bob, None, None, None, None)
|
||||
|
||||
assert (for_alice == for_bob) is shared
|
||||
|
||||
|
||||
def _register_local_tool(name: str, description: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
async def _handler(**kwargs):
|
||||
return None
|
||||
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=name, description=description, input_schema={"type": "object"}, handler=_handler
|
||||
)
|
||||
|
||||
|
||||
def _openapi_server(name: str) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=f"{name}-id", name=name, alias=name, transport=MCPTransport.http, url=None, spec_path="/spec.yaml"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openapi_listing_ignores_overlapping_server_prefix() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
|
||||
for prefix in ("pet-", "petstore-"):
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix(prefix)
|
||||
_register_local_tool("pet-list", "Local pet tool")
|
||||
_register_local_tool("petstore-list", "Foreign petstore tool")
|
||||
try:
|
||||
prefixed: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=True)
|
||||
bare: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=False)
|
||||
finally:
|
||||
for prefix in ("pet-", "petstore-"):
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix(prefix)
|
||||
|
||||
assert [t.name for t in prefixed] == ["pet-list"]
|
||||
assert [t.name for t in bare] == ["list"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefix() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-")
|
||||
_register_local_tool("pet_store-list", "Pet store tool")
|
||||
try:
|
||||
listed: Final = await manager._get_tools_from_server(server=_openapi_server("pet store"), add_prefix=False)
|
||||
finally:
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-")
|
||||
|
||||
assert [t.name for t in listed] == ["list"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_retries_cancelled_fetches() -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue