From 9d31de2f20d444d02674c8ec3b8f25a2cfa3257e Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Fri, 11 Sep 2026 18:08:27 -0700 Subject: [PATCH] fix(mcp): bound discovery cache result bytes --- .../mcp_server/mcp_server_manager.py | 34 ++++++++++--------- .../mcp_server/test_mcp_server_manager.py | 25 ++++++++++---- 2 files changed, 37 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 35dfeb9ee74..9211d8ab003 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -46,7 +46,7 @@ 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 @@ -1684,15 +1684,13 @@ _DiscoveryKey: TypeAlias = tuple[str, str | None] _DISCOVERY_CACHE_LIMIT: Final = 1024 -@dataclass(frozen=True, slots=True) -class _DiscoveryEntry(Generic[_DiscoveryItem]): - items: tuple[_DiscoveryItem, ...] - - class _DiscoveryCache(Generic[_DiscoveryItem]): - def __init__(self, ttl: float, clock: Callable[[], float]) -> None: + def __init__( + self, ttl: float, clock: Callable[[], float], adapter: TypeAdapter[tuple[_DiscoveryItem, ...]] + ) -> None: self._ttl = ttl - self._entries = InMemoryCache(max_size_in_memory=_DISCOVERY_CACHE_LIMIT, 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[list[_DiscoveryItem]] ] = {} # mutable-ok: constant-time fetch registration @@ -1720,11 +1718,9 @@ class _DiscoveryCache(Generic[_DiscoveryItem]): ) -> tuple[_DiscoveryItem, ...]: if self._ttl <= 0: return tuple(await fetch()) - entry: Final = cast( # cast-ok: private cache contains only entries for this item type - "_DiscoveryEntry[_DiscoveryItem] | None", self._entries.get_cache(json.dumps(key)) - ) + entry: Final[object] = self._entries.get_cache(json.dumps(key)) if entry is not None: - return tuple(item.model_copy(deep=True) for item in entry.items) + return self._adapter.validate_python(entry) pending: Final = self._pending.get(key) if pending is not None: return await self._await_fetch(key, pending) @@ -1760,7 +1756,7 @@ class _DiscoveryCache(Generic[_DiscoveryItem]): if self._pending.get(key) is asyncio.current_task(): self._entries.set_cache( json.dumps(key), - _DiscoveryEntry(tuple(item.model_copy(deep=True) for item in items)), + self._adapter.dump_json(tuple(items)), ttl=self._ttl, ) return items @@ -1909,9 +1905,15 @@ class MCPServerManager: token_exchanger=build_token_exchanger(), ) discovery_ttl: Final = _mcp_discovery_cache_ttl() - self._prompt_discovery_cache = _DiscoveryCache[Prompt](discovery_ttl, discovery_clock) - self._resource_discovery_cache = _DiscoveryCache[Resource](discovery_ttl, discovery_clock) - self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](discovery_ttl, discovery_clock) + self._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] = {} 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 6630e184a16..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 ( @@ -13224,7 +13224,7 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> 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()) + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) async def cancelled() -> list[Prompt]: raise asyncio.CancelledError() @@ -13241,7 +13241,7 @@ async def test_discovery_cache_retries_cancelled_fetches() -> None: 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()) + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) entered: Final = asyncio.Event() stopped: Final = asyncio.Event() release: Final = asyncio.Event() @@ -13270,7 +13270,7 @@ async def test_discovery_cache_cancels_fetch_when_last_waiter_leaves() -> None: 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()) + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) entered: Final[asyncio.Queue[None]] = asyncio.Queue() release: Final = asyncio.Event() @@ -13389,7 +13389,7 @@ async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None 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()) + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) async def original() -> list[Prompt]: return [Prompt(name="original")] @@ -13407,7 +13407,7 @@ async def test_discovery_cache_evicts_results_at_capacity() -> None: 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()) + cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...])) entered: Final = asyncio.Event() release: Final = asyncio.Event() @@ -13432,3 +13432,16 @@ async def test_discovery_cache_invalidation_preserves_other_servers_and_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