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 b1f52bd5bcc..5647a44f814 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4555,8 +4555,13 @@ class MCPServerManager: 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) self._listed_tools_by_server_id.pop(server_id, None) + invalidate_oauth_metadata_cache(server_id) def _discovers_per_caller(self, server: MCPServer) -> bool: return ( 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..e6aba5f1727 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,39 @@ 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 + from litellm.proxy._types import 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)