diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..91da7204435 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 31896d9ddc5..d2a4b6ebad9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f9a0075e530..0a92b548d62 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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) 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 70ef4312f4c..85b4aaa11eb 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 @@ -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: