From c129ea4fc99a583fb212150b659d46818baeb7b7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 17:25:39 -0700 Subject: [PATCH] fix(mcp): scope OpenAPI listings to the exact server prefix and drop upstream OAuth metadata when a server is saved (#43608) * fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates Discovery-list cache identity now uses the hashed token instead of the raw api_key and treats MCPJWTSigner-signed servers as per caller. Server definition changes also drop the cached upstream OAuth metadata. OpenAPI listings look tools up under the normalized registry prefix with the separator, so an overlapping sibling prefix no longer leaks into the list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep the discovery cache digest call unchanged so CodeQL matches the existing alert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): guard OAuth metadata cache writes with a per-server generation and drop unproven per-caller discovery keys An upstream metadata fetch that started before a server edit could store its stale reply after invalidate_oauth_metadata_cache ran. Invalidation now bumps a per-server generation and the fetch only stores when the generation it captured before I/O is unchanged. The MCPJWTSigner-based per-caller discovery classification and the api_key to token key change had no reproduction (the signer only injects on tools/list, and UserAPIKeyAuth hashes api_key in place), so both go back to the merge-base behavior. Integration coverage under tests/integration/mcp: overlapping OpenAPI aliases, a config-declared server name with a space, OAuth metadata refetch after a save, and the in-flight stale-write race Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * 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> * 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> * fix(mcp): keep a held OAuth metadata lock registered even when no fetcher slot claims it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): prove a peer worker drops stale upstream OAuth metadata after a save elsewhere Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 85 +++++++--- .../mcp_server/mcp_server_manager.py | 26 +-- tests/integration/mcp/test_mcp_management.py | 52 ++++++ .../mcp/test_oauth_configuration.py | 139 +++++++++++++++- .../mcp_server/test_discoverable_endpoints.py | 148 ++++++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 54 +++++++ 6 files changed, 470 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..7a0f59c3c2b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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,14 @@ _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. +_OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( tags=["mcp"], @@ -130,13 +139,52 @@ 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: + if cache_key in _OAUTH_METADATA_CACHE or not _oauth_metadata_lock_idle(cache_key): continue - lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) - if lock is None or lock.locked(): + _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(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS) + + +def _oauth_metadata_lock_idle(cache_key: tuple[str, str]) -> bool: + if cache_key in _OAUTH_METADATA_FETCHERS: + return False + lock: Final = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) + return lock is None or not lock.locked() + + +@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: + """Drop cached upstream IdP metadata for a server whose definition changed.""" + 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]: + if not _oauth_metadata_lock_idle(cache_key): continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -2360,12 +2408,19 @@ 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: return cached[1] + generation: Final = _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) + + def store(payload: dict | None, ttl_seconds: int) -> None: + if _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) != generation: + return + stored_at: Final = time.time() + _OAUTH_METADATA_CACHE[cache_key] = (stored_at + ttl_seconds, payload) + _prune_oauth_metadata_cache(stored_at) host_base: Final = f"{upstream.scheme}://{upstream.netloc}" candidates: Final = [f"{host_base}/.well-known/oauth-protected-resource"] @@ -2407,12 +2462,7 @@ async def fetch_upstream_oauth_protected_resource( ) continue if isinstance(payload, dict): - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_CACHE_TTL_SECONDS, - payload, - ) - _prune_oauth_metadata_cache(now) + store(payload, _OAUTH_METADATA_CACHE_TTL_SECONDS) return payload if len(network_errors) == len(candidates): @@ -2421,12 +2471,7 @@ async def fetch_upstream_oauth_protected_resource( # Negative-result caching: when no candidate yielded a usable payload, # remember that for a shorter TTL so we don't re-fetch on every # subsequent discovery request (and so the per-key lock can be pruned). - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS, - None, - ) - _prune_oauth_metadata_cache(now) + store(None, _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS) return None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 490d3072955..c0792c32de2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2673,7 +2673,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) _warn_legacy_delegate_auth_if_applicable(new_server, source="config") _warn_config_id_jag_server_outruns_sso(new_server) - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2877,7 +2877,7 @@ class MCPServerManager: global_mcp_tool_registry, ) - self._invalidate_discovery_lists(server.server_id) + self._invalidate_server_definition_caches(server.server_id) prefix_root: Final = normalize_server_name(get_server_prefix(server)) if server.spec_path and prefix_root: openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR @@ -3285,7 +3285,7 @@ class MCPServerManager: # env_vars_are_encrypted=False. new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3322,7 +3322,7 @@ class MCPServerManager: previous_server=self.registry[mcp_server.server_id], ) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -4504,16 +4504,16 @@ class MCPServerManager: if server.spec_path: # OpenAPI tools were stored in the registry under the prefix # active at registration time — fetch by that same prefix. - registered_prefix: Final = f"{get_server_prefix(server)}{MCP_TOOL_PREFIX_SEPARATOR}" + registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR registered: Final = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type( - global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server)) + global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix) ) registered_names: Final = MappingProxyType( - {t.name.removeprefix(registered_prefix): t.name for t in registered} + {t.name.removeprefix(registry_prefix): t.name for t in registered} ) guarded_openapi: Final = await self._guard_tool_catalog( server=server, - tools=[t.model_copy(update={"name": t.name.removeprefix(registered_prefix)}) for t in registered], + tools=[t.model_copy(update={"name": t.name.removeprefix(registry_prefix)}) for t in registered], proxy_logging_obj=proxy_logging_obj, user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, @@ -4582,6 +4582,14 @@ class MCPServerManager: self._resource_discovery_cache.invalidate(server_id) self._template_discovery_cache.invalidate(server_id) + def _invalidate_server_definition_caches(self, server_id: str) -> None: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton + invalidate_oauth_metadata_cache, + ) + + self._invalidate_discovery_lists(server_id) + invalidate_oauth_metadata_cache(server_id) + def _discovery_key( self, server: MCPServer, @@ -6792,7 +6800,7 @@ class MCPServerManager: for server_id in previous_registry.keys() | registered_registry.keys(): if previous_registry.get(server_id) != registered_registry.get(server_id): - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.registry = registered_registry _warn_on_shared_identifier_prefixes(registered_registry.values()) # A discovery task may have published into ``previous_registry`` while diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index 66c30a62bde..67cdbbff5a4 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -1,3 +1,4 @@ +import itertools import uuid from pathlib import Path from typing import Final @@ -7,10 +8,13 @@ import yaml from integration._support.client import Gateway, eventually from integration._support.mcp import ( McpCaller, + McpPeer, call_tool, delete_mcp, forget_mcp, + listed_tools, mcp_peer, + openapi_peer, register_mcp, tool_calls, tool_names, @@ -189,6 +193,54 @@ def test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide(gateway: Ga scenario.cleanups.callback(forget_mcp, gateway, winner) +def _openapi_server_lists_and_calls_only_its_own_tools( + gateway: Gateway, key: str, peer: McpPeer, identity: str +) -> None: + listed: Final = set(listed_tools(gateway, key, identity)) + assert listed == {"getpet", "createpet"}, (identity, listed) + peer.drain() + called: Final = call_tool(gateway, key, identity, "getpet", {"petId": "7"}) + assert called.status_code == 200, called.text + assert [(item["method"], item["path"]) for item in peer.drain()] == [("GET", "/pets/7")], identity + + +def test_openapi_listing_is_scoped_to_the_exact_alias_when_aliases_overlap(gateway: Gateway) -> None: + with openapi_peer() as short, openapi_peer() as long, gateway.scenario() as scenario: + stem: Final = "pet" + uuid.uuid4().hex[:8] + servers: Final = tuple( + (peer, alias, register_mcp(scenario, peer, alias)) + for peer, alias in ((short, stem), (long, stem + "store")) + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity for _, _, identity in servers]}) + for peer, _, identity in servers: + _openapi_server_lists_and_calls_only_its_own_tools(gateway, key, peer, identity) + aggregate: Final = McpCaller(gateway, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + assert sorted(aggregate.tools) == sorted( + f"{prefix}-{tool}" for prefix, tool in itertools.product((stem, stem + "store"), ("getpet", "createpet")) + ), aggregate.tools + assert all(peer.drain() == () for peer, _, _ in servers), "listing must not reach any OpenAPI upstream" + + +def test_config_declared_openapi_server_with_a_space_in_its_name_lists_its_tools( + gateway: Gateway, tmp_path: Path +) -> None: + with openapi_peer() as peer: + config: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text()) + name: Final = "pet store " + uuid.uuid4().hex[:8] + config["mcp_servers"] = {name: peer.registration()} + path: Final = tmp_path / "openapi-space.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + identity: Final = next(i for i, s in _servers(candidate).items() if s["server_name"] == name) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + _openapi_server_lists_and_calls_only_its_own_tools(candidate, key, peer, identity) + aggregate: Final = McpCaller(candidate, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + prefix: Final = name.replace(" ", "_") + assert sorted(aggregate.tools) == [f"{prefix}-createpet", f"{prefix}-getpet"], aggregate.tools + + def test_invalid_registrations_are_rejected(gateway: Gateway) -> None: with mcp_peer() as peer, gateway.scenario() as scenario: alias: Final = "mgmt" + uuid.uuid4().hex[:8] diff --git a/tests/integration/mcp/test_oauth_configuration.py b/tests/integration/mcp/test_oauth_configuration.py index 4c46c706054..fe2b1069f04 100644 --- a/tests/integration/mcp/test_oauth_configuration.py +++ b/tests/integration/mcp/test_oauth_configuration.py @@ -1,17 +1,23 @@ import json import queue +import threading import uuid -from urllib.parse import parse_qs, urlsplit -from typing import Final, Literal +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field from pathlib import Path +from typing import Final, Literal +from urllib.parse import parse_qs, urlsplit import pytest - -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, Scenario, eventually from integration._support.database import read_rows from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names from integration._support.process import owned_proxy -from integration._support.wire import Reply, Request, wire_server +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +_Upstream = Callable[[Request], Reply] @pytest.mark.covers("other.mcp.oauth.discovery_cannot_erase_configured_authorization_endpoint") @@ -104,6 +110,129 @@ def test_partial_discovery_and_unrelated_edit_keep_actual_authorization_destinat assert updated.status_code == 202, updated.text +@dataclass(frozen=True, slots=True) +class _Hold: + armed: threading.Event = field(default_factory=threading.Event) + released: threading.Event = field(default_factory=threading.Event) + + +def _idp_upstream(origin: Callable[[], str], moved: threading.Event, hold: _Hold | None = None) -> _Upstream: + def issuer() -> str: + return origin() + ("/idp-after" if moved.is_set() else "/idp-before") + + def respond(request: Request) -> Reply: + if "oauth-authorization-server" in request.target or "openid-configuration" in request.target: + current: Final = issuer() + return Reply( + body=json.dumps( + { + "issuer": current, + "authorization_endpoint": current + "/authorize", + "token_endpoint": current + "/token", + } + ).encode() + ) + if request.target.startswith("/.well-known/oauth-protected-resource"): + body: Final = json.dumps({"resource": origin() + "/mcp", "authorization_servers": [issuer()]}).encode() + if hold is not None and hold.armed.is_set(): + assert hold.released.wait(timeout=15), "the held upstream metadata reply was never released" + return Reply(body=body) + return Reply(status=404, body=b'{"error":"unexpected"}') + + return respond + + +def _register_pass_through(scenario: Scenario, wire: Wire, alias: str) -> str: + return register_mcp(scenario, McpPeer(wire.url + "/mcp", queue.Queue()), alias, auth_type="true_passthrough") + + +def _wire_requests(wire: Wire, seen: list[Request]) -> Callable[[], tuple[Request, ...]]: + def observed() -> tuple[Request, ...]: + seen.extend(wire.drain()) + return tuple(seen) + + return observed + + +def _registration_discovery_settled(requests: tuple[Request, ...]) -> bool: + return any( + "oauth-authorization-server" in item.target or "openid-configuration" in item.target for item in requests + ) + + +def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str, ...]: + response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp") + assert response.status_code == 200, response.text + return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"])) + + +def _eventually_advertises(gateway: Gateway, alias: str, issuer: str) -> None: + eventually( + lambda: gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp"), + lambda response: response.status_code == 200 and response.json()["authorization_servers"] == [issuer], + seconds=40, + ) + + +def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + moved.set() + wire.drain() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + assert any(request.target.startswith("/.well-known/oauth-protected-resource") for request in wire.drain()), ( + "the save must send protected-resource discovery back to the upstream" + ) + + +def test_peer_worker_stops_advertising_the_old_idp_after_a_save_on_another_worker( + gateway: Gateway, peer: Gateway +) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + _eventually_advertises(peer, alias, wire.url + "/idp-before") + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + _eventually_advertises(peer, alias, wire.url + "/idp-after") + + +def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None: + moved: Final = threading.Event() + hold: Final = _Hold() + with ( + wire_server(_idp_upstream(lambda: wire.url, moved, hold)) as wire, + gateway.scenario() as scenario, + ThreadPoolExecutor(max_workers=1) as pool, + ): + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + seen: Final[list[Request]] = [] + observed: Final = _wire_requests(wire, seen) + eventually(observed, _registration_discovery_settled, seconds=10) + settled: Final = len(seen) + hold.armed.set() + stale: Final = pool.submit(_advertised_authorization_servers, gateway, alias) + eventually(observed, lambda requests: len(requests) > settled, seconds=10) + assert seen[settled].target.startswith("/.well-known/oauth-protected-resource"), seen[settled:] + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + hold.released.set() + assert stale.result(timeout=30) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + + @pytest.mark.covers("other.mcp.oauth.same_url_credentials_are_isolated_by_user_and_server") @pytest.mark.parametrize("transition", ("revoke", "expire")) def test_same_url_oauth_credentials_and_revocation_are_isolated_by_user_and_server( 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 f9a0075e530..b1f0b3fa67e 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 @@ -12611,3 +12611,151 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session( proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called() proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_server_drops_cached_upstream_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + server = MCPServer( + server_id="oauth-cache-server", + name="oauth_cache_server", + url="http://old-upstream/mcp", + transport=MCPTransport.http, + ) + manager.registry[server.server_id] = server + stale_key: Final = (server.server_id, server.url) + other_key: Final = ("other-server", "http://other/mcp") + discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"}) + discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"}) + try: + await manager.update_server( + LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.name, + url="http://new-upstream/mcp", + transport=MCPTransport.http, + ) + ) + assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) + + +@pytest.mark.asyncio +async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cache(): + 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="stale-write-server", name="stale_write", 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": ["old-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, + ): + in_flight: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + 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) + + +@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 + + 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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 0a7ea012b98..a1dc0e779da 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14064,6 +14064,60 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "second" not in str(second) +def _register_local_tool(name: str, description: str) -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + async def _handler(**kwargs): + return None + + global_mcp_tool_registry.register_tool( + name=name, description=description, input_schema={"type": "object"}, handler=_handler + ) + + +def _openapi_server(name: str) -> MCPServer: + return MCPServer( + server_id=f"{name}-id", name=name, alias=name, transport=MCPTransport.http, url=None, spec_path="/spec.yaml" + ) + + +@pytest.mark.asyncio +async def test_openapi_listing_ignores_overlapping_server_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + _register_local_tool("pet-list", "Local pet tool") + _register_local_tool("petstore-list", "Foreign petstore tool") + try: + prefixed: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=True) + bare: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=False) + finally: + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + + assert [t.name for t in prefixed] == ["pet-list"] + assert [t.name for t in bare] == ["list"] + + +@pytest.mark.asyncio +async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + _register_local_tool("pet_store-list", "Pet store tool") + try: + listed: Final = await manager._get_tools_from_server(server=_openapi_server("pet store"), add_prefix=False) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + + assert [t.name for t in listed] == ["list"] + + @pytest.mark.asyncio async def test_discovery_cache_retries_cancelled_fetches() -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache