mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(mcp): bound discovery cache result bytes
This commit is contained in:
parent
05d2c316f5
commit
9d31de2f20
2 changed files with 37 additions and 22 deletions
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue