diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d2851cf4fbf..d4e95e7ad70 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -108,7 +108,8 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} # Per-server_id generation, bumped on invalidation so a fetch that started before the server -# definition changed cannot repopulate the cache with the stale reply. +# definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch +# in flight carry an entry; the rest are pruned with the cache. _OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( @@ -143,10 +144,20 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + for server_id in [sid for sid in _OAUTH_METADATA_GENERATIONS if not _oauth_metadata_fetch_in_flight(sid)]: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) + + +def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: + return any(lock.locked() for cache_key, lock in _OAUTH_METADATA_FETCH_LOCKS.items() if cache_key[0] == server_id) + def invalidate_oauth_metadata_cache(server_id: str) -> None: """Drop cached upstream IdP metadata for a server whose definition changed.""" - _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 + if _oauth_metadata_fetch_in_flight(server_id): + _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 + else: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) 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]: 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 7a2af8dfcca..d9135a999b6 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 @@ -12686,6 +12686,22 @@ async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cach release.set() assert await in_flight == {"authorization_servers": ["old-idp"]} assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + discoverable_endpoints._prune_oauth_metadata_cache() + assert server.server_id not in discoverable_endpoints._OAUTH_METADATA_GENERATIONS finally: discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) + + +def test_invalidating_an_idle_server_leaves_no_generation_behind(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import invalidate_oauth_metadata_cache + + server_ids: Final = tuple(f"churned-server-{i}" for i in range(50)) + try: + for server_id in server_ids: + invalidate_oauth_metadata_cache(server_id) + assert not set(server_ids) & set(discoverable_endpoints._OAUTH_METADATA_GENERATIONS) + finally: + for server_id in server_ids: + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server_id, None)