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 <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 17:25:39 -07:00 • committed by GitHub
parent ffb15f946f
commit c129ea4fc9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 470 additions and 34 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,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

View file

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

View file

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

View file

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

View file

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

View file

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