mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
c0099a45de
commit
9786e3509a
2 changed files with 80 additions and 12 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue