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