diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 4058a3d72dd..56c9147e066 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -13,6 +13,7 @@ import json import sys import threading import time +from collections.abc import Callable from typing import TYPE_CHECKING, Any, Final if TYPE_CHECKING: @@ -34,6 +35,7 @@ class InMemoryCache(BaseCache): default_ttl: int | None = 600, # default ttl is 10 minutes. At maximum litellm rate limiting logic requires objects to be in memory for 1 minute max_size_per_item: int | None = 1024, # 1MB = 1024KB + clock: Callable[[], float] | None = None, ): """ max_size_in_memory [int]: Maximum number of items in cache. done to prevent memory leaks. Use 200 items as a default @@ -49,6 +51,7 @@ class InMemoryCache(BaseCache): self.ttl_dict: dict = {} self.expiration_heap: list[tuple[float, str]] = [] self._increment_lock = threading.Lock() + self._clock = clock if clock is not None else lambda: time.time() def check_value_size(self, value: Any): """ @@ -91,7 +94,7 @@ class InMemoryCache(BaseCache): """ Check if a specific key is expired """ - return key in self.ttl_dict and time.time() > self.ttl_dict[key] + return key in self.ttl_dict and self._clock() > self.ttl_dict[key] def _remove_key(self, key: str) -> None: """ @@ -113,7 +116,7 @@ class InMemoryCache(BaseCache): - 3. the size of in-memory cache is bounded """ - current_time: Final = time.time() + current_time: Final = self._clock() # Step 1: Remove expired or outdated items while self.expiration_heap: @@ -147,7 +150,7 @@ class InMemoryCache(BaseCache): Check if ttl is set for a key """ ttl_time: Final = self.ttl_dict.get(key) - if ttl_time is None or float(ttl_time) < time.time(): # if ttl is not set, allow override + if ttl_time is None or float(ttl_time) < self._clock(): # if ttl is not set, allow override return True else: return False @@ -167,10 +170,10 @@ class InMemoryCache(BaseCache): self.cache_dict[key] = value if self.allow_ttl_override(key): # if ttl is not set, set it to default ttl if "ttl" in kwargs and kwargs["ttl"] is not None: - self.ttl_dict[key] = time.time() + float(kwargs["ttl"]) + self.ttl_dict[key] = self._clock() + float(kwargs["ttl"]) heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key)) else: - self.ttl_dict[key] = time.time() + self.default_ttl + self.ttl_dict[key] = self._clock() + self.default_ttl heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key)) async def async_set_cache(self, key, value, **kwargs): diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 6cecc1e6157..ee01a53ecb3 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -4,6 +4,8 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers. import asyncio import base64 +import hashlib +import json import os from collections.abc import Awaitable, Callable, Generator from contextlib import AbstractAsyncContextManager @@ -343,6 +345,22 @@ class MCPClient: if auth_value: self.update_auth_value(auth_value) + async def discovery_auth_fingerprint(self) -> str: + request: Final = httpx.Request("POST", self.server_url or "http://localhost/", headers=self._get_auth_headers()) + if self._resolved_auth is None: + return self._hash_discovery_auth(request) + flow: Final = self._resolved_auth.async_auth_flow(request) + try: + authenticated: Final = await flow.__anext__() + return self._hash_discovery_auth(authenticated) + finally: + await flow.aclose() + + @staticmethod + def _hash_discovery_auth(request: httpx.Request) -> str: + material: Final = json.dumps((str(request.url), tuple(sorted(request.headers.multi_items())))) + return hashlib.sha256(material.encode()).hexdigest() + def _create_transport_context( self, ) -> tuple[_TransportContext, httpx.AsyncClient | None]: @@ -781,7 +799,7 @@ class MCPClient: # Return a default error result instead of raising return self.error_tool_result(e) - async def list_prompts(self) -> list[Prompt]: + async def list_prompts(self, *, raise_on_error: bool = False) -> list[Prompt]: """List available prompts from the server.""" verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio") @@ -811,6 +829,8 @@ class MCPClient: verbose_logger.warning("MCP client list_prompts was cancelled") raise except Exception as e: + if raise_on_error: + raise error_type: Final = type(e).__name__ verbose_logger.error( "MCP client list_prompts failed - Error Type: %s, Error: %s, Server: %s, Transport: %s", @@ -869,7 +889,7 @@ class MCPClient: ) raise - async def list_resources(self) -> list[Resource]: + async def list_resources(self, *, raise_on_error: bool = False) -> list[Resource]: """List available resources from the server.""" verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio") @@ -899,6 +919,8 @@ class MCPClient: verbose_logger.warning("MCP client list_resources was cancelled") raise except Exception as e: + if raise_on_error: + raise error_type: Final = type(e).__name__ verbose_logger.error( "MCP client list_resources failed - Error Type: %s, Error: %s, Server: %s, Transport: %s", @@ -916,7 +938,7 @@ class MCPClient: # Return empty list instead of raising to allow graceful degradation return [] - async def list_resource_templates(self) -> list[ResourceTemplate]: + async def list_resource_templates(self, *, raise_on_error: bool = False) -> list[ResourceTemplate]: """List available resource templates from the server.""" verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio") @@ -949,6 +971,8 @@ class MCPClient: verbose_logger.warning("MCP client list_resource_templates was cancelled") raise except Exception as e: + if raise_on_error: + raise error_type: Final = type(e).__name__ verbose_logger.error( "MCP client list_resource_templates failed - Error Type: %s, Error: %s, Server: %s, Transport: %s", diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index cc7ea0c1ea5..9211d8ab003 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -10,6 +10,7 @@ import asyncio import datetime import hashlib import json +import math import os import re import time @@ -26,8 +27,9 @@ from collections.abc import ( from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache +from itertools import chain from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast +from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse import anyio @@ -44,11 +46,12 @@ from mcp.types import ( ResourceTemplate, ) from mcp.types import Tool as MCPTool -from pydantic import AnyUrl, BaseModel +from pydantic import AnyUrl, BaseModel, TypeAdapter from typing_extensions import ReadOnly import litellm from litellm._logging import verbose_logger +from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import ( MCP_CLIENT_TIMEOUT, MCP_HEALTH_CHECK_TIMEOUT, @@ -193,7 +196,6 @@ if TYPE_CHECKING: from mcp.shared.context import RequestContext from mcp.types import CreateMessageRequestParams - from litellm.caching.caching import InMemoryCache from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.mcp_server.mcp_toolset import MCPToolset @@ -1677,6 +1679,105 @@ def _record_mcp_guardrail_evaluations( verbose_logger.warning("Failed to record MCP guardrail evaluation for logging: %s", e) +_DiscoveryItem = TypeVar("_DiscoveryItem", bound=BaseModel) +_DiscoveryKey: TypeAlias = tuple[str, str | None] +_DISCOVERY_CACHE_LIMIT: Final = 1024 + + +class _DiscoveryCache(Generic[_DiscoveryItem]): + def __init__( + self, ttl: float, clock: Callable[[], float], adapter: TypeAdapter[tuple[_DiscoveryItem, ...]] + ) -> None: + self._ttl = ttl + self._adapter = adapter + self._entries = InMemoryCache(max_size_in_memory=_DISCOVERY_CACHE_LIMIT, max_size_per_item=64, clock=clock) + self._pending: dict[ + _DiscoveryKey, asyncio.Task[list[_DiscoveryItem]] + ] = {} # mutable-ok: constant-time fetch registration + self._waiters: dict[asyncio.Task[list[_DiscoveryItem]], int] = {} # mutable-ok: constant-time waiter accounting + + def invalidate(self, server_id: str) -> None: + prefix: Final = f"[{json.dumps(server_id)}," + keys: Final = cast( # cast-ok: private cache contains only JSON string keys + "tuple[str, ...]", tuple(self._entries.cache_dict) + ) + for entry_key in keys: + if entry_key.startswith(prefix): + self._entries.delete_cache(entry_key) + for key in tuple(self._pending): + if key[0] == server_id: + self._pending.pop(key) + + @staticmethod + def _observe_completion(task: asyncio.Task[list[_DiscoveryItem]]) -> None: + if not task.cancelled(): + task.exception() + + async def get( + self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[list[_DiscoveryItem]]] + ) -> tuple[_DiscoveryItem, ...]: + if self._ttl <= 0: + return tuple(await fetch()) + entry: Final[object] = self._entries.get_cache(json.dumps(key)) + if entry is not None: + return self._adapter.validate_python(entry) + pending: Final = self._pending.get(key) + if pending is not None: + return await self._await_fetch(key, pending) + if len(self._pending) >= _DISCOVERY_CACHE_LIMIT: + return tuple(await fetch()) + task: Final = asyncio.create_task(self._fetch(key, fetch)) + self._pending[key] = task + task.add_done_callback(self._observe_completion) + return await self._await_fetch(key, task) + + async def _await_fetch( + self, key: _DiscoveryKey, task: asyncio.Task[list[_DiscoveryItem]] + ) -> tuple[_DiscoveryItem, ...]: + self._waiters[task] = self._waiters.get(task, 0) + 1 + try: + return tuple(item.model_copy(deep=True) for item in await asyncio.shield(task)) + finally: + remaining: Final = self._waiters[task] - 1 + if remaining: + self._waiters[task] = remaining + else: + self._waiters.pop(task) + if self._pending.get(key) is task: + self._pending.pop(key) + if not task.done(): + task.cancel() + + async def _fetch( + self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[list[_DiscoveryItem]]] + ) -> list[_DiscoveryItem]: + try: + items: Final = await fetch() + if self._pending.get(key) is asyncio.current_task(): + self._entries.set_cache( + json.dumps(key), + self._adapter.dump_json(tuple(items)), + ttl=self._ttl, + ) + return items + finally: + if self._pending.get(key) is asyncio.current_task(): + self._pending.pop(key) + + +def _mcp_discovery_cache_ttl() -> float: + raw: Final = os.environ.get("LITELLM_MCP_DISCOVERY_CACHE_TTL", "60") + try: + ttl: Final = float(raw) + except ValueError: + verbose_logger.warning("Invalid LITELLM_MCP_DISCOVERY_CACHE_TTL; using 60 seconds") + return 60.0 + if not math.isfinite(ttl) or ttl < 0: + verbose_logger.warning("Invalid LITELLM_MCP_DISCOVERY_CACHE_TTL; using 60 seconds") + return 60.0 + return ttl + + class MCPServerManager: _STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$") @@ -1793,6 +1894,7 @@ class MCPServerManager: cred_provider: UpstreamCredentialProvider | None = None, per_user_oauth_token_store: InvalidatableOAuthTokenStore | None = None, per_user_token_cache: MCPPerUserTokenCache | None = None, + discovery_clock: Callable[[], float] = time.monotonic, ): self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore( self.get_mcp_server_by_id @@ -1802,6 +1904,16 @@ class MCPServerManager: oauth_token_store=self._per_user_oauth_token_store, token_exchanger=build_token_exchanger(), ) + discovery_ttl: Final = _mcp_discovery_cache_ttl() + self._prompt_discovery_cache = _DiscoveryCache[Prompt]( + discovery_ttl, discovery_clock, TypeAdapter(tuple[Prompt, ...]) + ) + self._resource_discovery_cache = _DiscoveryCache[Resource]( + discovery_ttl, discovery_clock, TypeAdapter(tuple[Resource, ...]) + ) + self._template_discovery_cache = _DiscoveryCache[ResourceTemplate]( + discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...]) + ) self.registry: dict[str, MCPServer] = {} self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe) self.config_mcp_servers: dict[str, MCPServer] = {} @@ -2529,6 +2641,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) _warn_internal_delegate_pkce_if_applicable(new_server, source="config") _warn_config_id_jag_server_outruns_sso(new_server) + self._invalidate_discovery_lists(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2730,6 +2843,7 @@ class MCPServerManager: global_mcp_tool_registry, ) + self._invalidate_discovery_lists(server.server_id) prefix_root: Final = normalize_server_name(get_server_prefix(server)) if server.spec_path and prefix_root: openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR @@ -3106,6 +3220,7 @@ class MCPServerManager: # env_vars_are_encrypted=False. new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) + self._invalidate_discovery_lists(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3142,6 +3257,7 @@ class MCPServerManager: previous_server=self.registry[mcp_server.server_id], ) self._assign_unique_short_prefix(new_server) + self._invalidate_discovery_lists(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -4375,6 +4491,41 @@ class MCPServerManager: ) raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge) + def _invalidate_discovery_lists(self, server_id: str) -> None: + self._prompt_discovery_cache.invalidate(server_id) + self._resource_discovery_cache.invalidate(server_id) + self._template_discovery_cache.invalidate(server_id) + + def _discovery_key( + self, + server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | dict[str, str] | None, + extra_headers: dict[str, str] | None, + stdio_env: dict[str, str] | None, + subject_token: str | None, + credential_fingerprint: str | None = None, + ) -> _DiscoveryKey: + per_user: Final = ( + server.requires_per_user_auth + or self._references_per_user_env_var(server) + or server.delegate_auth_to_upstream + or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) + ) + if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): + return server.server_id, None + identity: Final = ( + (user_api_key_auth.user_id, user_api_key_auth.api_key) + if per_user and user_api_key_auth is not None + else None + ) + material: Final = json.dumps( + (identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint), + sort_keys=True, + separators=(",", ":"), + ) + return server.server_id, hashlib.sha256(material.encode()).hexdigest() + async def get_prompts_from_server( self, server: MCPServer, @@ -4384,47 +4535,38 @@ class MCPServerManager: add_prefix: bool = True, raw_headers: dict[str, str] | None = None, ) -> list[Prompt]: - """ - Helper method to get prompts from a single MCP server with prefixed names. - - Args: - server (MCPServer): The server to query prompts from - mcp_auth_header: Optional auth header for MCP server - - Returns: - List[Prompt]: List of prompts available on the server with prefixed names - """ - - verbose_logger.debug("Connecting to url: %s", server.url) - verbose_logger.info("get_prompts_from_server for %s...", server.name) - - client = None - try: - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) - + headers: Final = ( + dict( + chain( + extra_headers.items() if extra_headers else (), + server.static_headers.items() if server.static_headers else (), + ) + ) + or None + ) stdio_env: Final = self._build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - - client = await self._create_mcp_client( + client: Final = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, + extra_headers=headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, + ) + credential_fingerprint: Final = await client.discovery_auth_fingerprint() + key: Final = self._discovery_key( + server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint ) - prompts: Final = await client.list_prompts() + async def fetch() -> list[Prompt]: + return await client.list_prompts(raise_on_error=True) - prefixed_or_original_prompts: Final = self._create_prefixed_prompts(prompts, server, add_prefix=add_prefix) - - return prefixed_or_original_prompts - - except Exception as e: - verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, e) + items: Final = await self._prompt_discovery_cache.get(key, fetch) + return self._create_prefixed_prompts(items, server, add_prefix=add_prefix) + except Exception as error: + verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error) return [] async def get_resources_from_server( @@ -4436,38 +4578,38 @@ class MCPServerManager: add_prefix: bool = True, raw_headers: dict[str, str] | None = None, ) -> list[Resource]: - """Fetch available resources from a single MCP server.""" - - verbose_logger.debug("Connecting to url: %s", server.url) - verbose_logger.info("get_resources_from_server for %s...", server.name) - - client = None - try: - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) - + headers: Final = ( + dict( + chain( + extra_headers.items() if extra_headers else (), + server.static_headers.items() if server.static_headers else (), + ) + ) + or None + ) stdio_env: Final = self._build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - - client = await self._create_mcp_client( + client: Final = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, + extra_headers=headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, + ) + credential_fingerprint: Final = await client.discovery_auth_fingerprint() + key: Final = self._discovery_key( + server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint ) - resources: Final = await client.list_resources() + async def fetch() -> list[Resource]: + return await client.list_resources(raise_on_error=True) - prefixed_resources: Final = self._create_prefixed_resources(resources, server, add_prefix=add_prefix) - - return prefixed_resources - - except Exception as e: - verbose_logger.warning("Failed to get resources from server %s: %s", server.name, e) + items: Final = await self._resource_discovery_cache.get(key, fetch) + return self._create_prefixed_resources(items, server, add_prefix=add_prefix) + except Exception as error: + verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error) return [] async def get_resource_templates_from_server( @@ -4479,40 +4621,38 @@ class MCPServerManager: add_prefix: bool = True, raw_headers: dict[str, str] | None = None, ) -> list[ResourceTemplate]: - """Fetch available resource templates from a single MCP server.""" - - verbose_logger.debug("Connecting to url: %s", server.url) - verbose_logger.info("get_resource_templates_from_server for %s...", server.name) - - client = None - try: - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) - + headers: Final = ( + dict( + chain( + extra_headers.items() if extra_headers else (), + server.static_headers.items() if server.static_headers else (), + ) + ) + or None + ) stdio_env: Final = self._build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - - client = await self._create_mcp_client( + client: Final = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, + extra_headers=headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, + ) + credential_fingerprint: Final = await client.discovery_auth_fingerprint() + key: Final = self._discovery_key( + server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint ) - resource_templates: Final = await client.list_resource_templates() + async def fetch() -> list[ResourceTemplate]: + return await client.list_resource_templates(raise_on_error=True) - prefixed_templates: Final = self._create_prefixed_resource_templates( - resource_templates, server, add_prefix=add_prefix - ) - - return prefixed_templates - - except Exception as e: - verbose_logger.warning("Failed to get resource templates from server %s: %s", server.name, e) + items: Final = await self._template_discovery_cache.get(key, fetch) + return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix) + except Exception as error: + verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error) return [] async def read_resource_from_server( @@ -5220,7 +5360,7 @@ class MCPServerManager: return prefixed_tools def _create_prefixed_prompts( - self, prompts: list[Prompt], server: MCPServer, add_prefix: bool = True + self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True ) -> list[Prompt]: """ Create prefixed prompts and update prompt mapping. @@ -5247,7 +5387,7 @@ class MCPServerManager: return prefixed_prompts def _create_prefixed_resources( - self, resources: list[Resource], server: MCPServer, add_prefix: bool = True + self, resources: Sequence[Resource], server: MCPServer, add_prefix: bool = True ) -> list[Resource]: """Prefix resource names and track origin server for read requests.""" @@ -5264,7 +5404,7 @@ class MCPServerManager: def _create_prefixed_resource_templates( self, - resource_templates: list[ResourceTemplate], + resource_templates: Sequence[ResourceTemplate], server: MCPServer, add_prefix: bool = True, ) -> list[ResourceTemplate]: @@ -6001,6 +6141,7 @@ class MCPServerManager: failure is logged, never raised, because the DB write already succeeded and the TTL remains the backstop. """ + self._invalidate_discovery_lists(server_id) try: await self._per_user_oauth_token_store.invalidate(user_id, server_id) except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop @@ -6466,6 +6607,9 @@ class MCPServerManager: for registry_key in dropped_registry_keys: self._invalidate_oauth_discovery_state(previous_registry[registry_key].server_id) + for server_id in previous_registry.keys() | registered_registry.keys(): + if previous_registry.get(server_id) != registered_registry.get(server_id): + self._invalidate_discovery_lists(server_id) self.registry = registered_registry # A discovery task may have published into ``previous_registry`` while # this replacement was being staged. Reconcile every published entry diff --git a/tests/test_litellm/caching/test_in_memory_cache.py b/tests/test_litellm/caching/test_in_memory_cache.py index 85e8308ae91..40ad4f0c6f0 100644 --- a/tests/test_litellm/caching/test_in_memory_cache.py +++ b/tests/test_litellm/caching/test_in_memory_cache.py @@ -250,3 +250,27 @@ def test_in_memory_cache_prunes_expired_heap_entries_below_capacity(): assert len(in_memory_cache.cache_dict) == 5 assert len(in_memory_cache.ttl_dict) == 5 assert len(in_memory_cache.expiration_heap) == 5 + + +def test_in_memory_cache_injected_clock_controls_expiry_and_eviction() -> None: + class Clock: + now = 0.0 + + def __call__(self) -> float: + return self.now + + clock = Clock() + cache = InMemoryCache(max_size_in_memory=2, default_ttl=60, clock=clock) + cache.set_cache("first", "original", ttl=10) + clock.now = 9.0 + cache.set_cache("second", "survivor") + assert cache.get_cache("first") == "original" + clock.now = 10.001 + assert cache.get_cache("first") is None + cache.set_cache("third", "replacement") + assert cache.get_cache("second") == "survivor" + clock.now = 69.001 + cache.set_cache("fourth", "new") + assert cache.get_cache("second") is None + assert cache.get_cache("third") == "replacement" + assert cache.get_cache("fourth") == "new" diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 3283af26cf7..f72316f5d5e 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -1704,8 +1704,9 @@ async def test_empty_http_event_stream_uses_the_existing_request_deadline() -> N "initialize_not_found", ), ) +@pytest.mark.parametrize("raise_on_error", (False, True)) async def test_optional_discovery_capabilities_and_errors( - method: str, outcome: str, caplog: pytest.LogCaptureFixture + method: str, outcome: str, caplog: pytest.LogCaptureFixture, raise_on_error: bool ) -> None: import logging from unittest.mock import Mock @@ -1783,7 +1784,11 @@ async def test_optional_discovery_capabilities_and_errors( "resources/list": client.list_resources, "resources/templates/list": client.list_resource_templates, }[method] - result: Final = await operation() + if raise_on_error and outcome in ("internal_error", "unauthorized", "timeout", "initialize_not_found"): + with pytest.raises((McpError, httpx.HTTPError)): + await operation(raise_on_error=True) + return + result: Final = await operation(raise_on_error=raise_on_error) requests: Final = tuple( JSONRPCMessage.model_validate_json(call.args[0].content).root @@ -1907,3 +1912,25 @@ def test_client_import_before_proxy_credentials_succeeds_in_fresh_process(): ) assert result.returncode == 0, result.stderr assert result.stdout.strip() == "MCPServerManager" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("resolved", (False, True)) +async def test_discovery_auth_fingerprint_tracks_effective_credentials(resolved: bool) -> None: + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth + + def client(token: str) -> MCPClient: + return MCPClient( + server_url="https://example.com/mcp", + auth_type=MCPAuth.api_key, + auth_value=None if resolved else token, + resolved_auth=StaticHeaderAuth(token) if resolved else None, + ) + + original: Final = await client("private-original-credential").discovery_auth_fingerprint() + repeated: Final = await client("private-original-credential").discovery_auth_fingerprint() + replaced: Final = await client("private-replaced-credential").discovery_auth_fingerprint() + assert original == repeated + assert original != replaced + assert len(original) == 64 + assert "private-original-credential" not in original diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index b637586d7ff..d2987c5112e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -30,7 +30,7 @@ from mcp.types import ( TextResourceContents, ) from mcp.types import Tool as MCPTool -from pydantic import AnyUrl +from pydantic import AnyUrl, TypeAdapter from litellm.constants import MCP_METADATA_TIMEOUT from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -3708,6 +3708,7 @@ class TestMCPServerManager: mock_prompt = Prompt(name="hello", description="Say hi") mock_client = AsyncMock() mock_client.list_prompts = AsyncMock(return_value=[mock_prompt]) + mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash") with patch.object( manager, @@ -3779,6 +3780,7 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_resources = [Resource(name="file", uri="https://example.com/file")] mock_client.list_resources = AsyncMock(return_value=mock_resources) + mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash") prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")] with ( @@ -3788,11 +3790,6 @@ class TestMCPServerManager: new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client, - patch.object( - manager, - "_create_prefixed_resources", - return_value=prefixed_resources, - ) as mock_prefix, ): result = await manager.get_resources_from_server( server=server, @@ -3808,7 +3805,6 @@ class TestMCPServerManager: assert called_kwargs["mcp_auth_header"] == "auth" assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "static"} mock_client.list_resources.assert_awaited_once() - mock_prefix.assert_called_once_with(mock_resources, server, add_prefix=True) assert result == prefixed_resources @pytest.mark.asyncio @@ -3832,9 +3828,10 @@ class TestMCPServerManager: ) ] mock_client.list_resource_templates = AsyncMock(return_value=mock_templates) - prefixed_templates = [ + mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash") + expected_templates = [ ResourceTemplate( - name="alias-server-template", + name="template", uriTemplate="https://example.com/{id}", ) ] @@ -3846,11 +3843,6 @@ class TestMCPServerManager: new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client, - patch.object( - manager, - "_create_prefixed_resource_templates", - return_value=prefixed_templates, - ) as mock_prefix, ): result = await manager.get_resource_templates_from_server( server=server, @@ -3866,10 +3858,10 @@ class TestMCPServerManager: extra_headers=None, stdio_env=None, subject_token=None, + user_api_key_auth=None, ) mock_client.list_resource_templates.assert_awaited_once() - mock_prefix.assert_called_once_with(mock_templates, server, add_prefix=False) - assert result == prefixed_templates + assert result == expected_templates @pytest.mark.asyncio async def test_read_resource_from_server_success(self): @@ -13020,3 +13012,436 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon assert cached.status == "healthy" assert len(attempts) == 2 assert route.call_count == 1 + + +class _DiscoveryClock: + def __init__(self) -> None: + self.now = 0.0 + + def __call__(self) -> float: + return self.now + + +class _DiscoveryUpstream: + def __init__(self) -> None: + self.requests: tuple[tuple[str, str], ...] = () + self.outcome = "supported" + self.entered = asyncio.Event() + self.release = asyncio.Event() + self.release.set() + + async def respond(self, request: httpx.Request) -> httpx.Response: + from mcp.types import JSONRPCMessage, JSONRPCRequest + + if request.method == "DELETE": + return httpx.Response(200) + payload: Final = JSONRPCMessage.model_validate_json(request.content).root + if not isinstance(payload, JSONRPCRequest): + return httpx.Response(202) + self.requests = (*self.requests, (payload.method, request.headers.get("authorization", ""))) + if payload.method == "initialize": + return httpx.Response(200, json={ + "jsonrpc": "2.0", "id": payload.id, + "result": {"protocolVersion": "2025-03-26", "serverInfo": {"name": "discovery", "version": "1"}, + "capabilities": {} if self.outcome == "unsupported" else {"prompts": {}, "resources": {}}}, + }) + self.entered.set() + await self.release.wait() + if self.outcome == "failure": + return httpx.Response(503) + if self.outcome == "cancelled": + raise asyncio.CancelledError() + if self.outcome == "rejected": + return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, + "error": {"code": -32601, "message": "Unsupported"}}) + result: Final = { + "prompts/list": {"prompts": [{"name": "example", "description": "original"}]}, + "resources/list": {"resources": [{"name": "example", "uri": "test://example", "description": "original"}]}, + "resources/templates/list": {"resourceTemplates": [{"name": "example", "uriTemplate": "test://{name}", "description": "original"}]}, + "tools/list": {"tools": []}, + }[payload.method] + return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result}) + + @property + def initializes(self) -> int: + return sum(method == "initialize" for method, _auth in self.requests) + + +def _discovery_server() -> MCPServer: + return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ("prompts", "resources", "templates")) +async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None: + import respx + + clock: Final = _DiscoveryClock() + manager: Final = MCPServerManager(discovery_clock=clock) + upstream: Final = _DiscoveryUpstream() + operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server, + "templates": manager.get_resource_templates_from_server}[kind] + server: Final = _discovery_server() + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=upstream.respond) + first: Final = await operation(server, None) + assert len(first) == 1 + assert first[0].name == "discovery-example" + first[0].description = "caller changed it" + second: Final = await operation(server, None, add_prefix=False) + assert second[0].name == "example" + assert second[0].description == "original" + assert upstream.initializes == 1 + clock.now = 59.999 + assert (await operation(server, None))[0].name == "discovery-example" + assert upstream.initializes == 1 + clock.now = 60.001 + assert (await operation(server, None))[0].name == "discovery-example" + assert upstream.initializes == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ("prompts", "resources", "templates")) +@pytest.mark.parametrize("outcome", ("unsupported", "rejected", "failure")) +async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: str) -> None: + import respx + + manager: Final = MCPServerManager() + upstream: Final = _DiscoveryUpstream() + upstream.outcome = outcome + operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server, + "templates": manager.get_resource_templates_from_server}[kind] + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=upstream.respond) + assert await operation(_discovery_server(), None) == [] + assert await operation(_discovery_server(), None) == [] + assert upstream.initializes == (2 if outcome == "failure" else 1) + if outcome == "failure": + upstream.outcome = "supported" + assert (await operation(_discovery_server(), None))[0].name == "discovery-example" + assert upstream.initializes == 3 + + +@pytest.mark.asyncio +async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_auth() -> None: + import respx + + manager: Final = MCPServerManager() + upstream: Final = _DiscoveryUpstream() + server: Final = _discovery_server() + first_user: Final = UserAPIKeyAuth(user_id="first") + second_user: Final = UserAPIKeyAuth(user_id="second") + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=upstream.respond) + for user in (first_user, second_user): + assert len(await manager.get_prompts_from_server(server, user)) == 1 + assert upstream.initializes == 1 + for credential in ("first-secret", "second-secret", "first-secret"): + assert len(await manager.get_prompts_from_server(server, first_user, extra_headers={"Authorization": credential})) == 1 + assert upstream.initializes == 3 + assert {auth for method, auth in upstream.requests if method == "prompts/list"} == {"", "first-secret", "second-secret"} + + +@pytest.mark.asyncio +async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> None: + import respx + + manager: Final = MCPServerManager() + upstream: Final = _DiscoveryUpstream() + upstream.release.clear() + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=upstream.respond) + tasks: Final = tuple(asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None)) for _ in range(10)) + await asyncio.wait_for(upstream.entered.wait(), timeout=5) + tasks[0].cancel() + with pytest.raises(asyncio.CancelledError): + await tasks[0] + upstream.release.set() + results: Final = await asyncio.wait_for(asyncio.gather(*tasks[1:]), timeout=5) + assert all(result[0].name == "discovery-example" for result in results) + assert upstream.initializes == 1 + assert results[0][0] is not results[1][0] + assert (await manager.get_prompts_from_server(_discovery_server(), None))[0].name == "discovery-example" + assert upstream.initializes == 1 + + +@pytest.mark.asyncio +async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old_results() -> None: + import respx + + manager: Final = MCPServerManager() + upstream: Final = _DiscoveryUpstream() + upstream.release.clear() + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=upstream.respond) + task: Final = asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None)) + await asyncio.wait_for(upstream.entered.wait(), timeout=5) + manager._invalidate_discovery_lists("discovery") + upstream.release.set() + assert (await task)[0].name == "discovery-example" + assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1 + assert upstream.initializes == 2 + manager._invalidate_discovery_lists("discovery") + assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1 + assert upstream.initializes == 3 + + +@pytest.mark.asyncio +async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) -> None: + import respx + + monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", "0") + manager: Final = MCPServerManager() + upstream: Final = _DiscoveryUpstream() + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=upstream.respond) + assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1 + assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1 + assert upstream.initializes == 2 + + +@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5))) +def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl + + monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", value) + assert _mcp_discovery_cache_ttl() == expected + + +@pytest.mark.parametrize("auth_type", (MCPAuth.oauth2, MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)) +def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> None: + manager: Final = MCPServerManager() + server: Final = _discovery_server().model_copy(update={"auth_type": auth_type}) + first: Final = manager._discovery_key(server, UserAPIKeyAuth(user_id="first"), None, None, None, None) + second: Final = manager._discovery_key(server, UserAPIKeyAuth(user_id="second"), None, None, None, None) + anonymous: Final = manager._discovery_key(server, None, None, None, None, None) + assert len({first, second, anonymous}) == 3 + assert "first" not in str(first) + assert "second" not in str(second) + + +@pytest.mark.asyncio +async def test_discovery_cache_retries_cancelled_fetches() -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache + + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) + + async def cancelled() -> list[Prompt]: + raise asyncio.CancelledError() + + async def supported() -> list[Prompt]: + return [Prompt(name="recovered")] + + with pytest.raises(asyncio.CancelledError): + await cache.get(("server", None), cancelled) + assert [item.name for item in await cache.get(("server", None), supported)] == ["recovered"] + + +@pytest.mark.asyncio +async def test_discovery_cache_cancels_fetch_when_last_waiter_leaves() -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache + + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) + entered: Final = asyncio.Event() + stopped: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def fetch() -> list[Prompt]: + entered.set() + try: + await release.wait() + return [Prompt(name="result")] + finally: + stopped.set() + + tasks: Final = tuple(asyncio.create_task(cache.get(("server", None), fetch)) for _ in range(3)) + await asyncio.wait_for(entered.wait(), timeout=5) + for task in tasks: + task.cancel() + outcomes: Final = await asyncio.gather(*tasks, return_exceptions=True) + assert all(isinstance(outcome, asyncio.CancelledError) for outcome in outcomes) + try: + await asyncio.wait_for(stopped.wait(), timeout=1) + finally: + release.set() + + +@pytest.mark.asyncio +async def test_discovery_cache_bounds_detached_fetches_without_dropping_results() -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache + + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) + entered: Final[asyncio.Queue[None]] = asyncio.Queue() + release: Final = asyncio.Event() + + async def blocked() -> list[Prompt]: + await entered.put(None) + await release.wait() + return [Prompt(name="blocked")] + + tasks: Final = tuple(asyncio.create_task(cache.get((str(index), None), blocked)) for index in range(1024)) + try: + for _ in tasks: + await asyncio.wait_for(entered.get(), timeout=5) + active_tasks: Final = frozenset(asyncio.all_tasks()) + + async def overflow() -> list[Prompt]: + assert frozenset(asyncio.all_tasks()) <= active_tasks + return [Prompt(name="overflow")] + + result: Final = await cache.get(("overflow", None), overflow) + assert [item.name for item in result] == ["overflow"] + finally: + release.set() + outcomes: Final = await asyncio.gather(*tasks) + assert all(result[0].name == "blocked" for result in outcomes) + + +@pytest.mark.asyncio +async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> None: + import respx + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth + from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import UpstreamCredentialProvider + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result + from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject + + class CredentialSource(UpstreamCredentialProvider): + def __init__(self) -> None: + super().__init__() + self.token: str | None = "token-a" + + async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Result[httpx.Auth, CredError]: + if self.token is None: + return Error(CredError.of_unauthorized("Credential revoked")) + return Ok(StaticHeaderAuth("Bearer " + self.token)) + + source: Final = CredentialSource() + managers: Final = (MCPServerManager(cred_provider=source), MCPServerManager(cred_provider=source)) + server: Final = MCPServer( + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="discovery-client", + authorization_url="https://discovery.example/authorize", token_url="https://discovery.example/token", + ) + user: Final = UserAPIKeyAuth(user_id="same-user", api_key="same-key") + upstream: Final = _DiscoveryUpstream() + + async def respond(request: httpx.Request) -> httpx.Response: + response: Final = await upstream.respond(request) + if '"prompts/list"' not in request.content.decode(): + return response + from mcp.types import JSONRPCMessage, JSONRPCRequest + + payload: Final = JSONRPCMessage.model_validate_json(request.content).root + assert isinstance(payload, JSONRPCRequest) + name: Final = {"Bearer token-a": "account-a", "Bearer token-b": "account-b"}[request.headers["authorization"]] + return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"prompts": [{"name": name}]}}) + + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=respond) + for manager in managers: + assert [item.name for item in await manager.get_prompts_from_server(server, user)] == ["discovery-account-a"] + assert upstream.initializes == 2 + source.token = "token-b" + for manager in managers: + assert [item.name for item in await manager.get_prompts_from_server(server, user)] == ["discovery-account-b"] + assert upstream.initializes == 4 + source.token = None + for manager in managers: + assert await manager.get_prompts_from_server(server, user) == [] + assert upstream.initializes == 4 + + +@pytest.mark.asyncio +async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None: + import respx + from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken + + class TokenStore: + def __init__(self) -> None: + self.calls: tuple[tuple[str, str], ...] = () + + async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: + self.calls = (*self.calls, (user_id, server_id)) + return OAuthToken(access_token="stored-token") + + async def invalidate(self, user_id: str, server_id: str) -> None: + return None + + store: Final = TokenStore() + manager: Final = MCPServerManager(per_user_oauth_token_store=store) + server: Final = MCPServer( + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="discovery-client", + authorization_url="https://discovery.example/authorize", token_url="https://discovery.example/token", + ) + user: Final = UserAPIKeyAuth(user_id="requesting-user") + upstream: Final = _DiscoveryUpstream() + with respx.mock(base_url="https://discovery.example") as router: + router.route().mock(side_effect=upstream.respond) + assert len(await manager.get_prompts_from_server(server, user)) == 1 + assert len(await manager.get_prompts_from_server(server, user)) == 1 + assert store.calls == (("requesting-user", "discovery"), ("requesting-user", "discovery")) + assert upstream.initializes == 1 + assert ("prompts/list", "Bearer stored-token") in upstream.requests + + +@pytest.mark.asyncio +async def test_discovery_cache_evicts_results_at_capacity() -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache + + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) + + async def original() -> list[Prompt]: + return [Prompt(name="original")] + + async def refetched() -> list[Prompt]: + return [Prompt(name="refetched")] + + for index in range(1025): + assert (await cache.get((f"server-{index:04}", None), original))[0].name == "original" + assert (await cache.get(("server-1024", None), refetched))[0].name == "original" + assert (await cache.get(("server-0000", None), refetched))[0].name == "refetched" + + +@pytest.mark.asyncio +async def test_discovery_cache_invalidation_preserves_other_servers_and_pending_fetches() -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache + + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) + entered: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def original() -> list[Prompt]: + return [Prompt(name="original")] + + async def blocked() -> list[Prompt]: + entered.set() + await release.wait() + return [Prompt(name="pending")] + + async def refetched() -> list[Prompt]: + return [Prompt(name="refetched")] + + assert (await cache.get(("server", None), original))[0].name == "original" + assert (await cache.get(("server-extra", None), original))[0].name == "original" + task: Final = asyncio.create_task(cache.get(("other", None), blocked)) + await asyncio.wait_for(entered.wait(), timeout=5) + cache.invalidate("server") + release.set() + assert (await asyncio.wait_for(task, timeout=5))[0].name == "pending" + assert (await cache.get(("other", None), refetched))[0].name == "pending" + assert (await cache.get(("server-extra", None), refetched))[0].name == "original" + assert (await cache.get(("server", None), refetched))[0].name == "refetched" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("description", ("x" * 96_000, "é" * 40_000), ids=("ascii", "unicode")) +async def test_discovery_cache_returns_oversized_results_without_retaining_them(description: str) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache + + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) + fetch: Final = AsyncMock(return_value=[Prompt(name="large", description=description)]) + for _ in range(2): + result: Final = await cache.get(("server", None), fetch) + assert result[0].description == description + assert fetch.await_count == 2