fix(mcp): cache upstream discovery lists

This commit is contained in:
Joshua Valluru 2026-09-11 14:20:03 -07:00
parent d51a7af655
commit c038aaf622
5 changed files with 451 additions and 98 deletions

View file

@ -2,5 +2,10 @@
LiteLLM MCP Client is a client that allows you to use MCP tools with LiteLLM.
## Gateway discovery caching
The MCP gateway caches each upstream server's prompt, resource, and resource-template lists for 60 seconds per worker. Set `LITELLM_MCP_DISCOVERY_CACHE_TTL` to a nonnegative number of seconds to change the lifetime, or `0` to disable caching. Invalid values use the 60-second default
Discovery results may remain unchanged until that lifetime expires. Server configuration updates invalidate the affected server's entries. Concurrent requests for the same list share one upstream fetch. Each list cache holds at most 1,024 entries per worker
User-dependent upstream authentication uses separate cache entries. Gateway access checks still run for every request. Successful empty lists and unsupported capabilities are cached; failed requests retain the existing empty-list response and are retried on the next request

View file

@ -781,7 +781,7 @@ class MCPClient:
# Return a default error result instead of raising
return self.error_tool_result(e)
async def list_prompts(self) -> list[Prompt]:
async def list_prompts(self, *, raise_on_error: bool = False) -> list[Prompt]:
"""List available prompts from the server."""
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
@ -811,6 +811,8 @@ class MCPClient:
verbose_logger.warning("MCP client list_prompts was cancelled")
raise
except Exception as e:
if raise_on_error:
raise
error_type: Final = type(e).__name__
verbose_logger.error(
"MCP client list_prompts failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
@ -869,7 +871,7 @@ class MCPClient:
)
raise
async def list_resources(self) -> list[Resource]:
async def list_resources(self, *, raise_on_error: bool = False) -> list[Resource]:
"""List available resources from the server."""
verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio")
@ -899,6 +901,8 @@ class MCPClient:
verbose_logger.warning("MCP client list_resources was cancelled")
raise
except Exception as e:
if raise_on_error:
raise
error_type: Final = type(e).__name__
verbose_logger.error(
"MCP client list_resources failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
@ -916,7 +920,7 @@ class MCPClient:
# Return empty list instead of raising to allow graceful degradation
return []
async def list_resource_templates(self) -> list[ResourceTemplate]:
async def list_resource_templates(self, *, raise_on_error: bool = False) -> list[ResourceTemplate]:
"""List available resource templates from the server."""
verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio")
@ -949,6 +953,8 @@ class MCPClient:
verbose_logger.warning("MCP client list_resource_templates was cancelled")
raise
except Exception as e:
if raise_on_error:
raise
error_type: Final = type(e).__name__
verbose_logger.error(
"MCP client list_resource_templates failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",

View file

@ -10,6 +10,7 @@ import asyncio
import datetime
import hashlib
import json
import math
import os
import re
import time
@ -26,8 +27,9 @@ from collections.abc import (
from contextlib import asynccontextmanager
from dataclasses import dataclass, replace
from functools import lru_cache
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
from urllib.parse import ParseResult, urlparse
import anyio
@ -1677,6 +1679,91 @@ def _record_mcp_guardrail_evaluations(
verbose_logger.warning("Failed to record MCP guardrail evaluation for logging: %s", e)
_DiscoveryItem = TypeVar("_DiscoveryItem", bound=BaseModel)
_DiscoveryKey: TypeAlias = tuple[str, str | None]
@dataclass(frozen=True, slots=True)
class _DiscoveryEntry(Generic[_DiscoveryItem]):
expires_at: float
items: tuple[_DiscoveryItem, ...]
class _DiscoveryCache(Generic[_DiscoveryItem]):
def __init__(self, ttl: float, clock: Callable[[], float]) -> None:
self._ttl = ttl
self._clock = clock
self._entries: Mapping[_DiscoveryKey, _DiscoveryEntry[_DiscoveryItem]] = MappingProxyType({})
self._pending: Mapping[_DiscoveryKey, asyncio.Task[list[_DiscoveryItem]]] = MappingProxyType({})
def invalidate(self, server_id: str) -> None:
self._entries = MappingProxyType({key: entry for key, entry in self._entries.items() if key[0] != server_id})
self._pending = MappingProxyType({key: task for key, task in self._pending.items() if key[0] != server_id})
@staticmethod
def _observe_completion(task: asyncio.Task[list[_DiscoveryItem]]) -> None:
if not task.cancelled():
task.exception()
async def get(
self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[list[_DiscoveryItem]]]
) -> tuple[_DiscoveryItem, ...]:
if self._ttl <= 0:
return tuple(await fetch())
entry: Final = self._entries.get(key)
if entry is not None and entry.expires_at > self._clock():
return tuple(item.model_copy(deep=True) for item in entry.items)
pending: Final = self._pending.get(key)
if pending is not None:
return tuple(item.model_copy(deep=True) for item in await asyncio.shield(pending))
task: Final = asyncio.create_task(self._fetch(key, fetch))
self._pending = MappingProxyType({**self._pending, key: task})
task.add_done_callback(self._observe_completion)
return tuple(item.model_copy(deep=True) for item in await asyncio.shield(task))
async def _fetch(
self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[list[_DiscoveryItem]]]
) -> list[_DiscoveryItem]:
try:
items: Final = await fetch()
if self._pending.get(key) is asyncio.current_task():
now: Final = self._clock()
live_entries: Final = tuple(
(entry_key, entry) for entry_key, entry in self._entries.items() if entry.expires_at > now
)
self._entries = MappingProxyType(
{
entry_key: entry
for entry_key, entry in (
*live_entries[-1023:],
(
key,
_DiscoveryEntry(now + self._ttl, tuple(item.model_copy(deep=True) for item in items)),
),
)
}
)
return items
finally:
if self._pending.get(key) is asyncio.current_task():
self._pending = MappingProxyType(
{entry_key: task for entry_key, task in self._pending.items() if entry_key != key}
)
def _mcp_discovery_cache_ttl() -> float:
raw: Final = os.environ.get("LITELLM_MCP_DISCOVERY_CACHE_TTL", "60")
try:
ttl: Final = float(raw)
except ValueError:
verbose_logger.warning("Invalid LITELLM_MCP_DISCOVERY_CACHE_TTL; using 60 seconds")
return 60.0
if not math.isfinite(ttl) or ttl < 0:
verbose_logger.warning("Invalid LITELLM_MCP_DISCOVERY_CACHE_TTL; using 60 seconds")
return 60.0
return ttl
class MCPServerManager:
_STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$")
@ -1793,6 +1880,7 @@ class MCPServerManager:
cred_provider: UpstreamCredentialProvider | None = None,
per_user_oauth_token_store: InvalidatableOAuthTokenStore | None = None,
per_user_token_cache: MCPPerUserTokenCache | None = None,
discovery_clock: Callable[[], float] = time.monotonic,
):
self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore(
self.get_mcp_server_by_id
@ -1802,6 +1890,10 @@ class MCPServerManager:
oauth_token_store=self._per_user_oauth_token_store,
token_exchanger=build_token_exchanger(),
)
discovery_ttl: Final = _mcp_discovery_cache_ttl()
self._prompt_discovery_cache = _DiscoveryCache[Prompt](discovery_ttl, discovery_clock)
self._resource_discovery_cache = _DiscoveryCache[Resource](discovery_ttl, discovery_clock)
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](discovery_ttl, discovery_clock)
self.registry: dict[str, MCPServer] = {}
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
self.config_mcp_servers: dict[str, MCPServer] = {}
@ -2529,6 +2621,7 @@ class MCPServerManager:
self._assign_unique_short_prefix(new_server)
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
_warn_config_id_jag_server_outruns_sso(new_server)
self._invalidate_discovery_lists(server_id)
self.config_mcp_servers[server_id] = new_server
self._set_oauth_discovery_deferred(
server_id,
@ -2730,6 +2823,7 @@ class MCPServerManager:
global_mcp_tool_registry,
)
self._invalidate_discovery_lists(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
@ -3106,6 +3200,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.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
self.prime_oauth_metadata_discovery(new_server)
@ -3142,6 +3237,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.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
self.prime_oauth_metadata_discovery(new_server)
@ -4375,6 +4471,38 @@ class MCPServerManager:
)
raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge)
def _invalidate_discovery_lists(self, server_id: str) -> None:
self._prompt_discovery_cache.invalidate(server_id)
self._resource_discovery_cache.invalidate(server_id)
self._template_discovery_cache.invalidate(server_id)
def _discovery_key(
self,
server: MCPServer,
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | dict[str, str] | None,
extra_headers: dict[str, str] | None,
stdio_env: dict[str, str] | None,
subject_token: str | None,
) -> _DiscoveryKey:
per_user: Final = (
server.requires_per_user_auth
or self._references_per_user_env_var(server)
or server.delegate_auth_to_upstream
or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)
)
if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token):
return server.server_id, None
identity: Final = (
(user_api_key_auth.user_id, user_api_key_auth.api_key)
if per_user and user_api_key_auth is not None
else None
)
material: Final = json.dumps(
(identity, mcp_auth_header, extra_headers, stdio_env, subject_token), sort_keys=True, separators=(",", ":")
)
return server.server_id, hashlib.sha256(material.encode()).hexdigest()
async def get_prompts_from_server(
self,
server: MCPServer,
@ -4384,47 +4512,36 @@ class MCPServerManager:
add_prefix: bool = True,
raw_headers: dict[str, str] | None = None,
) -> list[Prompt]:
"""
Helper method to get prompts from a single MCP server with prefixed names.
Args:
server (MCPServer): The server to query prompts from
mcp_auth_header: Optional auth header for MCP server
Returns:
List[Prompt]: List of prompts available on the server with prefixed names
"""
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("get_prompts_from_server for %s...", server.name)
client = None
try:
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
headers: Final = (
dict(
chain(
extra_headers.items() if extra_headers else (),
server.static_headers.items() if server.static_headers else (),
)
)
or None
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token
)
prompts: Final = await client.list_prompts()
async def fetch() -> list[Prompt]:
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
stdio_env=stdio_env,
subject_token=subject_token,
)
return await client.list_prompts(raise_on_error=True)
prefixed_or_original_prompts: Final = self._create_prefixed_prompts(prompts, server, add_prefix=add_prefix)
return prefixed_or_original_prompts
except Exception as e:
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, e)
items: Final = await self._prompt_discovery_cache.get(key, fetch)
return self._create_prefixed_prompts(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error)
return []
async def get_resources_from_server(
@ -4436,38 +4553,36 @@ class MCPServerManager:
add_prefix: bool = True,
raw_headers: dict[str, str] | None = None,
) -> list[Resource]:
"""Fetch available resources from a single MCP server."""
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("get_resources_from_server for %s...", server.name)
client = None
try:
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
headers: Final = (
dict(
chain(
extra_headers.items() if extra_headers else (),
server.static_headers.items() if server.static_headers else (),
)
)
or None
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token
)
resources: Final = await client.list_resources()
async def fetch() -> list[Resource]:
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
stdio_env=stdio_env,
subject_token=subject_token,
)
return await client.list_resources(raise_on_error=True)
prefixed_resources: Final = self._create_prefixed_resources(resources, server, add_prefix=add_prefix)
return prefixed_resources
except Exception as e:
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, e)
items: Final = await self._resource_discovery_cache.get(key, fetch)
return self._create_prefixed_resources(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error)
return []
async def get_resource_templates_from_server(
@ -4479,40 +4594,36 @@ class MCPServerManager:
add_prefix: bool = True,
raw_headers: dict[str, str] | None = None,
) -> list[ResourceTemplate]:
"""Fetch available resource templates from a single MCP server."""
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("get_resource_templates_from_server for %s...", server.name)
client = None
try:
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
headers: Final = (
dict(
chain(
extra_headers.items() if extra_headers else (),
server.static_headers.items() if server.static_headers else (),
)
)
or None
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token
)
resource_templates: Final = await client.list_resource_templates()
async def fetch() -> list[ResourceTemplate]:
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
stdio_env=stdio_env,
subject_token=subject_token,
)
return await client.list_resource_templates(raise_on_error=True)
prefixed_templates: Final = self._create_prefixed_resource_templates(
resource_templates, server, add_prefix=add_prefix
)
return prefixed_templates
except Exception as e:
verbose_logger.warning("Failed to get resource templates from server %s: %s", server.name, e)
items: Final = await self._template_discovery_cache.get(key, fetch)
return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error)
return []
async def read_resource_from_server(
@ -5220,7 +5331,7 @@ class MCPServerManager:
return prefixed_tools
def _create_prefixed_prompts(
self, prompts: list[Prompt], server: MCPServer, add_prefix: bool = True
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
) -> list[Prompt]:
"""
Create prefixed prompts and update prompt mapping.
@ -5247,7 +5358,7 @@ class MCPServerManager:
return prefixed_prompts
def _create_prefixed_resources(
self, resources: list[Resource], server: MCPServer, add_prefix: bool = True
self, resources: Sequence[Resource], server: MCPServer, add_prefix: bool = True
) -> list[Resource]:
"""Prefix resource names and track origin server for read requests."""
@ -5264,7 +5375,7 @@ class MCPServerManager:
def _create_prefixed_resource_templates(
self,
resource_templates: list[ResourceTemplate],
resource_templates: Sequence[ResourceTemplate],
server: MCPServer,
add_prefix: bool = True,
) -> list[ResourceTemplate]:
@ -6466,6 +6577,9 @@ class MCPServerManager:
for registry_key in dropped_registry_keys:
self._invalidate_oauth_discovery_state(previous_registry[registry_key].server_id)
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.registry = registered_registry
# A discovery task may have published into ``previous_registry`` while
# this replacement was being staged. Reconcile every published entry

View file

@ -1704,8 +1704,9 @@ async def test_empty_http_event_stream_uses_the_existing_request_deadline() -> N
"initialize_not_found",
),
)
@pytest.mark.parametrize("raise_on_error", (False, True))
async def test_optional_discovery_capabilities_and_errors(
method: str, outcome: str, caplog: pytest.LogCaptureFixture
method: str, outcome: str, caplog: pytest.LogCaptureFixture, raise_on_error: bool
) -> None:
import logging
from unittest.mock import Mock
@ -1783,7 +1784,11 @@ async def test_optional_discovery_capabilities_and_errors(
"resources/list": client.list_resources,
"resources/templates/list": client.list_resource_templates,
}[method]
result: Final = await operation()
if raise_on_error and outcome in ("internal_error", "unauthorized", "timeout", "initialize_not_found"):
with pytest.raises((McpError, httpx.HTTPError)):
await operation(raise_on_error=True)
return
result: Final = await operation(raise_on_error=raise_on_error)
requests: Final = tuple(
JSONRPCMessage.model_validate_json(call.args[0].content).root

View file

@ -13020,3 +13020,226 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon
assert cached.status == "healthy"
assert len(attempts) == 2
assert route.call_count == 1
class _DiscoveryClock:
def __init__(self) -> None:
self.now = 0.0
def __call__(self) -> float:
return self.now
class _DiscoveryUpstream:
def __init__(self) -> None:
self.requests: tuple[tuple[str, str], ...] = ()
self.outcome = "supported"
self.entered = asyncio.Event()
self.release = asyncio.Event()
self.release.set()
async def respond(self, request: httpx.Request) -> httpx.Response:
from mcp.types import JSONRPCMessage, JSONRPCRequest
if request.method == "DELETE":
return httpx.Response(200)
payload: Final = JSONRPCMessage.model_validate_json(request.content).root
if not isinstance(payload, JSONRPCRequest):
return httpx.Response(202)
self.requests = (*self.requests, (payload.method, request.headers.get("authorization", "")))
if payload.method == "initialize":
return httpx.Response(200, json={
"jsonrpc": "2.0", "id": payload.id,
"result": {"protocolVersion": "2025-03-26", "serverInfo": {"name": "discovery", "version": "1"},
"capabilities": {} if self.outcome == "unsupported" else {"prompts": {}, "resources": {}}},
})
self.entered.set()
await self.release.wait()
if self.outcome == "failure":
return httpx.Response(503)
if self.outcome == "cancelled":
raise asyncio.CancelledError()
if self.outcome == "rejected":
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id,
"error": {"code": -32601, "message": "Unsupported"}})
result: Final = {
"prompts/list": {"prompts": [{"name": "example", "description": "original"}]},
"resources/list": {"resources": [{"name": "example", "uri": "test://example", "description": "original"}]},
"resources/templates/list": {"resourceTemplates": [{"name": "example", "uriTemplate": "test://{name}", "description": "original"}]},
"tools/list": {"tools": []},
}[payload.method]
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result})
@property
def initializes(self) -> int:
return sum(method == "initialize" for method, _auth in self.requests)
def _discovery_server() -> MCPServer:
return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None:
import respx
clock: Final = _DiscoveryClock()
manager: Final = MCPServerManager(discovery_clock=clock)
upstream: Final = _DiscoveryUpstream()
operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server,
"templates": manager.get_resource_templates_from_server}[kind]
server: Final = _discovery_server()
with respx.mock(base_url="https://discovery.example") as router:
router.route().mock(side_effect=upstream.respond)
first: Final = await operation(server, None)
assert len(first) == 1
assert first[0].name == "discovery-example"
first[0].description = "caller changed it"
second: Final = await operation(server, None, add_prefix=False)
assert second[0].name == "example"
assert second[0].description == "original"
assert upstream.initializes == 1
clock.now = 59.999
assert (await operation(server, None))[0].name == "discovery-example"
assert upstream.initializes == 1
clock.now = 60.0
assert (await operation(server, None))[0].name == "discovery-example"
assert upstream.initializes == 2
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
@pytest.mark.parametrize("outcome", ("unsupported", "rejected", "failure"))
async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: str) -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
upstream.outcome = outcome
operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server,
"templates": manager.get_resource_templates_from_server}[kind]
with respx.mock(base_url="https://discovery.example") as router:
router.route().mock(side_effect=upstream.respond)
assert await operation(_discovery_server(), None) == []
assert await operation(_discovery_server(), None) == []
assert upstream.initializes == (2 if outcome == "failure" else 1)
if outcome == "failure":
upstream.outcome = "supported"
assert (await operation(_discovery_server(), None))[0].name == "discovery-example"
assert upstream.initializes == 3
@pytest.mark.asyncio
async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_auth() -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
server: Final = _discovery_server()
first_user: Final = UserAPIKeyAuth(user_id="first")
second_user: Final = UserAPIKeyAuth(user_id="second")
with respx.mock(base_url="https://discovery.example") as router:
router.route().mock(side_effect=upstream.respond)
for user in (first_user, second_user):
assert len(await manager.get_prompts_from_server(server, user)) == 1
assert upstream.initializes == 1
for credential in ("first-secret", "second-secret", "first-secret"):
assert len(await manager.get_prompts_from_server(server, first_user, extra_headers={"Authorization": credential})) == 1
assert upstream.initializes == 3
assert {auth for method, auth in upstream.requests if method == "prompts/list"} == {"", "first-secret", "second-secret"}
@pytest.mark.asyncio
async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
upstream.release.clear()
with respx.mock(base_url="https://discovery.example") as router:
router.route().mock(side_effect=upstream.respond)
tasks: Final = tuple(asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None)) for _ in range(10))
await asyncio.wait_for(upstream.entered.wait(), timeout=5)
tasks[0].cancel()
with pytest.raises(asyncio.CancelledError):
await tasks[0]
upstream.release.set()
results: Final = await asyncio.wait_for(asyncio.gather(*tasks[1:]), timeout=5)
assert all(result[0].name == "discovery-example" for result in results)
assert upstream.initializes == 1
assert results[0][0] is not results[1][0]
assert (await manager.get_prompts_from_server(_discovery_server(), None))[0].name == "discovery-example"
assert upstream.initializes == 1
@pytest.mark.asyncio
async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old_results() -> None:
import respx
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
upstream.release.clear()
with respx.mock(base_url="https://discovery.example") as router:
router.route().mock(side_effect=upstream.respond)
task: Final = asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None))
await asyncio.wait_for(upstream.entered.wait(), timeout=5)
manager._invalidate_discovery_lists("discovery")
upstream.release.set()
assert (await task)[0].name == "discovery-example"
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
assert upstream.initializes == 2
manager._invalidate_discovery_lists("discovery")
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
assert upstream.initializes == 3
@pytest.mark.asyncio
async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) -> None:
import respx
monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", "0")
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
with respx.mock(base_url="https://discovery.example") as router:
router.route().mock(side_effect=upstream.respond)
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
assert upstream.initializes == 2
@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)))
def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl
monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", value)
assert _mcp_discovery_cache_ttl() == expected
@pytest.mark.parametrize("auth_type", (MCPAuth.oauth2, MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag))
def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> None:
manager: Final = MCPServerManager()
server: Final = _discovery_server().model_copy(update={"auth_type": auth_type})
first: Final = manager._discovery_key(server, UserAPIKeyAuth(user_id="first"), None, None, None, None)
second: Final = manager._discovery_key(server, UserAPIKeyAuth(user_id="second"), None, None, None, None)
anonymous: Final = manager._discovery_key(server, None, None, None, None, None)
assert len({first, second, anonymous}) == 3
assert "first" not in str(first)
assert "second" not in str(second)
@pytest.mark.asyncio
async def test_discovery_cache_retries_cancelled_fetches() -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock())
async def cancelled() -> list[Prompt]:
raise asyncio.CancelledError()
async def supported() -> list[Prompt]:
return [Prompt(name="recovered")]
with pytest.raises(asyncio.CancelledError):
await cache.get(("server", None), cancelled)
assert [item.name for item in await cache.get(("server", None), supported)] == ["recovered"]