fix(mcp): honor scoped cache freshness (#45165)

* fix(mcp): honor upstream freshness and caller scope in discovery caches

* test(mcp): type scoped freshness regression helpers

* fix(mcp): age discovery freshness through cleanup

* fix(mcp): preserve keyless discovery caller isolation

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-10-07 17:19:24 -07:00 • committed by GitHub
parent e9cfba2c17
commit 6498bef8db
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 701 additions and 221 deletions

View file

@ -7,6 +7,7 @@ import base64
import hashlib
import json
import os
import time
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
from contextlib import AbstractAsyncContextManager
from functools import partial
@ -35,6 +36,7 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams]
from mcp.types import (
METHOD_NOT_FOUND,
REQUEST_TIMEOUT,
CacheableResult,
ClientCapabilities,
DiscoverResult,
ElicitationCapability,
@ -55,7 +57,6 @@ from mcp.types import (
ListToolsRequest,
ListToolsResult,
PaginatedRequestParams,
PaginatedResult,
Prompt,
ResourceTemplate,
SamplingCapability,
@ -77,7 +78,11 @@ from litellm.constants import (
from litellm.experimental_mcp_client.tools import list_tools_with_pagination
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
from litellm.proxy._experimental.mcp_server.result_conversion import (
age_freshness,
aggregate_freshness,
error_text_result,
)
from litellm.types.llms.custom_http import VerifyTypes
from litellm.types.mcp import (
MCP_LEGACY_VERSIONS,
@ -178,7 +183,7 @@ def as_mcp_read_timeout(exc: BaseException) -> TimeoutError | None:
TSessionResult = TypeVar("TSessionResult")
_ListPage = TypeVar("_ListPage", bound=PaginatedResult)
_ListPage = TypeVar("_ListPage", ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult)
_ListItem = TypeVar("_ListItem")
@ -1031,12 +1036,21 @@ class MCPClient:
# Return a default error result instead of raising
return self.error_tool_result(e)
async def _run_optional_discovery(self, operation: Callable[[ClientSession], Awaitable[_ListPage]]) -> _ListPage:
async def timed_operation(session: ClientSession) -> tuple[_ListPage, float]:
result: Final = await operation(session)
return result, time.monotonic()
result, received = await self.run_with_session(timed_operation)
return age_freshness(result, time.monotonic() - received)
async def _list_optional_pages(
self,
fetch_page: Callable[[PaginatedRequestParams | None], Awaitable[_ListPage]],
items_of: Callable[[_ListPage], Sequence[_ListItem]],
) -> list[_ListItem]: # mutable-ok: existing list discovery API
) -> tuple[list[_ListItem], CacheableResult]:
items: Final[list[_ListItem]] = [] # mutable-ok: bounded iterative page accumulation
pages: Final[list[tuple[CacheableResult, float]]] = [] # mutable-ok: bounded pagination evidence
cursors: Final[set[str]] = set() # mutable-ok: constant-time detection of cursor cycles
cursor: str | None = None # rebind-ok: iterative traversal avoids recursion at the existing page cap
with anyio.fail_after(max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)):
@ -1047,9 +1061,13 @@ class MCPClient:
if page_index > 0 and error.error.code == METHOD_NOT_FOUND:
raise RuntimeError("MCP list operation became unavailable during pagination") from error
raise
pages.append((page, time.monotonic()))
items.extend(items_of(page))
if not page.next_cursor:
return items
now: Final = time.monotonic()
return items, aggregate_freshness(
tuple(age_freshness(value, now - received) for value, received in pages)
)
if page.next_cursor in cursors:
raise RuntimeError("MCP list pagination repeated a cursor")
cursors.add(page.next_cursor)
@ -1057,6 +1075,9 @@ class MCPClient:
raise RuntimeError(f"MCP list pagination exceeded {MCP_TOOL_LISTING_MAX_PAGES} pages")
async def list_prompts(self, *, raise_on_error: bool = False) -> list[Prompt]:
return (await self.list_prompts_result(raise_on_error=raise_on_error)).prompts
async def list_prompts_result(self, *, raise_on_error: bool = False) -> ListPromptsResult:
"""List available prompts from the server."""
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
@ -1065,11 +1086,10 @@ class MCPClient:
if capabilities is not None and capabilities.prompts is None:
return ListPromptsResult(prompts=[])
try:
return ListPromptsResult(
prompts=await self._list_optional_pages(
lambda params: session.list_prompts(params=params), lambda page: page.prompts
)
items, freshness = await self._list_optional_pages(
lambda params: session.list_prompts(params=params), lambda page: page.prompts
)
return ListPromptsResult(prompts=items, ttl_ms=freshness.ttl_ms, cache_scope=freshness.cache_scope)
except MCPError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
@ -1079,13 +1099,13 @@ class MCPClient:
return ListPromptsResult(prompts=[])
try:
result: Final = await self.run_with_session(_list_prompts_operation)
result: Final = await self._run_optional_discovery(_list_prompts_operation)
prompt_count: Final = len(result.prompts)
prompt_names: Final = [prompt.name for prompt in result.prompts]
verbose_logger.info(
"MCP client listed %s tools from %s: %s", prompt_count, self.server_url or "stdio", prompt_names
)
return result.prompts
return result
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_prompts was cancelled")
raise
@ -1107,7 +1127,7 @@ class MCPClient:
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
return ListPromptsResult(prompts=[])
async def get_prompt(self, get_prompt_request_params: GetPromptRequestParams) -> GetPromptResult:
"""Fetch a prompt definition from the MCP server."""
@ -1151,6 +1171,9 @@ class MCPClient:
raise
async def list_resources(self, *, raise_on_error: bool = False) -> list[Resource]:
return (await self.list_resources_result(raise_on_error=raise_on_error)).resources
async def list_resources_result(self, *, raise_on_error: bool = False) -> ListResourcesResult:
"""List available resources from the server."""
verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio")
@ -1159,11 +1182,10 @@ class MCPClient:
if capabilities is not None and capabilities.resources is None:
return ListResourcesResult(resources=[])
try:
return ListResourcesResult(
resources=await self._list_optional_pages(
lambda params: session.list_resources(params=params), lambda page: page.resources
)
items, freshness = await self._list_optional_pages(
lambda params: session.list_resources(params=params), lambda page: page.resources
)
return ListResourcesResult(resources=items, ttl_ms=freshness.ttl_ms, cache_scope=freshness.cache_scope)
except MCPError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
@ -1173,13 +1195,13 @@ class MCPClient:
return ListResourcesResult(resources=[])
try:
result: Final = await self.run_with_session(_list_resources_operation)
result: Final = await self._run_optional_discovery(_list_resources_operation)
resource_count: Final = len(result.resources)
resource_names: Final = [resource.name for resource in result.resources]
verbose_logger.info(
"MCP client listed %s resources from %s: %s", resource_count, self.server_url or "stdio", resource_names
)
return result.resources
return result
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_resources was cancelled")
raise
@ -1201,9 +1223,12 @@ class MCPClient:
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
return ListResourcesResult(resources=[])
async def list_resource_templates(self, *, raise_on_error: bool = False) -> list[ResourceTemplate]:
return (await self.list_resource_templates_result(raise_on_error=raise_on_error)).resource_templates
async def list_resource_templates_result(self, *, raise_on_error: bool = False) -> ListResourceTemplatesResult:
"""List available resource templates from the server."""
verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio")
@ -1212,11 +1237,11 @@ class MCPClient:
if capabilities is not None and capabilities.resources is None:
return ListResourceTemplatesResult(resource_templates=[])
try:
items, freshness = await self._list_optional_pages(
lambda params: session.list_resource_templates(params=params), lambda page: page.resource_templates
)
return ListResourceTemplatesResult(
resource_templates=await self._list_optional_pages(
lambda params: session.list_resource_templates(params=params),
lambda page: page.resource_templates,
)
resource_templates=items, ttl_ms=freshness.ttl_ms, cache_scope=freshness.cache_scope
)
except MCPError as error:
if error.error.code != METHOD_NOT_FOUND:
@ -1227,7 +1252,7 @@ class MCPClient:
return ListResourceTemplatesResult(resource_templates=[])
try:
result: Final = await self.run_with_session(_list_resource_templates_operation)
result: Final = await self._run_optional_discovery(_list_resource_templates_operation)
resource_template_count: Final = len(result.resource_templates)
resource_template_names: Final = [resource_template.name for resource_template in result.resource_templates]
verbose_logger.info(
@ -1236,7 +1261,7 @@ class MCPClient:
self.server_url or "stdio",
resource_template_names,
)
return result.resource_templates
return result
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_resource_templates was cancelled")
raise
@ -1258,7 +1283,7 @@ class MCPClient:
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
return ListResourceTemplatesResult(resource_templates=[])
async def read_resource(self, url: AnyUrl) -> ReadResourceResult:
"""Fetch resource contents from the MCP server."""

View file

@ -6,6 +6,7 @@ import asyncio
import base64
import hashlib
import json
import time
from collections import UserDict
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, MutableMapping, Sequence
from contextlib import ExitStack, asynccontextmanager
@ -14,11 +15,14 @@ from dataclasses import dataclass, replace
from functools import partial, wraps
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
from typing import TYPE_CHECKING, Final, Generic, ParamSpec, TypeAlias, TypeVar, cast
from mcp.types import CacheableResult
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.proxy._experimental.mcp_server.result_conversion import age_freshness, aggregate_freshness
if TYPE_CHECKING:
from mcp.types import ListToolsResult, PaginatedRequestParams, PaginatedResult
@ -659,6 +663,7 @@ async def paginate_catalog(
result,
)
started: Final = time.monotonic()
tasks: Final = tuple(asyncio.create_task(advance(position)) for position in state.positions)
try:
results: Final = await asyncio.gather(*tasks)
@ -668,7 +673,12 @@ async def paginate_catalog(
await asyncio.gather(*tasks, return_exceptions=True)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
pages: Final = tuple(result for _, result in results if result is not None)
elapsed: Final = time.monotonic() - started
pages: Final = tuple(
age_freshness(result, elapsed) if isinstance(result, CacheableResult) else result
for _, result in results
if result is not None
)
page_outcomes: Final = (
_OUTCOME_VALUES.validate_python((result.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})) for result in pages
)
@ -720,6 +730,10 @@ async def list_tools_page(
)
return ListToolsResult(
tools=list(chain.from_iterable(page.tools for page in pages)),
ttl_ms=0
if any(isinstance(value, dict) and value.get("tag") != "ok" for value in outcomes.values())
else aggregate_freshness(pages).ttl_ms,
cache_scope="private",
next_cursor=next_cursor,
_meta={SERVER_OUTCOMES_META_KEY: dict(outcomes)} if outcomes else None,
)
@ -1036,6 +1050,7 @@ async def aggregate_gateway_tools(
(result.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})
),
next_cursor=result.next_cursor,
ttl_ms=result.ttl_ms,
)
@ -1065,6 +1080,8 @@ async def list_gateway_tools(
return ListToolsResult(
tools=listing.tools,
next_cursor=listing.next_cursor,
ttl_ms=listing.ttl_ms,
cache_scope="private",
_meta={
SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()}
}
@ -1240,10 +1257,19 @@ def combine_optional_catalog(
ListResourceTemplatesResult,
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
outcomes: Final = (meta or {}).get(SERVER_OUTCOMES_META_KEY)
incomplete: Final = isinstance(outcomes, dict) and any(
isinstance(value, dict) and value.get("tag") != "ok" for value in outcomes.values()
)
ttl_ms: Final = 0 if incomplete else aggregate_freshness(pages).ttl_ms
if isinstance(request, ListPromptsRequest):
return ListPromptsResult(
prompts=list(chain.from_iterable(page.prompts for page in pages if isinstance(page, ListPromptsResult))),
next_cursor=next_cursor,
ttl_ms=ttl_ms,
cache_scope="private",
_meta=dict(meta) if meta is not None else None,
)
if isinstance(request, ListResourcesRequest):
@ -1252,6 +1278,8 @@ def combine_optional_catalog(
chain.from_iterable(page.resources for page in pages if isinstance(page, ListResourcesResult))
),
next_cursor=next_cursor,
ttl_ms=ttl_ms,
cache_scope="private",
_meta=dict(meta) if meta is not None else None,
)
return ListResourceTemplatesResult(
@ -1261,5 +1289,90 @@ def combine_optional_catalog(
)
),
next_cursor=next_cursor,
ttl_ms=ttl_ms,
cache_scope="private",
_meta=dict(meta) if meta is not None else None,
)
_DiscoveryPage = TypeVar("_DiscoveryPage", bound=CacheableResult)
_DiscoveryKey: TypeAlias = tuple[str, str | None]
_DISCOVERY_CACHE_LIMIT: Final = 1024
_DISCOVERY_ENTRY: Final = TypeAdapter(tuple[float, bytes])
class _DiscoveryCache(Generic[_DiscoveryPage]):
def __init__(self, ttl: float, clock: Callable[[], float], adapter: TypeAdapter[_DiscoveryPage]) -> None:
self._ttl = ttl
self._clock = clock
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[_DiscoveryPage]] = {}
self._waiters: dict[asyncio.Task[_DiscoveryPage], 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[_DiscoveryPage]) -> None:
if not task.cancelled():
task.exception()
async def get(self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[_DiscoveryPage]]) -> _DiscoveryPage:
if self._ttl <= 0:
return await fetch()
entry: Final[object] = self._entries.get_cache(json.dumps(key))
if entry is not None:
expires_at, payload = _DISCOVERY_ENTRY.validate_python(entry)
remaining: Final = max(0, int((expires_at - self._clock()) * 1000))
if remaining > 0:
return self._adapter.validate_json(payload).model_copy(update={"ttl_ms": remaining})
self._entries.delete_cache(json.dumps(key))
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 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[_DiscoveryPage]) -> _DiscoveryPage:
self._waiters[task] = self._waiters.get(task, 0) + 1
try:
return (await asyncio.shield(task)).model_copy(deep=True)
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[_DiscoveryPage]]) -> _DiscoveryPage:
try:
items: Final = await fetch()
ttl: Final = min(self._ttl, items.ttl_ms / 1000)
if ttl > 0 and self._pending.get(key) is asyncio.current_task():
self._entries.set_cache(
json.dumps(key),
_DISCOVERY_ENTRY.dump_json((self._clock() + ttl, self._adapter.dump_json(items))),
ttl=ttl,
)
return items
finally:
if self._pending.get(key) is asyncio.current_task():
self._pending.pop(key)

