Merge pull request #40790 from BerriAI/litellm_fix_mcp_discovery_cache_6585

fix(mcp): cache upstream discovery lists
This commit is contained in:
joshua-berri 2026-09-11 20:16:36 -07:00 committed by GitHub
commit 108cacdd2f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 756 additions and 109 deletions

View file

@ -13,6 +13,7 @@ import json
import sys
import threading
import time
from collections.abc import Callable
from typing import TYPE_CHECKING, Any, Final
if TYPE_CHECKING:
@ -34,6 +35,7 @@ class InMemoryCache(BaseCache):
default_ttl: int
| None = 600, # default ttl is 10 minutes. At maximum litellm rate limiting logic requires objects to be in memory for 1 minute
max_size_per_item: int | None = 1024, # 1MB = 1024KB
clock: Callable[[], float] | None = None,
):
"""
max_size_in_memory [int]: Maximum number of items in cache. done to prevent memory leaks. Use 200 items as a default
@ -49,6 +51,7 @@ class InMemoryCache(BaseCache):
self.ttl_dict: dict = {}
self.expiration_heap: list[tuple[float, str]] = []
self._increment_lock = threading.Lock()
self._clock = clock if clock is not None else lambda: time.time()
def check_value_size(self, value: Any):
"""
@ -91,7 +94,7 @@ class InMemoryCache(BaseCache):
"""
Check if a specific key is expired
"""
return key in self.ttl_dict and time.time() > self.ttl_dict[key]
return key in self.ttl_dict and self._clock() > self.ttl_dict[key]
def _remove_key(self, key: str) -> None:
"""
@ -113,7 +116,7 @@ class InMemoryCache(BaseCache):
- 3. the size of in-memory cache is bounded
"""
current_time: Final = time.time()
current_time: Final = self._clock()
# Step 1: Remove expired or outdated items
while self.expiration_heap:
@ -147,7 +150,7 @@ class InMemoryCache(BaseCache):
Check if ttl is set for a key
"""
ttl_time: Final = self.ttl_dict.get(key)
if ttl_time is None or float(ttl_time) < time.time(): # if ttl is not set, allow override
if ttl_time is None or float(ttl_time) < self._clock(): # if ttl is not set, allow override
return True
else:
return False
@ -167,10 +170,10 @@ class InMemoryCache(BaseCache):
self.cache_dict[key] = value
if self.allow_ttl_override(key): # if ttl is not set, set it to default ttl
if "ttl" in kwargs and kwargs["ttl"] is not None:
self.ttl_dict[key] = time.time() + float(kwargs["ttl"])
self.ttl_dict[key] = self._clock() + float(kwargs["ttl"])
heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key))
else:
self.ttl_dict[key] = time.time() + self.default_ttl
self.ttl_dict[key] = self._clock() + self.default_ttl
heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key))
async def async_set_cache(self, key, value, **kwargs):

View file

@ -4,6 +4,8 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
import asyncio
import base64
import hashlib
import json
import os
from collections.abc import Awaitable, Callable, Generator
from contextlib import AbstractAsyncContextManager
@ -343,6 +345,22 @@ class MCPClient:
if auth_value:
self.update_auth_value(auth_value)
async def discovery_auth_fingerprint(self) -> str:
request: Final = httpx.Request("POST", self.server_url or "http://localhost/", headers=self._get_auth_headers())
if self._resolved_auth is None:
return self._hash_discovery_auth(request)
flow: Final = self._resolved_auth.async_auth_flow(request)
try:
authenticated: Final = await flow.__anext__()
return self._hash_discovery_auth(authenticated)
finally:
await flow.aclose()
@staticmethod
def _hash_discovery_auth(request: httpx.Request) -> str:
material: Final = json.dumps((str(request.url), tuple(sorted(request.headers.multi_items()))))
return hashlib.sha256(material.encode()).hexdigest()
def _create_transport_context(
self,
) -> tuple[_TransportContext, httpx.AsyncClient | None]:
@ -781,7 +799,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 +829,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 +889,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 +919,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 +938,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 +971,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
@ -44,11 +46,12 @@ from mcp.types import (
ResourceTemplate,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, BaseModel
from pydantic import AnyUrl, BaseModel, TypeAdapter
from typing_extensions import ReadOnly
import litellm
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
MCP_CLIENT_TIMEOUT,
MCP_HEALTH_CHECK_TIMEOUT,
@ -193,7 +196,6 @@ if TYPE_CHECKING:
from mcp.shared.context import RequestContext
from mcp.types import CreateMessageRequestParams
from litellm.caching.caching import InMemoryCache
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.mcp_server.mcp_toolset import MCPToolset
@ -1677,6 +1679,105 @@ 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]
_DISCOVERY_CACHE_LIMIT: Final = 1024
class _DiscoveryCache(Generic[_DiscoveryItem]):
def __init__(
self, ttl: float, clock: Callable[[], float], adapter: TypeAdapter[tuple[_DiscoveryItem, ...]]
) -> None:
self._ttl = ttl
self._adapter = adapter
self._entries = InMemoryCache(max_size_in_memory=_DISCOVERY_CACHE_LIMIT, max_size_per_item=64, clock=clock)
self._pending: dict[
_DiscoveryKey, asyncio.Task[list[_DiscoveryItem]]
] = {} # mutable-ok: constant-time fetch registration
self._waiters: dict[asyncio.Task[list[_DiscoveryItem]], int] = {} # mutable-ok: constant-time waiter accounting
def invalidate(self, server_id: str) -> None:
prefix: Final = f"[{json.dumps(server_id)},"
keys: Final = cast( # cast-ok: private cache contains only JSON string keys
"tuple[str, ...]", tuple(self._entries.cache_dict)
)
for entry_key in keys:
if entry_key.startswith(prefix):
self._entries.delete_cache(entry_key)
for key in tuple(self._pending):
if key[0] == server_id:
self._pending.pop(key)
@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[object] = self._entries.get_cache(json.dumps(key))
if entry is not None:
return self._adapter.validate_python(entry)
pending: Final = self._pending.get(key)
if pending is not None:
return await self._await_fetch(key, pending)
if len(self._pending) >= _DISCOVERY_CACHE_LIMIT:
return tuple(await fetch())
task: Final = asyncio.create_task(self._fetch(key, fetch))
self._pending[key] = task
task.add_done_callback(self._observe_completion)
return await self._await_fetch(key, task)
async def _await_fetch(
self, key: _DiscoveryKey, task: asyncio.Task[list[_DiscoveryItem]]
) -> tuple[_DiscoveryItem, ...]:
self._waiters[task] = self._waiters.get(task, 0) + 1
try:
return tuple(item.model_copy(deep=True) for item in await asyncio.shield(task))
finally:
remaining: Final = self._waiters[task] - 1
if remaining:
self._waiters[task] = remaining
else:
self._waiters.pop(task)
if self._pending.get(key) is task:
self._pending.pop(key)
if not task.done():
task.cancel()
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():
self._entries.set_cache(
json.dumps(key),
self._adapter.dump_json(tuple(items)),
ttl=self._ttl,
)
return items
finally:
if self._pending.get(key) is asyncio.current_task():
self._pending.pop(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 +1894,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 +1904,16 @@ 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, TypeAdapter(tuple[Prompt, ...])
)
self._resource_discovery_cache = _DiscoveryCache[Resource](
discovery_ttl, discovery_clock, TypeAdapter(tuple[Resource, ...])
)
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
)
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 +2641,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 +2843,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 +3220,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 +3257,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 +4491,41 @@ 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,
credential_fingerprint: str | None = 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, credential_fingerprint),
sort_keys=True,
separators=(",", ":"),
)
return server.server_id, hashlib.sha256(material.encode()).hexdigest()
async def get_prompts_from_server(
self,
server: MCPServer,
@ -4384,47 +4535,38 @@ 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(
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
extra_headers=headers,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
)
prompts: Final = await client.list_prompts()
async def fetch() -> list[Prompt]:
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 +4578,38 @@ 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(
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
extra_headers=headers,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
)
resources: Final = await client.list_resources()
async def fetch() -> list[Resource]:
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 +4621,38 @@ 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(
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
extra_headers=headers,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
key: Final = self._discovery_key(
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
)
resource_templates: Final = await client.list_resource_templates()
async def fetch() -> list[ResourceTemplate]:
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 +5360,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 +5387,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 +5404,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]:
@ -6001,6 +6141,7 @@ class MCPServerManager:
failure is logged, never raised, because the DB write already succeeded and the TTL remains
the backstop.
"""
self._invalidate_discovery_lists(server_id)
try:
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
@ -6466,6 +6607,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

