From 6498bef8db6704cabd6b63a8f9bafe70622e1dc3 Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Wed, 7 Oct 2026 17:19:24 -0700 Subject: [PATCH] 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> --- litellm/experimental_mcp_client/client.py | 77 ++++-- .../proxy/_experimental/mcp_server/catalog.py | 117 ++++++++- .../mcp_server/faults/list_outcomes.py | 1 + .../mcp_server/mcp_server_manager.py | 185 +++++--------- .../mcp_server/result_conversion.py | 20 +- tests/integration/_support/mcp.py | 13 +- tests/integration/mcp/test_pagination.py | 58 ++++- .../test_mcp_client.py | 96 ++++++- .../_experimental/mcp_server/test_catalog.py | 86 ++++++- .../mcp_server/test_mcp_server_manager.py | 239 +++++++++++++----- .../mcp_server/test_result_conversion.py | 29 ++- tests/unit/test_internal_context.py | 1 + 12 files changed, 701 insertions(+), 221 deletions(-) diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 8f78253eaca..960b9dde53f 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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.""" diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 9331f017582..294bb6ac80f 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py index 6c74ef68118..aaeb50981b6 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py +++ b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py @@ -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]: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f057c471a26..65bb3a59324 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 [] diff --git a/litellm/proxy/_experimental/mcp_server/result_conversion.py b/litellm/proxy/_experimental/mcp_server/result_conversion.py index 29a33746b1e..61435f94c94 100644 --- a/litellm/proxy/_experimental/mcp_server/result_conversion.py +++ b/litellm/proxy/_experimental/mcp_server/result_conversion.py @@ -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", + ) diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index 5991c35c140..17b1e74124c 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -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 ], diff --git a/tests/integration/mcp/test_pagination.py b/tests/integration/mcp/test_pagination.py index b8b28bd39ce..267da498c92 100644 --- a/tests/integration/mcp/test_pagination.py +++ b/tests/integration/mcp/test_pagination.py @@ -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) diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 6c488d770c3..b06ed9468b0 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_catalog.py b/tests/unit/proxy/_experimental/mcp_server/test_catalog.py index 5c029d36c7d..b383078c0ed 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_catalog.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_catalog.py @@ -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" + diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5a608a77931..9aaba0e9356 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_result_conversion.py b/tests/unit/proxy/_experimental/mcp_server/test_result_conversion.py index d9b5063a811..9926fa131ab 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_result_conversion.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_result_conversion.py @@ -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) diff --git a/tests/unit/test_internal_context.py b/tests/unit/test_internal_context.py index 295d2e51023..665ed4f4a8f 100644 --- a/tests/unit/test_internal_context.py +++ b/tests/unit/test_internal_context.py @@ -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",