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:
yucheng 2026-09-28 20:01:28 +00:00
parent 9fd78ff6f4
commit 444894be2a
4 changed files with 180 additions and 28 deletions

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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: