fix(mcp): count queued OAuth metadata fetchers so invalidation survives lock handoff

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-28 22:58:32 +00:00
parent c0099a45de
commit 9786e3509a
2 changed files with 80 additions and 12 deletions

View file

@ -3,7 +3,8 @@ import html as _html
import json
import secrets
import time
from collections.abc import Callable, Mapping
from collections.abc import AsyncIterator, Callable, Mapping
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
@ -107,6 +108,10 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128
# Per-(server_id, resource_url) async locks so concurrent discovery requests
# coalesce onto a single upstream fetch instead of issuing N parallel calls.
_OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {}
# Callers inside ``_oauth_metadata_fetch_slot`` per cache key, lock waiters included. ``Lock.locked()``
# reads False between one holder's release and the next waiter's wake-up, so it cannot tell an
# idle lock from one being handed off.
_OAUTH_METADATA_FETCHERS: Final[dict[tuple[str, str], int]] = {}
# 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. Only servers with a fetch
# in flight carry an entry; the rest are pruned with the cache.
@ -134,13 +139,10 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None:
for cache_key in cache_keys_by_expiry[:overflow]:
_OAUTH_METADATA_CACHE.pop(cache_key, None)
# Drop locks whose cache entry has been evicted and that aren't currently
# held; held locks stay so in-flight callers continue to coalesce.
# Drop locks whose cache entry has been evicted and that nobody holds or
# waits on; the rest stay so in-flight callers continue to coalesce.
for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS):
if cache_key in _OAUTH_METADATA_CACHE:
continue
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
if lock is None or lock.locked():
if cache_key in _OAUTH_METADATA_CACHE or cache_key in _OAUTH_METADATA_FETCHERS:
continue
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
@ -149,7 +151,21 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> 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)
return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS)
@asynccontextmanager
async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]:
_OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1
try:
async with _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()):
yield
finally:
remaining: Final = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) - 1
if remaining > 0:
_OAUTH_METADATA_FETCHERS[cache_key] = remaining
else:
_OAUTH_METADATA_FETCHERS.pop(cache_key, None)
def invalidate_oauth_metadata_cache(server_id: str) -> None:
@ -161,8 +177,7 @@ def invalidate_oauth_metadata_cache(server_id: str) -> 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]:
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
if lock is None or lock.locked():
if cache_key in _OAUTH_METADATA_FETCHERS:
continue
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
@ -2386,8 +2401,7 @@ async def fetch_upstream_oauth_protected_resource(
if cached is not None and cached[0] > now:
return cached[1]
lock: Final = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock())
async with lock:
async with _oauth_metadata_fetch_slot(cache_key):
now = time.time()
cached = _OAUTH_METADATA_CACHE.get(cache_key)
if cached is not None and cached[0] > now:

View file

@ -12693,6 +12693,60 @@ async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cach
discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None)
@pytest.mark.asyncio
async def test_fetch_waiting_on_a_lock_handoff_stays_tracked_through_invalidation():
import asyncio
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
fetch_upstream_oauth_protected_resource,
invalidate_oauth_metadata_cache,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="handoff-server", name="handoff", url="http://upstream/mcp", transport=MCPTransport.http
)
cache_key: Final = (server.server_id, server.url)
started: Final = asyncio.Event()
release: Final = asyncio.Event()
async def slow_get(url: str, headers: dict[str, str]) -> MagicMock:
started.set()
await release.wait()
return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["pre-save-idp"]}))
client = MagicMock()
client.get = slow_get
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None)
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=client,
):
async with discoverable_endpoints._oauth_metadata_fetch_slot(cache_key):
shared_lock: Final = discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS[cache_key]
waiting: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server))
for _ in range(3):
await asyncio.sleep(0)
assert not started.is_set() and not waiting.done()
invalidate_oauth_metadata_cache(server.server_id)
assert discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.get(cache_key) is shared_lock
assert discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id)
await started.wait()
invalidate_oauth_metadata_cache(server.server_id)
release.set()
assert await waiting == {"authorization_servers": ["pre-save-idp"]}
assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE
assert not discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id)
finally:
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None)
discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
discoverable_endpoints._OAUTH_METADATA_FETCHERS.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