fix(mcp): drop the OAuth discovery cache when a server definition changes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-26 08:44:33 +00:00
parent bef557f3ed
commit 3650243b90
3 changed files with 52 additions and 0 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

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

View file

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