View file

@ -66,6 +66,7 @@ class AggregateToolListing(NamedTuple):
tools: list[MCPTool]
outcomes: dict[str, ServerOutcome]
next_cursor: str | None = None
ttl_ms: int = 0
def _iter_upstream_responses(exc: BaseException) -> Iterator[httpx.Response | httpx2.Response]:

View file

@ -29,7 +29,7 @@ from dataclasses import dataclass, replace
from functools import lru_cache
from itertools import chain, groupby
from types import EllipsisType, MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast
from urllib.parse import ParseResult, urlparse
import anyio
@ -45,15 +45,18 @@ from mcp.types import (
GetPromptResult,
InputRequiredResult,
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
ListResourcesResult,
ListResourceTemplatesRequest,
ListResourceTemplatesResult,
ListToolsResult,
PaginatedRequestParams,
Prompt,
ResourceTemplate,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, BaseModel, Field, TypeAdapter
from pydantic import AnyUrl, Field, TypeAdapter
from typing_extensions import ReadOnly, assert_never
import litellm
@ -78,6 +81,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPServerAccess,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.catalog import _configuration_identity, _DiscoveryCache, _DiscoveryKey
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
@ -1744,90 +1748,6 @@ 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]]] = {}
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:
@ -1968,14 +1888,14 @@ class MCPServerManager:
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._prompt_discovery_cache = _DiscoveryCache[ListPromptsResult](
discovery_ttl, discovery_clock, TypeAdapter(ListPromptsResult)
)
self._resource_discovery_cache = _DiscoveryCache[Resource](
discovery_ttl, discovery_clock, TypeAdapter(tuple[Resource, ...])
self._resource_discovery_cache = _DiscoveryCache[ListResourcesResult](
discovery_ttl, discovery_clock, TypeAdapter(ListResourcesResult)
)
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
self._template_discovery_cache = _DiscoveryCache[ListResourceTemplatesResult](
discovery_ttl, discovery_clock, TypeAdapter(ListResourceTemplatesResult)
)
from litellm.proxy._experimental.mcp_server.catalog import CatalogSnapshots
@ -4608,24 +4528,34 @@ class MCPServerManager:
stdio_env: dict[str, str] | None,
subject_token: str | None,
credential_fingerprint: str | None = None,
per_caller: bool = False,
raw_headers: Mapping[str, str] | None = None,
) -> _DiscoveryKey:
per_user: Final = (
per_caller
or 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
user_api_key_auth.model_dump(
include={
"end_user_id",
"user_role",
"object_permission_id",
"team_object_permission_id",
"team_object_permission",
"end_user_object_permission",
},
mode="json",
)
if 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),
(
_configuration_identity(server),
_admission_identity(user_api_key_auth, raw_headers) if user_api_key_auth is not None else None,
identity,
mcp_auth_header,
extra_headers,
stdio_env,
subject_token,
credential_fingerprint,
),
sort_keys=True,
separators=(",", ":"),
)
@ -4741,14 +4671,21 @@ class MCPServerManager:
)
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
server,
user_api_key_auth,
mcp_auth_header,
headers,
stdio_env,
subject_token,
credential_fingerprint,
raw_headers=raw_headers,
)
async def fetch() -> list[Prompt]:
return await client.list_prompts(raise_on_error=True)
async def fetch() -> ListPromptsResult:
return await client.list_prompts_result(raise_on_error=True)
items: Final = await self._prompt_discovery_cache.get(key, fetch)
return self._create_prefixed_prompts(items, server, add_prefix=add_prefix)
return self._create_prefixed_prompts(items.prompts, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error)
return []
@ -4789,14 +4726,21 @@ class MCPServerManager:
)
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
server,
user_api_key_auth,
mcp_auth_header,
headers,
stdio_env,
subject_token,
credential_fingerprint,
raw_headers=raw_headers,
)
async def fetch() -> list[Resource]:
return await client.list_resources(raise_on_error=True)
async def fetch() -> ListResourcesResult:
return await client.list_resources_result(raise_on_error=True)
items: Final = await self._resource_discovery_cache.get(key, fetch)
return self._create_prefixed_resources(items, server, add_prefix=add_prefix)
return self._create_prefixed_resources(items.resources, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error)
return []
@ -4837,14 +4781,21 @@ class MCPServerManager:
)
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
server,
user_api_key_auth,
mcp_auth_header,
headers,
stdio_env,
subject_token,
credential_fingerprint,
raw_headers=raw_headers,
)
async def fetch() -> list[ResourceTemplate]:
return await client.list_resource_templates(raise_on_error=True)
async def fetch() -> ListResourceTemplatesResult:
return await client.list_resource_templates_result(raise_on_error=True)
items: Final = await self._template_discovery_cache.get(key, fetch)
return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix)
return self._create_prefixed_resource_templates(items.resource_templates, 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 []

View file

@ -11,9 +11,11 @@ revisions (``2024-11-05`` .. ``2025-11-25``) and admits any JSON value, plus
from __future__ import annotations
import json
from typing import Final, TypeAlias
import math
from collections.abc import Sequence
from typing import Final, TypeAlias, TypeVar
from mcp.types import CallToolResult, ContentBlock, InputRequiredResult, TextContent, Tool
from mcp.types import CacheableResult, CallToolResult, ContentBlock, InputRequiredResult, TextContent, Tool
from typing_extensions import ReadOnly, TypedDict, assert_never
from litellm.proxy._experimental.mcp_server.tool_outcome import (
@ -118,3 +120,17 @@ def _downgrade_structured_content(result: CallToolResult) -> CallToolResult:
def to_gateway_tool(tool: Tool, name: str) -> Tool:
update: Final[_Renamed] = {"name": name}
return tool.model_copy(deep=True, update=update)
_Cacheable = TypeVar("_Cacheable", bound=CacheableResult)
def age_freshness(result: _Cacheable, elapsed: float) -> _Cacheable:
return result.model_copy(update={"ttl_ms": max(0, result.ttl_ms - math.ceil(max(0.0, elapsed) * 1000))})
def aggregate_freshness(results: Sequence[CacheableResult]) -> CacheableResult:
return CacheableResult(
ttl_ms=min((result.ttl_ms for result in results), default=0),
cache_scope="public" if results and all(result.cache_scope == "public" for result in results) else "private",
)

View file

@ -628,6 +628,7 @@ def tool_calls(observed: tuple[dict[str, object], ...]) -> tuple[dict[str, objec
def paginated_mcp_peer(
*,
page_size: int = 1,
ttl_ms: int = 0,
repeat_cursor: bool = False,
fail_listing: bool = False,
fail_continuation: bool = False,
@ -664,6 +665,8 @@ def paginated_mcp_peer(
async def tools(context, params):
indexes, cursor = window(params)
return ListToolsResult(
ttl_ms=ttl_ms,
cache_scope="public",
tools=[
Tool(
name=f"add{index}",
@ -682,12 +685,18 @@ def paginated_mcp_peer(
async def prompts(context, params):
indexes, cursor = window(params)
return ListPromptsResult(
prompts=[Prompt(name=f"prompt{index}") for index in indexes], next_cursor=cursor, meta=metadata
ttl_ms=ttl_ms,
cache_scope="public",
prompts=[Prompt(name=f"prompt{index}") for index in indexes],
next_cursor=cursor,
meta=metadata,
)
async def resources(context, params):
indexes, cursor = window(params)
return ListResourcesResult(
ttl_ms=ttl_ms,
cache_scope="public",
resources=[Resource(name=f"resource{index}", uri=f"status://item{index}") for index in indexes],
next_cursor=cursor,
meta=metadata,
@ -696,6 +705,8 @@ def paginated_mcp_peer(
async def templates(context, params):
indexes, cursor = window(params)
return ListResourceTemplatesResult(
ttl_ms=ttl_ms,
cache_scope="public",
resource_templates=[
ResourceTemplate(name=f"template{index}", uri_template=f"status{index}://{{item}}") for index in indexes
],

View file

@ -11,7 +11,7 @@ from mcp.client.streamable_http import streamable_http_client
from mcp.types import CallToolRequest, CallToolRequestParams, CallToolResult, PaginatedRequestParams
from integration._support.client import Gateway
from integration._support.mcp import paginated_mcp_peer
from integration._support.mcp import McpPeer, paginated_mcp_peer
from integration._support.process import owned_proxy
from litellm.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPTransport
@ -562,3 +562,59 @@ def test_missing_user_keeps_explicit_key_and_team_grants(
assert not forbidden.ok, forbidden
assert tool_calls(allowed.drain()) == ()
assert private.drain() == ()
def test_aggregate_freshness_matches_real_upstream() -> None:
from mcp.types import (
ListPromptsRequest,
ListResourcesRequest,
ListResourceTemplatesRequest,
ListToolsRequest,
ListToolsResult,
)
from litellm.proxy._experimental.mcp_server.catalog import combine_optional_catalog, list_tools_page
async def exercise(peer: McpPeer) -> None:
client: Final = MCPClient(server_url=peer.url, transport_type=MCPTransport.http, protocol_version="2026-07-28")
for request in (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest()):
page: Final = await client.list_page(request)
assert page.ttl_ms == 9000
result: Final = combine_optional_catalog(request, [page], None, None)
assert result.ttl_ms == 9000
assert result.cache_scope == "private"
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
return await client.list_page(ListToolsRequest())
tools_result: Final = await list_tools_page(
cursor=None, caller_scope="caller", snapshot="snapshot", server_ids=("server",), fetch=fetch, now=100
)
assert 0 < tools_result.ttl_ms <= 9000
assert tools_result.cache_scope == "private"
assert len(tools_result.tools) == 3
with paginated_mcp_peer(page_size=3, ttl_ms=9000) as peer:
asyncio.run(exercise(peer))
@pytest.mark.parametrize("ttl_ms", [0, 9000])
def test_discovery_cache_freshness_and_caller_isolation_over_http(ttl_ms: int) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
async def exercise(peer: McpPeer) -> None:
server: Final = MCPServer(
server_id="pages", name="pages", url=peer.url, transport=MCPTransport.http, protocol_version="2026-07-28"
)
for manager in (MCPServerManager(), MCPServerManager()):
for user in (UserAPIKeyAuth(user_id="one"), UserAPIKeyAuth(user_id="two")):
for _ in range(2):
result: Final = await manager.get_prompts_from_server(server, user)
assert [prompt.name for prompt in result] == ["pages-prompt0", "pages-prompt1", "pages-prompt2"]
with paginated_mcp_peer(page_size=3, ttl_ms=ttl_ms) as peer:
asyncio.run(exercise(peer))
observed: Final = peer.drain()
listings: Final = [item for item in observed if item.get("body", {}).get("method") == "prompts/list"]
assert len(listings) == (8 if ttl_ms == 0 else 4)

View file

@ -6,6 +6,7 @@ import os
import selectors
import sys
from collections.abc import AsyncIterator, Callable
from contextlib import asynccontextmanager
from pathlib import Path
from types import ModuleType
from typing import Final
@ -42,6 +43,7 @@ from litellm.experimental_mcp_client.client import (
MCPClient,
_first_non_cancelled_cause,
_TransportContext,
_TransportStreams,
as_mcp_read_timeout,
strip_auth_scheme,
)
@ -71,14 +73,26 @@ def _initialized(instructions: str | None = None) -> InitializeResult:
class _MockTransportClient(MCPClient):
"""An MCPClient whose streamable-HTTP transport runs on an httpx2 MockTransport."""
def __init__(self, respond, **kwargs):
def __init__(
self,
respond,
*,
http_transport: httpx2.AsyncBaseTransport | None = None,
transport_context: Callable[[str, httpx2.AsyncClient], _TransportContext] | None = None,
**kwargs,
):
super().__init__(**kwargs)
self._respond = respond
self._http_transport = http_transport
self._transport_context = transport_context
def _create_transport_context(self) -> tuple[_TransportContext, httpx2.AsyncClient]:
http_client: Final = self._create_httpx_client_factory(transport=httpx2.MockTransport(self._respond))(
transport: Final = self._http_transport or httpx2.MockTransport(self._respond)
http_client: Final = self._create_httpx_client_factory(transport=transport)(
headers=self._get_auth_headers(), timeout=httpx2.Timeout(self.timeout)
)
if self._transport_context is not None:
return self._transport_context(self.server_url, http_client), http_client
return streamable_http_client(self.server_url, http_client=http_client), http_client
@ -3520,3 +3534,81 @@ async def test_optional_catalog_distinguishes_absent_capability_from_failed_cont
assert result.next_cursor is None
if failure == "unadvertised":
assert method not in methods
@pytest.mark.asyncio
@pytest.mark.parametrize(
"kind,field,entry",
[
("prompts", "prompts", {"name": "example"}),
("resources", "resources", {"name": "example", "uri": "test://example"}),
("resource_templates", "resourceTemplates", {"name": "example", "uriTemplate": "test://{name}"}),
],
)
@pytest.mark.parametrize("ttl", [0, 5000])
@pytest.mark.parametrize("cleanup_phase", ["transport", "http_client"])
@pytest.mark.parametrize("cleanup_seconds", [0, 2, 6])
async def test_optional_discovery_retains_freshness_across_pages(
kind: str, field: str, entry: dict[str, str], ttl: int, cleanup_phase: str, cleanup_seconds: int
) -> None:
def respond(request: httpx2.Request) -> httpx2.Response:
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
assert isinstance(payload, JSONRPCRequest)
if payload.method == "server/discover":
discovery_result: Final = {
"supportedVersions": ["2026-07-28"],
"capabilities": {"prompts": {}, "resources": {}},
"ttlMs": 0,
"cacheScope": "private",
"resultType": "complete",
}
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": discovery_result})
following: Final = bool((payload.params or {}).get("cursor"))
listing_result: Final = {
field: [entry],
"ttlMs": ttl if following else 9000,
"cacheScope": "private" if following else "public",
"resultType": "complete",
**({} if following else {"nextCursor": "next"}),
}
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": listing_result})
clock: Final = _ManualClockLoop()
class CleanupTransport(httpx2.AsyncBaseTransport):
def __init__(self) -> None:
self._transport = httpx2.MockTransport(respond)
async def handle_async_request(self, request: httpx2.Request) -> httpx2.Response:
return await self._transport.handle_async_request(request)
async def aclose(self) -> None:
await self._transport.aclose()
if cleanup_phase == "http_client":
clock.advance(cleanup_seconds)
@asynccontextmanager
async def transport_with_cleanup(url: str, http_client: httpx2.AsyncClient) -> AsyncIterator[_TransportStreams]:
async with streamable_http_client(url, http_client=http_client) as streams:
yield streams
if cleanup_phase == "transport":
clock.advance(cleanup_seconds)
client: Final = _MockTransportClient(
respond,
http_transport=CleanupTransport(),
transport_context=transport_with_cleanup,
server_url="https://example.com/mcp",
protocol_version="2026-07-28",
)
try:
with patch.object(mcp_client_module, "time", Mock(monotonic=clock.time)):
result: Final = await getattr(client, "list_" + kind + "_result")(raise_on_error=True)
finally:
clock.close()
assert len(getattr(result, kind)) == 2
assert result.cache_scope == "private"
assert result.next_cursor is None
assert result.ttl_ms == max(0, ttl - cleanup_seconds * 1000)
assert len(await getattr(client, "list_" + kind)(raise_on_error=True)) == 2

View file

@ -230,14 +230,16 @@ async def test_failed_upstream_remains_visible_when_other_sources_continue(monke
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
if server_id == "a":
return ListToolsResult(tools=[], meta={SERVER_OUTCOMES_META_KEY: {"a": {"tag": "timeout"}}})
return page("b2" if cursor else "b1", None if cursor else "next")
return page("b2" if cursor else "b1", None if cursor else "next").model_copy(update={"ttl_ms": 9000})
first = await listing(fetch)
assert [tool.name for tool in first.tools] == ["b1"]
assert first.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
assert first.ttl_ms == 0
second = await listing(fetch, first.next_cursor)
assert [tool.name for tool in second.tools] == ["b2"]
assert second.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
assert second.ttl_ms == 0
@pytest.mark.asyncio
@ -473,6 +475,7 @@ async def test_optional_gateway_catalog_reports_initial_failure_and_rejects_fail
assert fetch.await_args.args[-1] == "next"
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["tools", "prompts", "resources", "templates"])
async def test_first_catalog_page_keeps_admitted_items_and_reports_rate_limited_servers(
@ -587,3 +590,84 @@ async def test_continuation_rate_limit_raises_without_refetching_completed_serve
charged_server_ids: Final = [call.args[1].server_id for call in limiter.await_args_list]
assert charged_server_ids.count("catalog-a") == 2
assert charged_server_ids.count("catalog-b") == 1
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
@pytest.mark.parametrize("ttls,expected", (((9000, 4000), 4000), ((9000, 0), 0)))
def test_optional_catalog_preserves_conservative_freshness(kind: str, ttls: tuple[int, int], expected: int) -> None:
from mcp.types import (
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
ListResourcesResult,
ListResourceTemplatesRequest,
ListResourceTemplatesResult,
)
request, result_type, field = {
"prompts": (ListPromptsRequest(), ListPromptsResult, "prompts"),
"resources": (ListResourcesRequest(), ListResourcesResult, "resources"),
"templates": (ListResourceTemplatesRequest(), ListResourceTemplatesResult, "resource_templates"),
}[kind]
pages: Final = tuple(result_type(**{field: []}, ttl_ms=ttl, cache_scope="public") for ttl in ttls)
result: Final = catalog.combine_optional_catalog(request, pages, None, None)
assert result.ttl_ms == expected
assert result.cache_scope == "private"
@pytest.mark.asyncio
async def test_tool_catalog_preserves_upstream_freshness() -> None:
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id).model_copy(update={"ttl_ms": 9000, "cache_scope": "public"})
result: Final = await listing(fetch)
assert 0 < result.ttl_ms <= 9000
assert result.cache_scope == "private"
@pytest.mark.asyncio
@pytest.mark.parametrize("upstream_ttl,limit,expected_ttl", [(1000, 60, 1000), (90000, 1, 1000), (0, 60, 0)])
async def test_discovery_cache_uses_upstream_freshness_and_configured_cap(
upstream_ttl: int, limit: float, expected_ttl: int
) -> None:
from unittest.mock import AsyncMock
from mcp.types import ListPromptsResult, Prompt
from pydantic import TypeAdapter
class Clock:
now: float = 0.0
def __call__(self) -> float:
return self.now
clock: Final = Clock()
cache: Final = catalog._DiscoveryCache(limit, clock, TypeAdapter(ListPromptsResult))
fetch: Final = AsyncMock(return_value=ListPromptsResult(prompts=[Prompt(name="fresh")], ttl_ms=upstream_ttl))
first: Final = await cache.get(("server", "caller"), fetch)
assert first.prompts[0].name == "fresh"
clock.now = 0.5 # rebind-ok: advance the injected test clock without sleeping
second: Final = await cache.get(("server", "caller"), fetch)
assert second.prompts[0].name == "fresh"
if expected_ttl:
assert fetch.await_count == 1
assert second.ttl_ms == 500
second.prompts[0].name = "caller edit" # rebind-ok: prove caller mutation cannot alter retained results
else:
assert fetch.await_count == 2
clock.now = 1.0 # rebind-ok: reach the exact expiry boundary without sleeping
assert (await cache.get(("server", "caller"), fetch)).prompts[0].name == "fresh"
assert fetch.await_count == (2 if expected_ttl else 3)
def test_partial_optional_catalog_never_advertises_freshness() -> None:
from mcp.types import ListPromptsRequest, ListPromptsResult
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
result: Final = catalog.combine_optional_catalog(
ListPromptsRequest(),
[ListPromptsResult(prompts=[], ttl_ms=9000)],
None,
{SERVER_OUTCOMES_META_KEY: {"failed-earlier": {"tag": "timeout"}}},
)
assert result.ttl_ms == 0
assert result.cache_scope == "private"

View file

@ -4102,7 +4102,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.list_prompts_result = AsyncMock(return_value=ListPromptsResult(prompts=[mock_prompt]))
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
with patch.object(
@ -4113,7 +4113,7 @@ class TestMCPServerManager:
):
prompts = await manager.get_prompts_from_server(server, user_api_key_auth=None, add_prefix=True)
mock_client.list_prompts.assert_awaited_once()
mock_client.list_prompts_result.assert_awaited_once()
assert len(prompts) == 1
assert prompts[0].name == "alias-server-hello"
@ -4174,7 +4174,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.list_resources_result = AsyncMock(return_value=ListResourcesResult(resources=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")]
@ -4199,7 +4199,7 @@ class TestMCPServerManager:
assert called_kwargs["server"] is server
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_client.list_resources_result.assert_awaited_once()
assert result == prefixed_resources
@pytest.mark.asyncio
@ -4222,7 +4222,9 @@ class TestMCPServerManager:
uriTemplate="https://example.com/{id}",
)
]
mock_client.list_resource_templates = AsyncMock(return_value=mock_templates)
mock_client.list_resource_templates_result = AsyncMock(
return_value=ListResourceTemplatesResult(resource_templates=mock_templates)
)
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
expected_templates = [
ResourceTemplate(
@ -4258,7 +4260,7 @@ class TestMCPServerManager:
raw_headers=None,
client_ip=None,
)
mock_client.list_resource_templates.assert_awaited_once()
mock_client.list_resource_templates_result.assert_awaited_once()
assert result == expected_templates
@pytest.mark.asyncio
@ -7459,10 +7461,10 @@ class TestMCPServerManager:
metadata_key: Final = (new.server_id, new.url)
prompt_fetches = 0
async def fetch_prompts() -> list[Prompt]:
async def fetch_prompts() -> ListPromptsResult:
nonlocal prompt_fetches
prompt_fetches += 1
return [Prompt(name="greet")]
return ListPromptsResult(prompts=[Prompt(name="greet")], ttl_ms=60000)
async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None:
manager.record_listed_tools(
@ -7488,7 +7490,7 @@ class TestMCPServerManager:
assert manager.registry["srv"] is new
assert manager.get_listed_tool(new, "search", caller) is None
prompts = await manager._prompt_discovery_cache.get((new.server_id, None), fetch_prompts)
assert [prompt.name for prompt in prompts] == ["greet"]
assert [prompt.name for prompt in prompts.prompts] == ["greet"]
assert prompt_fetches == 1, "the prompts list filled after the save was published went upstream again"
cached_metadata = discoverable_endpoints._OAUTH_METADATA_CACHE.get(metadata_key)
assert cached_metadata is not None and cached_metadata[1] == {"resource": new.url}
@ -14514,7 +14516,7 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken:
manager: Final = MCPServerManager()
client: Final = AsyncMock()
client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False))
client.list_prompts = AsyncMock(return_value=[])
client.list_prompts_result = AsyncMock(return_value=ListPromptsResult(prompts=[]))
client.read_resource = AsyncMock(return_value=ReadResourceResult(contents=[]))
manager._create_mcp_client = AsyncMock(return_value=client)
return manager
@ -15190,7 +15192,7 @@ class _DiscoveryClock:
from pydantic import TypeAdapter
from mcp.types import JSONRPCMessage
from mcp.types import JSONRPCMessage, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult
_JSONRPC_ADAPTER = TypeAdapter(JSONRPCMessage)
@ -15219,6 +15221,7 @@ class _DiscoveryUpstream:
def __init__(self) -> None:
self.requests: tuple[tuple[str, str], ...] = ()
self.outcome = "supported"
self.ttl_ms = 60000
self.entered = asyncio.Event()
self.release = asyncio.Event()
self.release.set()
@ -15232,6 +15235,21 @@ class _DiscoveryUpstream:
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
self.requests = (*self.requests, (payload.method, request.headers.get("authorization", "")))
if payload.method == "server/discover":
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"supportedVersions": ["2026-07-28"],
"capabilities": {} if self.outcome == "unsupported" else {"prompts": {}, "resources": {}},
"ttlMs": 0,
"cacheScope": "private",
"resultType": "complete",
},
},
)
if payload.method == "initialize":
return httpx2.Response(
200,
@ -15270,16 +15288,33 @@ class _DiscoveryUpstream:
if self.outcome in ("paged", "paged_failure") and not (payload.params or {}).get("cursor")
else {}
)
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {**result, **continuation}})
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
**result,
**continuation,
"ttlMs": self.ttl_ms,
"cacheScope": "private",
"resultType": "complete",
},
},
)
@property
def initializes(self) -> int:
return sum(method == "initialize" for method, _auth in self.requests)
return sum(method in ("initialize", "server/discover") 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
server_id="discovery",
name="discovery",
url="https://discovery.example/mcp",
transport=MCPTransport.http,
protocol_version="2026-07-28",
)
@ -15306,7 +15341,7 @@ async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None
assert second[0].name == "example"
assert second[0].description == "original"
assert upstream.initializes == 1
clock.now = 59.999
clock.now = 59.9
assert (await operation(server, None))[0].name == "discovery-example"
assert upstream.initializes == 1
clock.now = 60.001
@ -15331,7 +15366,7 @@ async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: st
with _mcp_upstream(upstream.respond):
assert await operation(_discovery_server(), None) == []
assert await operation(_discovery_server(), None) == []
assert upstream.initializes == (2 if outcome == "failure" else 1)
assert upstream.initializes == 2
if outcome == "failure":
upstream.outcome = "supported"
assert (await operation(_discovery_server(), None))[0].name == "discovery-example"
@ -15362,7 +15397,7 @@ async def test_discovery_cache_retries_failed_pagination_before_caching_complete
@pytest.mark.asyncio
async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_auth() -> None:
async def test_discovery_cache_isolates_forwarded_credentials_and_static_auth_callers() -> None:
import respx
manager: Final = MCPServerManager()
@ -15373,7 +15408,7 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_
with _mcp_upstream(upstream.respond):
for user in (first_user, second_user):
assert len(await manager.get_prompts_from_server(server, user)) == 1
assert upstream.initializes == 1
assert upstream.initializes == 2
for credential in ("first-secret", "second-secret", "first-secret"):
assert (
len(
@ -15383,7 +15418,7 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_
)
== 1
)
assert upstream.initializes == 3
assert upstream.initializes == 4
assert {auth for method, auth in upstream.requests if method == "prompts/list"} == {
"",
"first-secret",
@ -15391,6 +15426,28 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_
}
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
@pytest.mark.parametrize("header", ("Authorization", "X-LiteLLM-API-Key"))
@pytest.mark.parametrize("identified", (False, True))
async def test_discovery_cache_isolates_keyless_admission_credentials(
kind: str, header: str, identified: bool
) -> None:
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
server: Final = _discovery_server().model_copy(update={"static_headers": {"Authorization": "Bearer upstream"}})
auth: Final = UserAPIKeyAuth(team_id="shared-team", user_id="known-user" if identified else None)
operation: Final = {
"prompts": manager.get_prompts_from_server,
"resources": manager.get_resources_from_server,
"templates": manager.get_resource_templates_from_server,
}[kind]
with _mcp_upstream(upstream.respond):
for credential in ("Bearer first", "Bearer second", "Bearer first"):
assert len(await operation(server, auth, raw_headers={header: credential})) == 1
assert upstream.initializes == (1 if identified else 2)
@pytest.mark.asyncio
async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> None:
import respx
@ -15526,35 +15583,35 @@ async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefi
@pytest.mark.asyncio
async def test_discovery_cache_retries_cancelled_fetches() -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
async def cancelled() -> list[Prompt]:
async def cancelled() -> ListPromptsResult:
raise asyncio.CancelledError()
async def supported() -> list[Prompt]:
return [Prompt(name="recovered")]
async def supported() -> ListPromptsResult:
return ListPromptsResult(prompts=[Prompt(name="recovered")], ttl_ms=60000)
with pytest.raises(asyncio.CancelledError):
await cache.get(("server", None), cancelled)
assert [item.name for item in await cache.get(("server", None), supported)] == ["recovered"]
assert [item.name for item in (await cache.get(("server", None), supported)).prompts] == ["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
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
entered: Final = asyncio.Event()
stopped: Final = asyncio.Event()
release: Final = asyncio.Event()
async def fetch() -> list[Prompt]:
async def fetch() -> ListPromptsResult:
entered.set()
try:
await release.wait()
return [Prompt(name="result")]
return ListPromptsResult(prompts=[Prompt(name="result")], ttl_ms=60000)
finally:
stopped.set()
@ -15572,16 +15629,16 @@ async def test_discovery_cache_cancels_fetch_when_last_waiter_leaves() -> None:
@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
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
entered: Final[asyncio.Queue[None]] = asyncio.Queue()
release: Final = asyncio.Event()
async def blocked() -> list[Prompt]:
async def blocked() -> ListPromptsResult:
await entered.put(None)
await release.wait()
return [Prompt(name="blocked")]
return ListPromptsResult(prompts=[Prompt(name="blocked")], ttl_ms=60000)
tasks: Final = tuple(asyncio.create_task(cache.get((str(index), None), blocked)) for index in range(1024))
try:
@ -15589,16 +15646,16 @@ async def test_discovery_cache_bounds_detached_fetches_without_dropping_results(
await asyncio.wait_for(entered.get(), timeout=5)
active_tasks: Final = frozenset(asyncio.all_tasks())
async def overflow() -> list[Prompt]:
async def overflow() -> ListPromptsResult:
assert frozenset(asyncio.all_tasks()) <= active_tasks
return [Prompt(name="overflow")]
return ListPromptsResult(prompts=[Prompt(name="overflow")], ttl_ms=60000)
result: Final = await cache.get(("overflow", None), overflow)
assert [item.name for item in result] == ["overflow"]
assert [item.name for item in result.prompts] == ["overflow"]
finally:
release.set()
outcomes: Final = await asyncio.gather(*tasks)
assert all(result[0].name == "blocked" for result in outcomes)
assert all(result.prompts[0].name == "blocked" for result in outcomes)
@pytest.mark.asyncio
@ -15628,6 +15685,7 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
protocol_version="2026-07-28",
client_id="discovery-client",
authorization_url="https://discovery.example/authorize",
token_url="https://discovery.example/token",
@ -15644,16 +15702,28 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N
payload: Final = _JSONRPC_ADAPTER.validate_json(request.content)
assert isinstance(payload, JSONRPCRequest)
name: Final = {"Bearer token-a": "account-a", "Bearer token-b": "account-b"}[request.headers["authorization"]]
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"prompts": [{"name": name}]}})
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"prompts": [{"name": name}],
"ttlMs": 60000,
"cacheScope": "private",
"resultType": "complete",
},
},
)
with _mcp_upstream(respond):
for manager in managers:
for manager in (*managers, *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:
for manager in (*managers, *managers):
assert [item.name for item in await manager.get_prompts_from_server(server, user)] == [
"discovery-account-b"
]
@ -15699,69 +15769,71 @@ async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None
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 upstream.initializes == 2
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
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
async def original() -> list[Prompt]:
return [Prompt(name="original")]
async def original() -> ListPromptsResult:
return ListPromptsResult(prompts=[Prompt(name="original")], ttl_ms=60000)
async def refetched() -> list[Prompt]:
return [Prompt(name="refetched")]
async def refetched() -> ListPromptsResult:
return ListPromptsResult(prompts=[Prompt(name="refetched")], ttl_ms=60000)
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"
assert (await cache.get((f"server-{index:04}", None), original)).prompts[0].name == "original"
assert (await cache.get(("server-1024", None), refetched)).prompts[0].name == "original"
assert (await cache.get(("server-0000", None), refetched)).prompts[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
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
entered: Final = asyncio.Event()
release: Final = asyncio.Event()
async def original() -> list[Prompt]:
return [Prompt(name="original")]
async def original() -> ListPromptsResult:
return ListPromptsResult(prompts=[Prompt(name="original")], ttl_ms=60000)
async def blocked() -> list[Prompt]:
async def blocked() -> ListPromptsResult:
entered.set()
await release.wait()
return [Prompt(name="pending")]
return ListPromptsResult(prompts=[Prompt(name="pending")], ttl_ms=60000)
async def refetched() -> list[Prompt]:
return [Prompt(name="refetched")]
async def refetched() -> ListPromptsResult:
return ListPromptsResult(prompts=[Prompt(name="refetched")], ttl_ms=60000)
assert (await cache.get(("server", None), original))[0].name == "original"
assert (await cache.get(("server-extra", None), original))[0].name == "original"
assert (await cache.get(("server", None), original)).prompts[0].name == "original"
assert (await cache.get(("server-extra", None), original)).prompts[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"
assert (await asyncio.wait_for(task, timeout=5)).prompts[0].name == "pending"
assert (await cache.get(("other", None), refetched)).prompts[0].name == "pending"
assert (await cache.get(("server-extra", None), refetched)).prompts[0].name == "original"
assert (await cache.get(("server", None), refetched)).prompts[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
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
fetch: Final = AsyncMock(return_value=[Prompt(name="large", description=description)])
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
fetch: Final = AsyncMock(
return_value=ListPromptsResult(prompts=[Prompt(name="large", description=description)], ttl_ms=60000)
)
for _ in range(2):
result: Final = await cache.get(("server", None), fetch)
assert result[0].description == description
assert result.prompts[0].description == description
assert fetch.await_count == 2
@ -19005,3 +19077,34 @@ async def test_aggregate_publishes_complete_bare_routes_only_after_delivering_a_
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
await listing
assert manager._get_mcp_server_from_tool_name("first") is None
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
async def test_discovery_does_not_retain_unknown_freshness(kind: str) -> None:
manager: Final = MCPServerManager()
upstream: Final = _DiscoveryUpstream()
upstream.ttl_ms = 0
operation: Final = {
"prompts": manager.get_prompts_from_server,
"resources": manager.get_resources_from_server,
"templates": manager.get_resource_templates_from_server,
}[kind]
with _mcp_upstream(upstream.respond):
assert len(await operation(_discovery_server(), None)) == 1
assert len(await operation(_discovery_server(), None)) == 1
assert upstream.initializes == 2
def test_discovery_keys_bind_static_auth_to_caller_and_configuration() -> None:
manager: Final = MCPServerManager()
server: Final = _discovery_server()
first: Final = UserAPIKeyAuth(user_id="first", team_id="one")
second: Final = UserAPIKeyAuth(user_id="second", team_id="two")
updated: Final = server.model_copy(update={"url": "https://replacement.example/mcp"})
keys: Final = (
manager._discovery_key(server, first, None, None, None, None),
manager._discovery_key(server, second, None, None, None, None),
manager._discovery_key(updated, first, None, None, None, None),
)
assert len(set(keys)) == 3

View file

@ -1,5 +1,5 @@
import json
from typing import Final
from typing import Final, Literal
import pytest
from mcp.types import CallToolResult, ImageContent, InputRequiredResult, TextContent, Tool
@ -241,3 +241,30 @@ class TestToGatewayTool:
assert renamed.input_schema == tool.input_schema and renamed.input_schema is not tool.input_schema
assert renamed.meta == {"owner": "x"}
assert renamed.description == "d"
@pytest.mark.parametrize("elapsed,expected", [(0, 1000), (0.0001, 999), (0.5, 500), (1, 0), (2, 0), (-1, 1000)])
def test_freshness_aging_preserves_content_and_scope(elapsed: float, expected: int) -> None:
from mcp.types import ListPromptsResult, Prompt
from litellm.proxy._experimental.mcp_server.result_conversion import age_freshness
result: Final = ListPromptsResult(prompts=[Prompt(name="kept")], ttl_ms=1000, cache_scope="public")
aged: Final = age_freshness(result, elapsed)
assert aged.ttl_ms == expected
assert aged.cache_scope == "public"
assert aged.prompts == result.prompts
assert result.ttl_ms == 1000
@pytest.mark.parametrize(
"scopes,expected", [((), "private"), (("public", "public"), "public"), (("public", "private"), "private")]
)
def test_aggregate_freshness_never_broadens_sharing(
scopes: tuple[Literal["private", "public"], ...], expected: str
) -> None:
from mcp.types import CacheableResult
from litellm.proxy._experimental.mcp_server.result_conversion import aggregate_freshness
result: Final = aggregate_freshness(tuple(CacheableResult(ttl_ms=1000, cache_scope=scope) for scope in scopes))
assert result.cache_scope == expected
assert result.ttl_ms == (1000 if scopes else 0)

View file

@ -59,6 +59,7 @@ _IN_MEMORY_ONLY_CALLERS: Final = frozenset(
"litellm/llms/vertex_ai/vertex_ai_non_gemini.py",
"litellm/llms/watsonx/common_utils.py",
"litellm/proxy/_experimental/mcp_server/byok_credential_cache.py",
"litellm/proxy/_experimental/mcp_server/catalog.py",
"litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py",
"litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py",
"litellm/proxy/_experimental/mcp_server/operations.py",