fix(mcp): keep OAuth metadata generations only while a fetch is in flight

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-28 22:37:53 +00:00
parent 9ce2803f42
commit c0099a45de
2 changed files with 29 additions and 2 deletions

View file

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

View file

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