refactor(mcp): reuse in-memory discovery storage

This commit is contained in:
Joshua Valluru 2026-09-11 17:53:51 -07:00
parent ceb1f04988
commit 05d2c316f5
4 changed files with 112 additions and 43 deletions

View file

@ -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):

View file

@ -51,6 +51,7 @@ 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,
@ -195,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
@ -1686,21 +1686,29 @@ _DISCOVERY_CACHE_LIMIT: Final = 1024
@dataclass(frozen=True, slots=True)
class _DiscoveryEntry(Generic[_DiscoveryItem]):
expires_at: float
items: tuple[_DiscoveryItem, ...]
class _DiscoveryCache(Generic[_DiscoveryItem]):
def __init__(self, ttl: float, clock: Callable[[], float]) -> None:
self._ttl = ttl
self._clock = clock
self._entries: Mapping[_DiscoveryKey, _DiscoveryEntry[_DiscoveryItem]] = MappingProxyType({})
self._pending: Mapping[_DiscoveryKey, asyncio.Task[list[_DiscoveryItem]]] = MappingProxyType({})
self._waiters: Mapping[asyncio.Task[list[_DiscoveryItem]], int] = MappingProxyType({})
self._entries = InMemoryCache(max_size_in_memory=_DISCOVERY_CACHE_LIMIT, 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:
self._entries = MappingProxyType({key: entry for key, entry in self._entries.items() if key[0] != server_id})
self._pending = MappingProxyType({key: task for key, task in self._pending.items() if key[0] != server_id})
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:
@ -1712,8 +1720,10 @@ class _DiscoveryCache(Generic[_DiscoveryItem]):
) -> tuple[_DiscoveryItem, ...]:
if self._ttl <= 0:
return tuple(await fetch())
entry: Final = self._entries.get(key)
if entry is not None and entry.expires_at > self._clock():
entry: Final = cast( # cast-ok: private cache contains only entries for this item type
"_DiscoveryEntry[_DiscoveryItem] | None", self._entries.get_cache(json.dumps(key))
)
if entry is not None:
return tuple(item.model_copy(deep=True) for item in entry.items)
pending: Final = self._pending.get(key)
if pending is not None:
@ -1721,28 +1731,24 @@ class _DiscoveryCache(Generic[_DiscoveryItem]):
if len(self._pending) >= _DISCOVERY_CACHE_LIMIT:
return tuple(await fetch())
task: Final = asyncio.create_task(self._fetch(key, fetch))
self._pending = MappingProxyType({**self._pending, key: task})
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 = MappingProxyType({**self._waiters, task: self._waiters.get(task, 0) + 1})
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 = MappingProxyType({**self._waiters, task: remaining})
self._waiters[task] = remaining
else:
self._waiters = MappingProxyType(
{pending: count for pending, count in self._waiters.items() if pending is not task}
)
self._waiters.pop(task)
if self._pending.get(key) is task:
self._pending = MappingProxyType(
{entry_key: pending for entry_key, pending in self._pending.items() if entry_key != key}
)
self._pending.pop(key)
if not task.done():
task.cancel()
@ -1752,28 +1758,15 @@ class _DiscoveryCache(Generic[_DiscoveryItem]):
try:
items: Final = await fetch()
if self._pending.get(key) is asyncio.current_task():
now: Final = self._clock()
live_entries: Final = tuple(
(entry_key, entry) for entry_key, entry in self._entries.items() if entry.expires_at > now
)
self._entries = MappingProxyType(
{
entry_key: entry
for entry_key, entry in (
*live_entries[-(_DISCOVERY_CACHE_LIMIT - 1) :],
(
key,
_DiscoveryEntry(now + self._ttl, tuple(item.model_copy(deep=True) for item in items)),
),
)
}
self._entries.set_cache(
json.dumps(key),
_DiscoveryEntry(tuple(item.model_copy(deep=True) for item in items)),
ttl=self._ttl,
)
return items
finally:
if self._pending.get(key) is asyncio.current_task():
self._pending = MappingProxyType(
{entry_key: task for entry_key, task in self._pending.items() if entry_key != key}
)
self._pending.pop(key)
def _mcp_discovery_cache_ttl() -> float:

View file

@ -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"

View file

@ -13095,7 +13095,7 @@ async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None
clock.now = 59.999
assert (await operation(server, None))[0].name == "discovery-example"
assert upstream.initializes == 1
clock.now = 60.0
clock.now = 60.001
assert (await operation(server, None))[0].name == "discovery-example"
assert upstream.initializes == 2
@ -13383,3 +13383,52 @@ async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None
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())
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())
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"