fix(mcp): bound discovery cache result bytes

This commit is contained in:
Joshua Valluru 2026-09-11 18:08:27 -07:00
parent 05d2c316f5
commit 9d31de2f20
2 changed files with 37 additions and 22 deletions

View file

@ -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] = {}

View file

@ -30,7 +30,7 @@ from mcp.types import (
TextResourceContents,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from pydantic import AnyUrl, TypeAdapter
from litellm.constants import MCP_METADATA_TIMEOUT
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -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