@ -250,3 +250,27 @@ def test_in_memory_cache_prunes_expired_heap_entries_below_capacity():
assert len(in_memory_cache.cache_dict) == 5
assert len(in_memory_cache.ttl_dict) == 5
assert len(in_memory_cache.expiration_heap) == 5
def test_in_memory_cache_injected_clock_controls_expiry_and_eviction() -> None:
class Clock:
now = 0.0
def __call__(self) -> float:
return self.now
clock = Clock()
cache = InMemoryCache(max_size_in_memory=2, default_ttl=60, clock=clock)
cache.set_cache("first", "original", ttl=10)
clock.now = 9.0
cache.set_cache("second", "survivor")
assert cache.get_cache("first") == "original"
clock.now = 10.001
assert cache.get_cache("first") is None
cache.set_cache("third", "replacement")
assert cache.get_cache("second") == "survivor"
clock.now = 69.001
cache.set_cache("fourth", "new")
assert cache.get_cache("second") is None
assert cache.get_cache("third") == "replacement"
assert cache.get_cache("fourth") == "new"

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
@ -1907,3 +1912,25 @@ def test_client_import_before_proxy_credentials_succeeds_in_fresh_process():
)
assert result.returncode == 0, result.stderr
assert result.stdout.strip() == "MCPServerManager"
@pytest.mark.asyncio
@pytest.mark.parametrize("resolved", (False, True))
async def test_discovery_auth_fingerprint_tracks_effective_credentials(resolved: bool) -> None:
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
def client(token: str) -> MCPClient:
return MCPClient(
server_url="https://example.com/mcp",
auth_type=MCPAuth.api_key,
auth_value=None if resolved else token,
resolved_auth=StaticHeaderAuth(token) if resolved else None,
)
original: Final = await client("private-original-credential").discovery_auth_fingerprint()
repeated: Final = await client("private-original-credential").discovery_auth_fingerprint()
replaced: Final = await client("private-replaced-credential").discovery_auth_fingerprint()
assert original == repeated
assert original != replaced
assert len(original) == 64
assert "private-original-credential" not in original

