mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
refactor(mcp): reuse in-memory discovery storage
This commit is contained in:
parent
ceb1f04988
commit
05d2c316f5
4 changed files with 112 additions and 43 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue