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/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index ca81e8191f7..35dfeb9ee74 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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: 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/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f1343bd143d..6630e184a16 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 @@ -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"