View file

@ -30,7 +30,7 @@ from mcp.types import (
TextResourceContents,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from pydantic import AnyUrl, TypeAdapter
from litellm.constants import MCP_METADATA_TIMEOUT
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -3708,6 +3708,7 @@ class TestMCPServerManager:
mock_prompt = Prompt(name="hello", description="Say hi")
mock_client = AsyncMock()
mock_client.list_prompts = AsyncMock(return_value=[mock_prompt])
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
with patch.object(
manager,
@ -3779,6 +3780,7 @@ class TestMCPServerManager:
mock_client = AsyncMock()
mock_resources = [Resource(name="file", uri="https://example.com/file")]
mock_client.list_resources = AsyncMock(return_value=mock_resources)
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")]
with (
@ -3788,11 +3790,6 @@ class TestMCPServerManager:
new_callable=AsyncMock,
return_value=mock_client,
) as mock_create_client,
patch.object(
manager,
"_create_prefixed_resources",
return_value=prefixed_resources,
) as mock_prefix,
):
result = await manager.get_resources_from_server(
server=server,
@ -3808,7 +3805,6 @@ class TestMCPServerManager:
assert called_kwargs["mcp_auth_header"] == "auth"
assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "static"}
mock_client.list_resources.assert_awaited_once()
mock_prefix.assert_called_once_with(mock_resources, server, add_prefix=True)
assert result == prefixed_resources
@pytest.mark.asyncio
@ -3832,9 +3828,10 @@ class TestMCPServerManager:
)
]
mock_client.list_resource_templates = AsyncMock(return_value=mock_templates)
prefixed_templates = [
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
expected_templates = [
ResourceTemplate(
name="alias-server-template",
name="template",
uriTemplate="https://example.com/{id}",
)
]
@ -3846,11 +3843,6 @@ class TestMCPServerManager:
new_callable=AsyncMock,
return_value=mock_client,
) as mock_create_client,
patch.object(
manager,
"_create_prefixed_resource_templates",
return_value=prefixed_templates,
) as mock_prefix,
):
result = await manager.get_resource_templates_from_server(
server=server,
@ -3866,10 +3858,10 @@ class TestMCPServerManager:
extra_headers=None,
stdio_env=None,
subject_token=None,
user_api_key_auth=None,
)
mock_client.list_resource_templates.assert_awaited_once()
mock_prefix.assert_called_once_with(mock_templates, server, add_prefix=False)
assert result == prefixed_templates
assert result == expected_templates
@pytest.mark.asyncio
async def test_read_resource_from_server_success(self):
@ -13020,3 +13012,436 @@ 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.001
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(), TypeAdapter(tuple[Prompt, ...]))
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"]
@pytest.mark.asyncio
async def test_discovery_cache_cancels_fetch_when_last_waiter_leaves() -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
entered: Final = asyncio.Event()
stopped: Final = asyncio.Event()
release: Final = asyncio.Event()
async def fetch() -> list[Prompt]:
entered.set()
try:
await release.wait()
return [Prompt(name="result")]
finally:
stopped.set()
tasks: Final = tuple(asyncio.create_task(cache.get(("server", None), fetch)) for _ in range(3))
await asyncio.wait_for(entered.wait(), timeout=5)
for task in tasks:
task.cancel()
outcomes: Final = await asyncio.gather(*tasks, return_exceptions=True)
assert all(isinstance(outcome, asyncio.CancelledError) for outcome in outcomes)
try:
await asyncio.wait_for(stopped.wait(), timeout=1)
finally:
release.set()
@pytest.mark.asyncio
async def test_discovery_cache_bounds_detached_fetches_without_dropping_results() -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
entered: Final[asyncio.Queue[None]] = asyncio.Queue()
release: Final = asyncio.Event()
async def blocked() -> list[Prompt]:
await entered.put(None)
await release.wait()
return [Prompt(name="blocked")]
tasks: Final = tuple(asyncio.create_task(cache.get((str(index), None), blocked)) for index in range(1024))
try:
for _ in tasks:
await asyncio.wait_for(entered.get(), timeout=5)
active_tasks: Final = frozenset(asyncio.all_tasks())
async def overflow() -> list[Prompt]:
assert frozenset(asyncio.all_tasks()) <= active_tasks
return [Prompt(name="overflow")]
result: Final = await cache.get(("overflow", None), overflow)
assert [item.name for item in result] == ["overflow"]
finally:
release.set()
outcomes: Final = await asyncio.gather(*tasks)
assert all(result[0].name == "blocked" for result in outcomes)
@pytest.mark.asyncio
async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> None:
import respx
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import UpstreamCredentialProvider
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject
class CredentialSource(UpstreamCredentialProvider):
def __init__(self) -> None:
super().__init__()
self.token: str | None = "token-a"
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Result[httpx.Auth, CredError]:
if self.token is None:
return Error(CredError.of_unauthorized("Credential revoked"))
return Ok(StaticHeaderAuth("Bearer " + self.token))
source: Final = CredentialSource()
managers: Final = (MCPServerManager(cred_provider=source), MCPServerManager(cred_provider=source))
server: Final = MCPServer(
server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="discovery-client",
authorization_url="https://discovery.example/authorize", token_url="https://discovery.example/token",
)
user: Final = UserAPIKeyAuth(user_id="same-user", api_key="same-key")
upstream: Final = _DiscoveryUpstream()
async def respond(request: httpx.Request) -> httpx.Response:
response: Final = await upstream.respond(request)
if '"prompts/list"' not in request.content.decode():
return response
from mcp.types import JSONRPCMessage, JSONRPCRequest
payload: Final = JSONRPCMessage.model_validate_json(request.content).root
assert isinstance(payload, JSONRPCRequest)
name: Final = {"Bearer token-a": "account-a", "Bearer token-b": "account-b"}[request.headers["authorization"]]
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"prompts": [{"name": name}]}})
with respx.mock(base_url="https://discovery.example") as router:
router.route().mock(side_effect=respond)
for manager in managers:
assert [item.name for item in await manager.get_prompts_from_server(server, user)] == ["discovery-account-a"]
assert upstream.initializes == 2
source.token = "token-b"
for manager in managers:
assert [item.name for item in await manager.get_prompts_from_server(server, user)] == ["discovery-account-b"]
assert upstream.initializes == 4
source.token = None
for manager in managers:
assert await manager.get_prompts_from_server(server, user) == []
assert upstream.initializes == 4
@pytest.mark.asyncio
async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None:
import respx
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
class TokenStore:
def __init__(self) -> None:
self.calls: tuple[tuple[str, str], ...] = ()
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls = (*self.calls, (user_id, server_id))
return OAuthToken(access_token="stored-token")
async def invalidate(self, user_id: str, server_id: str) -> None:
return None
store: Final = TokenStore()
manager: Final = MCPServerManager(per_user_oauth_token_store=store)
server: Final = MCPServer(
server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="discovery-client",
authorization_url="https://discovery.example/authorize", token_url="https://discovery.example/token",
)
user: Final = UserAPIKeyAuth(user_id="requesting-user")
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(server, user)) == 1
assert len(await manager.get_prompts_from_server(server, user)) == 1
assert store.calls == (("requesting-user", "discovery"), ("requesting-user", "discovery"))
assert upstream.initializes == 1
assert ("prompts/list", "Bearer stored-token") in upstream.requests
@pytest.mark.asyncio
async def test_discovery_cache_evicts_results_at_capacity() -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
async def original() -> list[Prompt]:
return [Prompt(name="original")]
async def refetched() -> list[Prompt]:
return [Prompt(name="refetched")]
for index in range(1025):
assert (await cache.get((f"server-{index:04}", None), original))[0].name == "original"
assert (await cache.get(("server-1024", None), refetched))[0].name == "original"
assert (await cache.get(("server-0000", None), refetched))[0].name == "refetched"
@pytest.mark.asyncio
async def test_discovery_cache_invalidation_preserves_other_servers_and_pending_fetches() -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
entered: Final = asyncio.Event()
release: Final = asyncio.Event()
async def original() -> list[Prompt]:
return [Prompt(name="original")]
async def blocked() -> list[Prompt]:
entered.set()
await release.wait()
return [Prompt(name="pending")]
async def refetched() -> list[Prompt]:
return [Prompt(name="refetched")]
assert (await cache.get(("server", None), original))[0].name == "original"
assert (await cache.get(("server-extra", None), original))[0].name == "original"
task: Final = asyncio.create_task(cache.get(("other", None), blocked))
await asyncio.wait_for(entered.wait(), timeout=5)
cache.invalidate("server")
release.set()
assert (await asyncio.wait_for(task, timeout=5))[0].name == "pending"
assert (await cache.get(("other", None), refetched))[0].name == "pending"
assert (await cache.get(("server-extra", None), refetched))[0].name == "original"
assert (await cache.get(("server", None), refetched))[0].name == "refetched"
@pytest.mark.asyncio
@pytest.mark.parametrize("description", ("x" * 96_000, "é" * 40_000), ids=("ascii", "unicode"))
async def test_discovery_cache_returns_oversized_results_without_retaining_them(description: str) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
fetch: Final = AsyncMock(return_value=[Prompt(name="large", description=description)])
for _ in range(2):
result: Final = await cache.get(("server", None), fetch)
assert result[0].description == description
assert fetch.await_count == 2