mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): honor scoped cache freshness (#45165)
* fix(mcp): honor upstream freshness and caller scope in discovery caches * test(mcp): type scoped freshness regression helpers * fix(mcp): age discovery freshness through cleanup * fix(mcp): preserve keyless discovery caller isolation --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
e9cfba2c17
commit
6498bef8db
12 changed files with 701 additions and 221 deletions
|
|
@ -7,6 +7,7 @@ import base64
|
|||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from functools import partial
|
||||
|
|
@ -35,6 +36,7 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams]
|
|||
from mcp.types import (
|
||||
METHOD_NOT_FOUND,
|
||||
REQUEST_TIMEOUT,
|
||||
CacheableResult,
|
||||
ClientCapabilities,
|
||||
DiscoverResult,
|
||||
ElicitationCapability,
|
||||
|
|
@ -55,7 +57,6 @@ from mcp.types import (
|
|||
ListToolsRequest,
|
||||
ListToolsResult,
|
||||
PaginatedRequestParams,
|
||||
PaginatedResult,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
SamplingCapability,
|
||||
|
|
@ -77,7 +78,11 @@ from litellm.constants import (
|
|||
from litellm.experimental_mcp_client.tools import list_tools_with_pagination
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import (
|
||||
age_freshness,
|
||||
aggregate_freshness,
|
||||
error_text_result,
|
||||
)
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCP_LEGACY_VERSIONS,
|
||||
|
|
@ -178,7 +183,7 @@ def as_mcp_read_timeout(exc: BaseException) -> TimeoutError | None:
|
|||
|
||||
|
||||
TSessionResult = TypeVar("TSessionResult")
|
||||
_ListPage = TypeVar("_ListPage", bound=PaginatedResult)
|
||||
_ListPage = TypeVar("_ListPage", ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult)
|
||||
_ListItem = TypeVar("_ListItem")
|
||||
|
||||
|
||||
|
|
@ -1031,12 +1036,21 @@ class MCPClient:
|
|||
# Return a default error result instead of raising
|
||||
return self.error_tool_result(e)
|
||||
|
||||
async def _run_optional_discovery(self, operation: Callable[[ClientSession], Awaitable[_ListPage]]) -> _ListPage:
|
||||
async def timed_operation(session: ClientSession) -> tuple[_ListPage, float]:
|
||||
result: Final = await operation(session)
|
||||
return result, time.monotonic()
|
||||
|
||||
result, received = await self.run_with_session(timed_operation)
|
||||
return age_freshness(result, time.monotonic() - received)
|
||||
|
||||
async def _list_optional_pages(
|
||||
self,
|
||||
fetch_page: Callable[[PaginatedRequestParams | None], Awaitable[_ListPage]],
|
||||
items_of: Callable[[_ListPage], Sequence[_ListItem]],
|
||||
) -> list[_ListItem]: # mutable-ok: existing list discovery API
|
||||
) -> tuple[list[_ListItem], CacheableResult]:
|
||||
items: Final[list[_ListItem]] = [] # mutable-ok: bounded iterative page accumulation
|
||||
pages: Final[list[tuple[CacheableResult, float]]] = [] # mutable-ok: bounded pagination evidence
|
||||
cursors: Final[set[str]] = set() # mutable-ok: constant-time detection of cursor cycles
|
||||
cursor: str | None = None # rebind-ok: iterative traversal avoids recursion at the existing page cap
|
||||
with anyio.fail_after(max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)):
|
||||
|
|
@ -1047,9 +1061,13 @@ class MCPClient:
|
|||
if page_index > 0 and error.error.code == METHOD_NOT_FOUND:
|
||||
raise RuntimeError("MCP list operation became unavailable during pagination") from error
|
||||
raise
|
||||
pages.append((page, time.monotonic()))
|
||||
items.extend(items_of(page))
|
||||
if not page.next_cursor:
|
||||
return items
|
||||
now: Final = time.monotonic()
|
||||
return items, aggregate_freshness(
|
||||
tuple(age_freshness(value, now - received) for value, received in pages)
|
||||
)
|
||||
if page.next_cursor in cursors:
|
||||
raise RuntimeError("MCP list pagination repeated a cursor")
|
||||
cursors.add(page.next_cursor)
|
||||
|
|
@ -1057,6 +1075,9 @@ class MCPClient:
|
|||
raise RuntimeError(f"MCP list pagination exceeded {MCP_TOOL_LISTING_MAX_PAGES} pages")
|
||||
|
||||
async def list_prompts(self, *, raise_on_error: bool = False) -> list[Prompt]:
|
||||
return (await self.list_prompts_result(raise_on_error=raise_on_error)).prompts
|
||||
|
||||
async def list_prompts_result(self, *, raise_on_error: bool = False) -> ListPromptsResult:
|
||||
"""List available prompts from the server."""
|
||||
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
|
||||
|
||||
|
|
@ -1065,11 +1086,10 @@ class MCPClient:
|
|||
if capabilities is not None and capabilities.prompts is None:
|
||||
return ListPromptsResult(prompts=[])
|
||||
try:
|
||||
return ListPromptsResult(
|
||||
prompts=await self._list_optional_pages(
|
||||
lambda params: session.list_prompts(params=params), lambda page: page.prompts
|
||||
)
|
||||
items, freshness = await self._list_optional_pages(
|
||||
lambda params: session.list_prompts(params=params), lambda page: page.prompts
|
||||
)
|
||||
return ListPromptsResult(prompts=items, ttl_ms=freshness.ttl_ms, cache_scope=freshness.cache_scope)
|
||||
except MCPError as error:
|
||||
if error.error.code != METHOD_NOT_FOUND:
|
||||
raise
|
||||
|
|
@ -1079,13 +1099,13 @@ class MCPClient:
|
|||
return ListPromptsResult(prompts=[])
|
||||
|
||||
try:
|
||||
result: Final = await self.run_with_session(_list_prompts_operation)
|
||||
result: Final = await self._run_optional_discovery(_list_prompts_operation)
|
||||
prompt_count: Final = len(result.prompts)
|
||||
prompt_names: Final = [prompt.name for prompt in result.prompts]
|
||||
verbose_logger.info(
|
||||
"MCP client listed %s tools from %s: %s", prompt_count, self.server_url or "stdio", prompt_names
|
||||
)
|
||||
return result.prompts
|
||||
return result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client list_prompts was cancelled")
|
||||
raise
|
||||
|
|
@ -1107,7 +1127,7 @@ class MCPClient:
|
|||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
)
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
return ListPromptsResult(prompts=[])
|
||||
|
||||
async def get_prompt(self, get_prompt_request_params: GetPromptRequestParams) -> GetPromptResult:
|
||||
"""Fetch a prompt definition from the MCP server."""
|
||||
|
|
@ -1151,6 +1171,9 @@ class MCPClient:
|
|||
raise
|
||||
|
||||
async def list_resources(self, *, raise_on_error: bool = False) -> list[Resource]:
|
||||
return (await self.list_resources_result(raise_on_error=raise_on_error)).resources
|
||||
|
||||
async def list_resources_result(self, *, raise_on_error: bool = False) -> ListResourcesResult:
|
||||
"""List available resources from the server."""
|
||||
verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio")
|
||||
|
||||
|
|
@ -1159,11 +1182,10 @@ class MCPClient:
|
|||
if capabilities is not None and capabilities.resources is None:
|
||||
return ListResourcesResult(resources=[])
|
||||
try:
|
||||
return ListResourcesResult(
|
||||
resources=await self._list_optional_pages(
|
||||
lambda params: session.list_resources(params=params), lambda page: page.resources
|
||||
)
|
||||
items, freshness = await self._list_optional_pages(
|
||||
lambda params: session.list_resources(params=params), lambda page: page.resources
|
||||
)
|
||||
return ListResourcesResult(resources=items, ttl_ms=freshness.ttl_ms, cache_scope=freshness.cache_scope)
|
||||
except MCPError as error:
|
||||
if error.error.code != METHOD_NOT_FOUND:
|
||||
raise
|
||||
|
|
@ -1173,13 +1195,13 @@ class MCPClient:
|
|||
return ListResourcesResult(resources=[])
|
||||
|
||||
try:
|
||||
result: Final = await self.run_with_session(_list_resources_operation)
|
||||
result: Final = await self._run_optional_discovery(_list_resources_operation)
|
||||
resource_count: Final = len(result.resources)
|
||||
resource_names: Final = [resource.name for resource in result.resources]
|
||||
verbose_logger.info(
|
||||
"MCP client listed %s resources from %s: %s", resource_count, self.server_url or "stdio", resource_names
|
||||
)
|
||||
return result.resources
|
||||
return result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client list_resources was cancelled")
|
||||
raise
|
||||
|
|
@ -1201,9 +1223,12 @@ class MCPClient:
|
|||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
)
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
return ListResourcesResult(resources=[])
|
||||
|
||||
async def list_resource_templates(self, *, raise_on_error: bool = False) -> list[ResourceTemplate]:
|
||||
return (await self.list_resource_templates_result(raise_on_error=raise_on_error)).resource_templates
|
||||
|
||||
async def list_resource_templates_result(self, *, raise_on_error: bool = False) -> ListResourceTemplatesResult:
|
||||
"""List available resource templates from the server."""
|
||||
verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio")
|
||||
|
||||
|
|
@ -1212,11 +1237,11 @@ class MCPClient:
|
|||
if capabilities is not None and capabilities.resources is None:
|
||||
return ListResourceTemplatesResult(resource_templates=[])
|
||||
try:
|
||||
items, freshness = await self._list_optional_pages(
|
||||
lambda params: session.list_resource_templates(params=params), lambda page: page.resource_templates
|
||||
)
|
||||
return ListResourceTemplatesResult(
|
||||
resource_templates=await self._list_optional_pages(
|
||||
lambda params: session.list_resource_templates(params=params),
|
||||
lambda page: page.resource_templates,
|
||||
)
|
||||
resource_templates=items, ttl_ms=freshness.ttl_ms, cache_scope=freshness.cache_scope
|
||||
)
|
||||
except MCPError as error:
|
||||
if error.error.code != METHOD_NOT_FOUND:
|
||||
|
|
@ -1227,7 +1252,7 @@ class MCPClient:
|
|||
return ListResourceTemplatesResult(resource_templates=[])
|
||||
|
||||
try:
|
||||
result: Final = await self.run_with_session(_list_resource_templates_operation)
|
||||
result: Final = await self._run_optional_discovery(_list_resource_templates_operation)
|
||||
resource_template_count: Final = len(result.resource_templates)
|
||||
resource_template_names: Final = [resource_template.name for resource_template in result.resource_templates]
|
||||
verbose_logger.info(
|
||||
|
|
@ -1236,7 +1261,7 @@ class MCPClient:
|
|||
self.server_url or "stdio",
|
||||
resource_template_names,
|
||||
)
|
||||
return result.resource_templates
|
||||
return result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client list_resource_templates was cancelled")
|
||||
raise
|
||||
|
|
@ -1258,7 +1283,7 @@ class MCPClient:
|
|||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
)
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
return ListResourceTemplatesResult(resource_templates=[])
|
||||
|
||||
async def read_resource(self, url: AnyUrl) -> ReadResourceResult:
|
||||
"""Fetch resource contents from the MCP server."""
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import asyncio
|
|||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from collections import UserDict
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, MutableMapping, Sequence
|
||||
from contextlib import ExitStack, asynccontextmanager
|
||||
|
|
@ -14,11 +15,14 @@ from dataclasses import dataclass, replace
|
|||
from functools import partial, wraps
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
|
||||
from typing import TYPE_CHECKING, Final, Generic, ParamSpec, TypeAlias, TypeVar, cast
|
||||
|
||||
from mcp.types import CacheableResult
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import age_freshness, aggregate_freshness
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.types import ListToolsResult, PaginatedRequestParams, PaginatedResult
|
||||
|
|
@ -659,6 +663,7 @@ async def paginate_catalog(
|
|||
result,
|
||||
)
|
||||
|
||||
started: Final = time.monotonic()
|
||||
tasks: Final = tuple(asyncio.create_task(advance(position)) for position in state.positions)
|
||||
try:
|
||||
results: Final = await asyncio.gather(*tasks)
|
||||
|
|
@ -668,7 +673,12 @@ async def paginate_catalog(
|
|||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
|
||||
|
||||
pages: Final = tuple(result for _, result in results if result is not None)
|
||||
elapsed: Final = time.monotonic() - started
|
||||
pages: Final = tuple(
|
||||
age_freshness(result, elapsed) if isinstance(result, CacheableResult) else result
|
||||
for _, result in results
|
||||
if result is not None
|
||||
)
|
||||
page_outcomes: Final = (
|
||||
_OUTCOME_VALUES.validate_python((result.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})) for result in pages
|
||||
)
|
||||
|
|
@ -720,6 +730,10 @@ async def list_tools_page(
|
|||
)
|
||||
return ListToolsResult(
|
||||
tools=list(chain.from_iterable(page.tools for page in pages)),
|
||||
ttl_ms=0
|
||||
if any(isinstance(value, dict) and value.get("tag") != "ok" for value in outcomes.values())
|
||||
else aggregate_freshness(pages).ttl_ms,
|
||||
cache_scope="private",
|
||||
next_cursor=next_cursor,
|
||||
_meta={SERVER_OUTCOMES_META_KEY: dict(outcomes)} if outcomes else None,
|
||||
)
|
||||
|
|
@ -1036,6 +1050,7 @@ async def aggregate_gateway_tools(
|
|||
(result.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})
|
||||
),
|
||||
next_cursor=result.next_cursor,
|
||||
ttl_ms=result.ttl_ms,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1065,6 +1080,8 @@ async def list_gateway_tools(
|
|||
return ListToolsResult(
|
||||
tools=listing.tools,
|
||||
next_cursor=listing.next_cursor,
|
||||
ttl_ms=listing.ttl_ms,
|
||||
cache_scope="private",
|
||||
_meta={
|
||||
SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()}
|
||||
}
|
||||
|
|
@ -1240,10 +1257,19 @@ def combine_optional_catalog(
|
|||
ListResourceTemplatesResult,
|
||||
)
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
|
||||
|
||||
outcomes: Final = (meta or {}).get(SERVER_OUTCOMES_META_KEY)
|
||||
incomplete: Final = isinstance(outcomes, dict) and any(
|
||||
isinstance(value, dict) and value.get("tag") != "ok" for value in outcomes.values()
|
||||
)
|
||||
ttl_ms: Final = 0 if incomplete else aggregate_freshness(pages).ttl_ms
|
||||
if isinstance(request, ListPromptsRequest):
|
||||
return ListPromptsResult(
|
||||
prompts=list(chain.from_iterable(page.prompts for page in pages if isinstance(page, ListPromptsResult))),
|
||||
next_cursor=next_cursor,
|
||||
ttl_ms=ttl_ms,
|
||||
cache_scope="private",
|
||||
_meta=dict(meta) if meta is not None else None,
|
||||
)
|
||||
if isinstance(request, ListResourcesRequest):
|
||||
|
|
@ -1252,6 +1278,8 @@ def combine_optional_catalog(
|
|||
chain.from_iterable(page.resources for page in pages if isinstance(page, ListResourcesResult))
|
||||
),
|
||||
next_cursor=next_cursor,
|
||||
ttl_ms=ttl_ms,
|
||||
cache_scope="private",
|
||||
_meta=dict(meta) if meta is not None else None,
|
||||
)
|
||||
return ListResourceTemplatesResult(
|
||||
|
|
@ -1261,5 +1289,90 @@ def combine_optional_catalog(
|
|||
)
|
||||
),
|
||||
next_cursor=next_cursor,
|
||||
ttl_ms=ttl_ms,
|
||||
cache_scope="private",
|
||||
_meta=dict(meta) if meta is not None else None,
|
||||
)
|
||||
|
||||
|
||||
_DiscoveryPage = TypeVar("_DiscoveryPage", bound=CacheableResult)
|
||||
_DiscoveryKey: TypeAlias = tuple[str, str | None]
|
||||
_DISCOVERY_CACHE_LIMIT: Final = 1024
|
||||
_DISCOVERY_ENTRY: Final = TypeAdapter(tuple[float, bytes])
|
||||
|
||||
|
||||
class _DiscoveryCache(Generic[_DiscoveryPage]):
|
||||
def __init__(self, ttl: float, clock: Callable[[], float], adapter: TypeAdapter[_DiscoveryPage]) -> None:
|
||||
self._ttl = ttl
|
||||
self._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[_DiscoveryPage]] = {}
|
||||
self._waiters: dict[asyncio.Task[_DiscoveryPage], int] = {} # mutable-ok: constant-time waiter accounting
|
||||
|
||||
def invalidate(self, server_id: str) -> None:
|
||||
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[_DiscoveryPage]) -> None:
|
||||
if not task.cancelled():
|
||||
task.exception()
|
||||
|
||||
async def get(self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[_DiscoveryPage]]) -> _DiscoveryPage:
|
||||
if self._ttl <= 0:
|
||||
return await fetch()
|
||||
entry: Final[object] = self._entries.get_cache(json.dumps(key))
|
||||
if entry is not None:
|
||||
expires_at, payload = _DISCOVERY_ENTRY.validate_python(entry)
|
||||
remaining: Final = max(0, int((expires_at - self._clock()) * 1000))
|
||||
if remaining > 0:
|
||||
return self._adapter.validate_json(payload).model_copy(update={"ttl_ms": remaining})
|
||||
self._entries.delete_cache(json.dumps(key))
|
||||
pending: Final = self._pending.get(key)
|
||||
if pending is not None:
|
||||
return await self._await_fetch(key, pending)
|
||||
if len(self._pending) >= _DISCOVERY_CACHE_LIMIT:
|
||||
return await fetch()
|
||||
task: Final = asyncio.create_task(self._fetch(key, fetch))
|
||||
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[_DiscoveryPage]) -> _DiscoveryPage:
|
||||
self._waiters[task] = self._waiters.get(task, 0) + 1
|
||||
try:
|
||||
return (await asyncio.shield(task)).model_copy(deep=True)
|
||||
finally:
|
||||
remaining: Final = self._waiters[task] - 1
|
||||
if remaining:
|
||||
self._waiters[task] = remaining
|
||||
else:
|
||||
self._waiters.pop(task)
|
||||
if self._pending.get(key) is task:
|
||||
self._pending.pop(key)
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
|
||||
async def _fetch(self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[_DiscoveryPage]]) -> _DiscoveryPage:
|
||||
try:
|
||||
items: Final = await fetch()
|
||||
ttl: Final = min(self._ttl, items.ttl_ms / 1000)
|
||||
if ttl > 0 and self._pending.get(key) is asyncio.current_task():
|
||||
self._entries.set_cache(
|
||||
json.dumps(key),
|
||||
_DISCOVERY_ENTRY.dump_json((self._clock() + ttl, self._adapter.dump_json(items))),
|
||||
ttl=ttl,
|
||||
)
|
||||
return items
|
||||
finally:
|
||||
if self._pending.get(key) is asyncio.current_task():
|
||||
self._pending.pop(key)
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ class AggregateToolListing(NamedTuple):
|
|||
tools: list[MCPTool]
|
||||
outcomes: dict[str, ServerOutcome]
|
||||
next_cursor: str | None = None
|
||||
ttl_ms: int = 0
|
||||
|
||||
|
||||
def _iter_upstream_responses(exc: BaseException) -> Iterator[httpx.Response | httpx2.Response]:
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ from dataclasses import dataclass, replace
|
|||
from functools import lru_cache
|
||||
from itertools import chain, groupby
|
||||
from types import EllipsisType, MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast
|
||||
from urllib.parse import ParseResult, urlparse
|
||||
|
||||
import anyio
|
||||
|
|
@ -45,15 +45,18 @@ from mcp.types import (
|
|||
GetPromptResult,
|
||||
InputRequiredResult,
|
||||
ListPromptsRequest,
|
||||
ListPromptsResult,
|
||||
ListResourcesRequest,
|
||||
ListResourcesResult,
|
||||
ListResourceTemplatesRequest,
|
||||
ListResourceTemplatesResult,
|
||||
ListToolsResult,
|
||||
PaginatedRequestParams,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl, BaseModel, Field, TypeAdapter
|
||||
from pydantic import AnyUrl, Field, TypeAdapter
|
||||
from typing_extensions import ReadOnly, assert_never
|
||||
|
||||
import litellm
|
||||
|
|
@ -78,6 +81,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
|||
MCPServerAccess,
|
||||
_is_mcp_admitted_user_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import _configuration_identity, _DiscoveryCache, _DiscoveryKey
|
||||
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
MCP_ELICITATION_AVAILABLE,
|
||||
|
|
@ -1744,90 +1748,6 @@ def _record_mcp_guardrail_evaluations(
|
|||
verbose_logger.warning("Failed to record MCP guardrail evaluation for logging: %s", e)
|
||||
|
||||
|
||||
_DiscoveryItem = TypeVar("_DiscoveryItem", bound=BaseModel)
|
||||
_DiscoveryKey: TypeAlias = tuple[str, str | None]
|
||||
_DISCOVERY_CACHE_LIMIT: Final = 1024
|
||||
|
||||
|
||||
class _DiscoveryCache(Generic[_DiscoveryItem]):
|
||||
def __init__(
|
||||
self, ttl: float, clock: Callable[[], float], adapter: TypeAdapter[tuple[_DiscoveryItem, ...]]
|
||||
) -> None:
|
||||
self._ttl = ttl
|
||||
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]]] = {}
|
||||
self._waiters: dict[asyncio.Task[list[_DiscoveryItem]], int] = {} # mutable-ok: constant-time waiter accounting
|
||||
|
||||
def invalidate(self, server_id: str) -> None:
|
||||
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:
|
||||
if not task.cancelled():
|
||||
task.exception()
|
||||
|
||||
async def get(
|
||||
self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[list[_DiscoveryItem]]]
|
||||
) -> tuple[_DiscoveryItem, ...]:
|
||||
if self._ttl <= 0:
|
||||
return tuple(await fetch())
|
||||
entry: Final[object] = self._entries.get_cache(json.dumps(key))
|
||||
if entry is not None:
|
||||
return self._adapter.validate_python(entry)
|
||||
pending: Final = self._pending.get(key)
|
||||
if pending is not None:
|
||||
return await self._await_fetch(key, pending)
|
||||
if len(self._pending) >= _DISCOVERY_CACHE_LIMIT:
|
||||
return tuple(await fetch())
|
||||
task: Final = asyncio.create_task(self._fetch(key, fetch))
|
||||
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[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[task] = remaining
|
||||
else:
|
||||
self._waiters.pop(task)
|
||||
if self._pending.get(key) is task:
|
||||
self._pending.pop(key)
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
|
||||
async def _fetch(
|
||||
self, key: _DiscoveryKey, fetch: Callable[[], Awaitable[list[_DiscoveryItem]]]
|
||||
) -> list[_DiscoveryItem]:
|
||||
try:
|
||||
items: Final = await fetch()
|
||||
if self._pending.get(key) is asyncio.current_task():
|
||||
self._entries.set_cache(
|
||||
json.dumps(key),
|
||||
self._adapter.dump_json(tuple(items)),
|
||||
ttl=self._ttl,
|
||||
)
|
||||
return items
|
||||
finally:
|
||||
if self._pending.get(key) is asyncio.current_task():
|
||||
self._pending.pop(key)
|
||||
|
||||
|
||||
def _mcp_discovery_cache_ttl() -> float:
|
||||
raw: Final = os.environ.get("LITELLM_MCP_DISCOVERY_CACHE_TTL", "60")
|
||||
try:
|
||||
|
|
@ -1968,14 +1888,14 @@ class MCPServerManager:
|
|||
token_exchanger=build_token_exchanger(),
|
||||
)
|
||||
discovery_ttl: Final = _mcp_discovery_cache_ttl()
|
||||
self._prompt_discovery_cache = _DiscoveryCache[Prompt](
|
||||
discovery_ttl, discovery_clock, TypeAdapter(tuple[Prompt, ...])
|
||||
self._prompt_discovery_cache = _DiscoveryCache[ListPromptsResult](
|
||||
discovery_ttl, discovery_clock, TypeAdapter(ListPromptsResult)
|
||||
)
|
||||
self._resource_discovery_cache = _DiscoveryCache[Resource](
|
||||
discovery_ttl, discovery_clock, TypeAdapter(tuple[Resource, ...])
|
||||
self._resource_discovery_cache = _DiscoveryCache[ListResourcesResult](
|
||||
discovery_ttl, discovery_clock, TypeAdapter(ListResourcesResult)
|
||||
)
|
||||
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
|
||||
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
|
||||
self._template_discovery_cache = _DiscoveryCache[ListResourceTemplatesResult](
|
||||
discovery_ttl, discovery_clock, TypeAdapter(ListResourceTemplatesResult)
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import CatalogSnapshots
|
||||
|
||||
|
|
@ -4608,24 +4528,34 @@ class MCPServerManager:
|
|||
stdio_env: dict[str, str] | None,
|
||||
subject_token: str | None,
|
||||
credential_fingerprint: str | None = None,
|
||||
per_caller: bool = False,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
) -> _DiscoveryKey:
|
||||
per_user: Final = (
|
||||
per_caller
|
||||
or server.requires_per_user_auth
|
||||
or self._references_per_user_env_var(server)
|
||||
or server.delegate_auth_to_upstream
|
||||
or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)
|
||||
)
|
||||
if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token):
|
||||
return server.server_id, None
|
||||
identity: Final = (
|
||||
(user_api_key_auth.user_id, user_api_key_auth.api_key)
|
||||
if per_user and user_api_key_auth is not None
|
||||
user_api_key_auth.model_dump(
|
||||
include={
|
||||
"end_user_id",
|
||||
"user_role",
|
||||
"object_permission_id",
|
||||
"team_object_permission_id",
|
||||
"team_object_permission",
|
||||
"end_user_object_permission",
|
||||
},
|
||||
mode="json",
|
||||
)
|
||||
if user_api_key_auth is not None
|
||||
else None
|
||||
)
|
||||
material: Final = json.dumps(
|
||||
(identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint),
|
||||
(
|
||||
_configuration_identity(server),
|
||||
_admission_identity(user_api_key_auth, raw_headers) if user_api_key_auth is not None else None,
|
||||
identity,
|
||||
mcp_auth_header,
|
||||
extra_headers,
|
||||
stdio_env,
|
||||
subject_token,
|
||||
credential_fingerprint,
|
||||
),
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
|
|
@ -4741,14 +4671,21 @@ class MCPServerManager:
|
|||
)
|
||||
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
|
||||
key: Final = self._discovery_key(
|
||||
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
|
||||
server,
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
headers,
|
||||
stdio_env,
|
||||
subject_token,
|
||||
credential_fingerprint,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
async def fetch() -> list[Prompt]:
|
||||
return await client.list_prompts(raise_on_error=True)
|
||||
async def fetch() -> ListPromptsResult:
|
||||
return await client.list_prompts_result(raise_on_error=True)
|
||||
|
||||
items: Final = await self._prompt_discovery_cache.get(key, fetch)
|
||||
return self._create_prefixed_prompts(items, server, add_prefix=add_prefix)
|
||||
return self._create_prefixed_prompts(items.prompts, server, add_prefix=add_prefix)
|
||||
except Exception as error:
|
||||
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error)
|
||||
return []
|
||||
|
|
@ -4789,14 +4726,21 @@ class MCPServerManager:
|
|||
)
|
||||
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
|
||||
key: Final = self._discovery_key(
|
||||
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
|
||||
server,
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
headers,
|
||||
stdio_env,
|
||||
subject_token,
|
||||
credential_fingerprint,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
async def fetch() -> list[Resource]:
|
||||
return await client.list_resources(raise_on_error=True)
|
||||
async def fetch() -> ListResourcesResult:
|
||||
return await client.list_resources_result(raise_on_error=True)
|
||||
|
||||
items: Final = await self._resource_discovery_cache.get(key, fetch)
|
||||
return self._create_prefixed_resources(items, server, add_prefix=add_prefix)
|
||||
return self._create_prefixed_resources(items.resources, server, add_prefix=add_prefix)
|
||||
except Exception as error:
|
||||
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error)
|
||||
return []
|
||||
|
|
@ -4837,14 +4781,21 @@ class MCPServerManager:
|
|||
)
|
||||
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
|
||||
key: Final = self._discovery_key(
|
||||
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token, credential_fingerprint
|
||||
server,
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
headers,
|
||||
stdio_env,
|
||||
subject_token,
|
||||
credential_fingerprint,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
async def fetch() -> list[ResourceTemplate]:
|
||||
return await client.list_resource_templates(raise_on_error=True)
|
||||
async def fetch() -> ListResourceTemplatesResult:
|
||||
return await client.list_resource_templates_result(raise_on_error=True)
|
||||
|
||||
items: Final = await self._template_discovery_cache.get(key, fetch)
|
||||
return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix)
|
||||
return self._create_prefixed_resource_templates(items.resource_templates, server, add_prefix=add_prefix)
|
||||
except Exception as error:
|
||||
verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error)
|
||||
return []
|
||||
|
|
|
|||
|
|
@ -11,9 +11,11 @@ revisions (``2024-11-05`` .. ``2025-11-25``) and admits any JSON value, plus
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final, TypeAlias
|
||||
import math
|
||||
from collections.abc import Sequence
|
||||
from typing import Final, TypeAlias, TypeVar
|
||||
|
||||
from mcp.types import CallToolResult, ContentBlock, InputRequiredResult, TextContent, Tool
|
||||
from mcp.types import CacheableResult, CallToolResult, ContentBlock, InputRequiredResult, TextContent, Tool
|
||||
from typing_extensions import ReadOnly, TypedDict, assert_never
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import (
|
||||
|
|
@ -118,3 +120,17 @@ def _downgrade_structured_content(result: CallToolResult) -> CallToolResult:
|
|||
def to_gateway_tool(tool: Tool, name: str) -> Tool:
|
||||
update: Final[_Renamed] = {"name": name}
|
||||
return tool.model_copy(deep=True, update=update)
|
||||
|
||||
|
||||
_Cacheable = TypeVar("_Cacheable", bound=CacheableResult)
|
||||
|
||||
|
||||
def age_freshness(result: _Cacheable, elapsed: float) -> _Cacheable:
|
||||
return result.model_copy(update={"ttl_ms": max(0, result.ttl_ms - math.ceil(max(0.0, elapsed) * 1000))})
|
||||
|
||||
|
||||
def aggregate_freshness(results: Sequence[CacheableResult]) -> CacheableResult:
|
||||
return CacheableResult(
|
||||
ttl_ms=min((result.ttl_ms for result in results), default=0),
|
||||
cache_scope="public" if results and all(result.cache_scope == "public" for result in results) else "private",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -628,6 +628,7 @@ def tool_calls(observed: tuple[dict[str, object], ...]) -> tuple[dict[str, objec
|
|||
def paginated_mcp_peer(
|
||||
*,
|
||||
page_size: int = 1,
|
||||
ttl_ms: int = 0,
|
||||
repeat_cursor: bool = False,
|
||||
fail_listing: bool = False,
|
||||
fail_continuation: bool = False,
|
||||
|
|
@ -664,6 +665,8 @@ def paginated_mcp_peer(
|
|||
async def tools(context, params):
|
||||
indexes, cursor = window(params)
|
||||
return ListToolsResult(
|
||||
ttl_ms=ttl_ms,
|
||||
cache_scope="public",
|
||||
tools=[
|
||||
Tool(
|
||||
name=f"add{index}",
|
||||
|
|
@ -682,12 +685,18 @@ def paginated_mcp_peer(
|
|||
async def prompts(context, params):
|
||||
indexes, cursor = window(params)
|
||||
return ListPromptsResult(
|
||||
prompts=[Prompt(name=f"prompt{index}") for index in indexes], next_cursor=cursor, meta=metadata
|
||||
ttl_ms=ttl_ms,
|
||||
cache_scope="public",
|
||||
prompts=[Prompt(name=f"prompt{index}") for index in indexes],
|
||||
next_cursor=cursor,
|
||||
meta=metadata,
|
||||
)
|
||||
|
||||
async def resources(context, params):
|
||||
indexes, cursor = window(params)
|
||||
return ListResourcesResult(
|
||||
ttl_ms=ttl_ms,
|
||||
cache_scope="public",
|
||||
resources=[Resource(name=f"resource{index}", uri=f"status://item{index}") for index in indexes],
|
||||
next_cursor=cursor,
|
||||
meta=metadata,
|
||||
|
|
@ -696,6 +705,8 @@ def paginated_mcp_peer(
|
|||
async def templates(context, params):
|
||||
indexes, cursor = window(params)
|
||||
return ListResourceTemplatesResult(
|
||||
ttl_ms=ttl_ms,
|
||||
cache_scope="public",
|
||||
resource_templates=[
|
||||
ResourceTemplate(name=f"template{index}", uri_template=f"status{index}://{{item}}") for index in indexes
|
||||
],
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from mcp.client.streamable_http import streamable_http_client
|
|||
from mcp.types import CallToolRequest, CallToolRequestParams, CallToolResult, PaginatedRequestParams
|
||||
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.mcp import paginated_mcp_peer
|
||||
from integration._support.mcp import McpPeer, paginated_mcp_peer
|
||||
from integration._support.process import owned_proxy
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
|
@ -562,3 +562,59 @@ def test_missing_user_keeps_explicit_key_and_team_grants(
|
|||
assert not forbidden.ok, forbidden
|
||||
assert tool_calls(allowed.drain()) == ()
|
||||
assert private.drain() == ()
|
||||
|
||||
|
||||
def test_aggregate_freshness_matches_real_upstream() -> None:
|
||||
from mcp.types import (
|
||||
ListPromptsRequest,
|
||||
ListResourcesRequest,
|
||||
ListResourceTemplatesRequest,
|
||||
ListToolsRequest,
|
||||
ListToolsResult,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import combine_optional_catalog, list_tools_page
|
||||
|
||||
async def exercise(peer: McpPeer) -> None:
|
||||
client: Final = MCPClient(server_url=peer.url, transport_type=MCPTransport.http, protocol_version="2026-07-28")
|
||||
for request in (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest()):
|
||||
page: Final = await client.list_page(request)
|
||||
assert page.ttl_ms == 9000
|
||||
result: Final = combine_optional_catalog(request, [page], None, None)
|
||||
assert result.ttl_ms == 9000
|
||||
assert result.cache_scope == "private"
|
||||
|
||||
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return await client.list_page(ListToolsRequest())
|
||||
|
||||
tools_result: Final = await list_tools_page(
|
||||
cursor=None, caller_scope="caller", snapshot="snapshot", server_ids=("server",), fetch=fetch, now=100
|
||||
)
|
||||
assert 0 < tools_result.ttl_ms <= 9000
|
||||
assert tools_result.cache_scope == "private"
|
||||
assert len(tools_result.tools) == 3
|
||||
|
||||
with paginated_mcp_peer(page_size=3, ttl_ms=9000) as peer:
|
||||
asyncio.run(exercise(peer))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ttl_ms", [0, 9000])
|
||||
def test_discovery_cache_freshness_and_caller_isolation_over_http(ttl_ms: int) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
async def exercise(peer: McpPeer) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="pages", name="pages", url=peer.url, transport=MCPTransport.http, protocol_version="2026-07-28"
|
||||
)
|
||||
for manager in (MCPServerManager(), MCPServerManager()):
|
||||
for user in (UserAPIKeyAuth(user_id="one"), UserAPIKeyAuth(user_id="two")):
|
||||
for _ in range(2):
|
||||
result: Final = await manager.get_prompts_from_server(server, user)
|
||||
assert [prompt.name for prompt in result] == ["pages-prompt0", "pages-prompt1", "pages-prompt2"]
|
||||
|
||||
with paginated_mcp_peer(page_size=3, ttl_ms=ttl_ms) as peer:
|
||||
asyncio.run(exercise(peer))
|
||||
observed: Final = peer.drain()
|
||||
listings: Final = [item for item in observed if item.get("body", {}).get("method") == "prompts/list"]
|
||||
assert len(listings) == (8 if ttl_ms == 0 else 4)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import os
|
|||
import selectors
|
||||
import sys
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
from typing import Final
|
||||
|
|
@ -42,6 +43,7 @@ from litellm.experimental_mcp_client.client import (
|
|||
MCPClient,
|
||||
_first_non_cancelled_cause,
|
||||
_TransportContext,
|
||||
_TransportStreams,
|
||||
as_mcp_read_timeout,
|
||||
strip_auth_scheme,
|
||||
)
|
||||
|
|
@ -71,14 +73,26 @@ def _initialized(instructions: str | None = None) -> InitializeResult:
|
|||
class _MockTransportClient(MCPClient):
|
||||
"""An MCPClient whose streamable-HTTP transport runs on an httpx2 MockTransport."""
|
||||
|
||||
def __init__(self, respond, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
respond,
|
||||
*,
|
||||
http_transport: httpx2.AsyncBaseTransport | None = None,
|
||||
transport_context: Callable[[str, httpx2.AsyncClient], _TransportContext] | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self._respond = respond
|
||||
self._http_transport = http_transport
|
||||
self._transport_context = transport_context
|
||||
|
||||
def _create_transport_context(self) -> tuple[_TransportContext, httpx2.AsyncClient]:
|
||||
http_client: Final = self._create_httpx_client_factory(transport=httpx2.MockTransport(self._respond))(
|
||||
transport: Final = self._http_transport or httpx2.MockTransport(self._respond)
|
||||
http_client: Final = self._create_httpx_client_factory(transport=transport)(
|
||||
headers=self._get_auth_headers(), timeout=httpx2.Timeout(self.timeout)
|
||||
)
|
||||
if self._transport_context is not None:
|
||||
return self._transport_context(self.server_url, http_client), http_client
|
||||
return streamable_http_client(self.server_url, http_client=http_client), http_client
|
||||
|
||||
|
||||
|
|
@ -3520,3 +3534,81 @@ async def test_optional_catalog_distinguishes_absent_capability_from_failed_cont
|
|||
assert result.next_cursor is None
|
||||
if failure == "unadvertised":
|
||||
assert method not in methods
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"kind,field,entry",
|
||||
[
|
||||
("prompts", "prompts", {"name": "example"}),
|
||||
("resources", "resources", {"name": "example", "uri": "test://example"}),
|
||||
("resource_templates", "resourceTemplates", {"name": "example", "uriTemplate": "test://{name}"}),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("ttl", [0, 5000])
|
||||
@pytest.mark.parametrize("cleanup_phase", ["transport", "http_client"])
|
||||
@pytest.mark.parametrize("cleanup_seconds", [0, 2, 6])
|
||||
async def test_optional_discovery_retains_freshness_across_pages(
|
||||
kind: str, field: str, entry: dict[str, str], ttl: int, cleanup_phase: str, cleanup_seconds: int
|
||||
) -> None:
|
||||
def respond(request: httpx2.Request) -> httpx2.Response:
|
||||
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
|
||||
assert isinstance(payload, JSONRPCRequest)
|
||||
if payload.method == "server/discover":
|
||||
discovery_result: Final = {
|
||||
"supportedVersions": ["2026-07-28"],
|
||||
"capabilities": {"prompts": {}, "resources": {}},
|
||||
"ttlMs": 0,
|
||||
"cacheScope": "private",
|
||||
"resultType": "complete",
|
||||
}
|
||||
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": discovery_result})
|
||||
following: Final = bool((payload.params or {}).get("cursor"))
|
||||
listing_result: Final = {
|
||||
field: [entry],
|
||||
"ttlMs": ttl if following else 9000,
|
||||
"cacheScope": "private" if following else "public",
|
||||
"resultType": "complete",
|
||||
**({} if following else {"nextCursor": "next"}),
|
||||
}
|
||||
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": listing_result})
|
||||
|
||||
clock: Final = _ManualClockLoop()
|
||||
|
||||
class CleanupTransport(httpx2.AsyncBaseTransport):
|
||||
def __init__(self) -> None:
|
||||
self._transport = httpx2.MockTransport(respond)
|
||||
|
||||
async def handle_async_request(self, request: httpx2.Request) -> httpx2.Response:
|
||||
return await self._transport.handle_async_request(request)
|
||||
|
||||
async def aclose(self) -> None:
|
||||
await self._transport.aclose()
|
||||
if cleanup_phase == "http_client":
|
||||
clock.advance(cleanup_seconds)
|
||||
|
||||
@asynccontextmanager
|
||||
async def transport_with_cleanup(url: str, http_client: httpx2.AsyncClient) -> AsyncIterator[_TransportStreams]:
|
||||
async with streamable_http_client(url, http_client=http_client) as streams:
|
||||
yield streams
|
||||
if cleanup_phase == "transport":
|
||||
clock.advance(cleanup_seconds)
|
||||
|
||||
client: Final = _MockTransportClient(
|
||||
respond,
|
||||
http_transport=CleanupTransport(),
|
||||
transport_context=transport_with_cleanup,
|
||||
server_url="https://example.com/mcp",
|
||||
protocol_version="2026-07-28",
|
||||
)
|
||||
|
||||
try:
|
||||
with patch.object(mcp_client_module, "time", Mock(monotonic=clock.time)):
|
||||
result: Final = await getattr(client, "list_" + kind + "_result")(raise_on_error=True)
|
||||
finally:
|
||||
clock.close()
|
||||
assert len(getattr(result, kind)) == 2
|
||||
assert result.cache_scope == "private"
|
||||
assert result.next_cursor is None
|
||||
assert result.ttl_ms == max(0, ttl - cleanup_seconds * 1000)
|
||||
assert len(await getattr(client, "list_" + kind)(raise_on_error=True)) == 2
|
||||
|
|
|
|||
|
|
@ -230,14 +230,16 @@ async def test_failed_upstream_remains_visible_when_other_sources_continue(monke
|
|||
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
if server_id == "a":
|
||||
return ListToolsResult(tools=[], meta={SERVER_OUTCOMES_META_KEY: {"a": {"tag": "timeout"}}})
|
||||
return page("b2" if cursor else "b1", None if cursor else "next")
|
||||
return page("b2" if cursor else "b1", None if cursor else "next").model_copy(update={"ttl_ms": 9000})
|
||||
|
||||
first = await listing(fetch)
|
||||
assert [tool.name for tool in first.tools] == ["b1"]
|
||||
assert first.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
|
||||
assert first.ttl_ms == 0
|
||||
second = await listing(fetch, first.next_cursor)
|
||||
assert [tool.name for tool in second.tools] == ["b2"]
|
||||
assert second.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
|
||||
assert second.ttl_ms == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -473,6 +475,7 @@ async def test_optional_gateway_catalog_reports_initial_failure_and_rejects_fail
|
|||
assert fetch.await_args.args[-1] == "next"
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ["tools", "prompts", "resources", "templates"])
|
||||
async def test_first_catalog_page_keeps_admitted_items_and_reports_rate_limited_servers(
|
||||
|
|
@ -587,3 +590,84 @@ async def test_continuation_rate_limit_raises_without_refetching_completed_serve
|
|||
charged_server_ids: Final = [call.args[1].server_id for call in limiter.await_args_list]
|
||||
assert charged_server_ids.count("catalog-a") == 2
|
||||
assert charged_server_ids.count("catalog-b") == 1
|
||||
|
||||
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
|
||||
@pytest.mark.parametrize("ttls,expected", (((9000, 4000), 4000), ((9000, 0), 0)))
|
||||
def test_optional_catalog_preserves_conservative_freshness(kind: str, ttls: tuple[int, int], expected: int) -> None:
|
||||
from mcp.types import (
|
||||
ListPromptsRequest,
|
||||
ListPromptsResult,
|
||||
ListResourcesRequest,
|
||||
ListResourcesResult,
|
||||
ListResourceTemplatesRequest,
|
||||
ListResourceTemplatesResult,
|
||||
)
|
||||
|
||||
request, result_type, field = {
|
||||
"prompts": (ListPromptsRequest(), ListPromptsResult, "prompts"),
|
||||
"resources": (ListResourcesRequest(), ListResourcesResult, "resources"),
|
||||
"templates": (ListResourceTemplatesRequest(), ListResourceTemplatesResult, "resource_templates"),
|
||||
}[kind]
|
||||
pages: Final = tuple(result_type(**{field: []}, ttl_ms=ttl, cache_scope="public") for ttl in ttls)
|
||||
result: Final = catalog.combine_optional_catalog(request, pages, None, None)
|
||||
assert result.ttl_ms == expected
|
||||
assert result.cache_scope == "private"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_catalog_preserves_upstream_freshness() -> None:
|
||||
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id).model_copy(update={"ttl_ms": 9000, "cache_scope": "public"})
|
||||
|
||||
result: Final = await listing(fetch)
|
||||
assert 0 < result.ttl_ms <= 9000
|
||||
assert result.cache_scope == "private"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("upstream_ttl,limit,expected_ttl", [(1000, 60, 1000), (90000, 1, 1000), (0, 60, 0)])
|
||||
async def test_discovery_cache_uses_upstream_freshness_and_configured_cap(
|
||||
upstream_ttl: int, limit: float, expected_ttl: int
|
||||
) -> None:
|
||||
from unittest.mock import AsyncMock
|
||||
from mcp.types import ListPromptsResult, Prompt
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
class Clock:
|
||||
now: float = 0.0
|
||||
|
||||
def __call__(self) -> float:
|
||||
return self.now
|
||||
|
||||
clock: Final = Clock()
|
||||
cache: Final = catalog._DiscoveryCache(limit, clock, TypeAdapter(ListPromptsResult))
|
||||
fetch: Final = AsyncMock(return_value=ListPromptsResult(prompts=[Prompt(name="fresh")], ttl_ms=upstream_ttl))
|
||||
first: Final = await cache.get(("server", "caller"), fetch)
|
||||
assert first.prompts[0].name == "fresh"
|
||||
clock.now = 0.5 # rebind-ok: advance the injected test clock without sleeping
|
||||
second: Final = await cache.get(("server", "caller"), fetch)
|
||||
assert second.prompts[0].name == "fresh"
|
||||
if expected_ttl:
|
||||
assert fetch.await_count == 1
|
||||
assert second.ttl_ms == 500
|
||||
second.prompts[0].name = "caller edit" # rebind-ok: prove caller mutation cannot alter retained results
|
||||
else:
|
||||
assert fetch.await_count == 2
|
||||
clock.now = 1.0 # rebind-ok: reach the exact expiry boundary without sleeping
|
||||
assert (await cache.get(("server", "caller"), fetch)).prompts[0].name == "fresh"
|
||||
assert fetch.await_count == (2 if expected_ttl else 3)
|
||||
|
||||
|
||||
def test_partial_optional_catalog_never_advertises_freshness() -> None:
|
||||
from mcp.types import ListPromptsRequest, ListPromptsResult
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
|
||||
|
||||
result: Final = catalog.combine_optional_catalog(
|
||||
ListPromptsRequest(),
|
||||
[ListPromptsResult(prompts=[], ttl_ms=9000)],
|
||||
None,
|
||||
{SERVER_OUTCOMES_META_KEY: {"failed-earlier": {"tag": "timeout"}}},
|
||||
)
|
||||
assert result.ttl_ms == 0
|
||||
assert result.cache_scope == "private"
|
||||
|
||||
|
|
|
|||
|
|
@ -4102,7 +4102,7 @@ class TestMCPServerManager:
|
|||
|
||||
mock_prompt = Prompt(name="hello", description="Say hi")
|
||||
mock_client = AsyncMock()
|
||||
mock_client.list_prompts = AsyncMock(return_value=[mock_prompt])
|
||||
mock_client.list_prompts_result = AsyncMock(return_value=ListPromptsResult(prompts=[mock_prompt]))
|
||||
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
|
||||
|
||||
with patch.object(
|
||||
|
|
@ -4113,7 +4113,7 @@ class TestMCPServerManager:
|
|||
):
|
||||
prompts = await manager.get_prompts_from_server(server, user_api_key_auth=None, add_prefix=True)
|
||||
|
||||
mock_client.list_prompts.assert_awaited_once()
|
||||
mock_client.list_prompts_result.assert_awaited_once()
|
||||
assert len(prompts) == 1
|
||||
assert prompts[0].name == "alias-server-hello"
|
||||
|
||||
|
|
@ -4174,7 +4174,7 @@ class TestMCPServerManager:
|
|||
|
||||
mock_client = AsyncMock()
|
||||
mock_resources = [Resource(name="file", uri="https://example.com/file")]
|
||||
mock_client.list_resources = AsyncMock(return_value=mock_resources)
|
||||
mock_client.list_resources_result = AsyncMock(return_value=ListResourcesResult(resources=mock_resources))
|
||||
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
|
||||
prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")]
|
||||
|
||||
|
|
@ -4199,7 +4199,7 @@ class TestMCPServerManager:
|
|||
assert called_kwargs["server"] is server
|
||||
assert called_kwargs["mcp_auth_header"] == "auth"
|
||||
assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "static"}
|
||||
mock_client.list_resources.assert_awaited_once()
|
||||
mock_client.list_resources_result.assert_awaited_once()
|
||||
assert result == prefixed_resources
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -4222,7 +4222,9 @@ class TestMCPServerManager:
|
|||
uriTemplate="https://example.com/{id}",
|
||||
)
|
||||
]
|
||||
mock_client.list_resource_templates = AsyncMock(return_value=mock_templates)
|
||||
mock_client.list_resource_templates_result = AsyncMock(
|
||||
return_value=ListResourceTemplatesResult(resource_templates=mock_templates)
|
||||
)
|
||||
mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
|
||||
expected_templates = [
|
||||
ResourceTemplate(
|
||||
|
|
@ -4258,7 +4260,7 @@ class TestMCPServerManager:
|
|||
raw_headers=None,
|
||||
client_ip=None,
|
||||
)
|
||||
mock_client.list_resource_templates.assert_awaited_once()
|
||||
mock_client.list_resource_templates_result.assert_awaited_once()
|
||||
assert result == expected_templates
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -7459,10 +7461,10 @@ class TestMCPServerManager:
|
|||
metadata_key: Final = (new.server_id, new.url)
|
||||
prompt_fetches = 0
|
||||
|
||||
async def fetch_prompts() -> list[Prompt]:
|
||||
async def fetch_prompts() -> ListPromptsResult:
|
||||
nonlocal prompt_fetches
|
||||
prompt_fetches += 1
|
||||
return [Prompt(name="greet")]
|
||||
return ListPromptsResult(prompts=[Prompt(name="greet")], ttl_ms=60000)
|
||||
|
||||
async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None:
|
||||
manager.record_listed_tools(
|
||||
|
|
@ -7488,7 +7490,7 @@ class TestMCPServerManager:
|
|||
assert manager.registry["srv"] is new
|
||||
assert manager.get_listed_tool(new, "search", caller) is None
|
||||
prompts = await manager._prompt_discovery_cache.get((new.server_id, None), fetch_prompts)
|
||||
assert [prompt.name for prompt in prompts] == ["greet"]
|
||||
assert [prompt.name for prompt in prompts.prompts] == ["greet"]
|
||||
assert prompt_fetches == 1, "the prompts list filled after the save was published went upstream again"
|
||||
cached_metadata = discoverable_endpoints._OAUTH_METADATA_CACHE.get(metadata_key)
|
||||
assert cached_metadata is not None and cached_metadata[1] == {"resource": new.url}
|
||||
|
|
@ -14514,7 +14516,7 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken:
|
|||
manager: Final = MCPServerManager()
|
||||
client: Final = AsyncMock()
|
||||
client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False))
|
||||
client.list_prompts = AsyncMock(return_value=[])
|
||||
client.list_prompts_result = AsyncMock(return_value=ListPromptsResult(prompts=[]))
|
||||
client.read_resource = AsyncMock(return_value=ReadResourceResult(contents=[]))
|
||||
manager._create_mcp_client = AsyncMock(return_value=client)
|
||||
return manager
|
||||
|
|
@ -15190,7 +15192,7 @@ class _DiscoveryClock:
|
|||
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from mcp.types import JSONRPCMessage
|
||||
from mcp.types import JSONRPCMessage, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult
|
||||
|
||||
_JSONRPC_ADAPTER = TypeAdapter(JSONRPCMessage)
|
||||
|
||||
|
|
@ -15219,6 +15221,7 @@ class _DiscoveryUpstream:
|
|||
def __init__(self) -> None:
|
||||
self.requests: tuple[tuple[str, str], ...] = ()
|
||||
self.outcome = "supported"
|
||||
self.ttl_ms = 60000
|
||||
self.entered = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
self.release.set()
|
||||
|
|
@ -15232,6 +15235,21 @@ class _DiscoveryUpstream:
|
|||
if not isinstance(payload, JSONRPCRequest):
|
||||
return httpx2.Response(202)
|
||||
self.requests = (*self.requests, (payload.method, request.headers.get("authorization", "")))
|
||||
if payload.method == "server/discover":
|
||||
return httpx2.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload.id,
|
||||
"result": {
|
||||
"supportedVersions": ["2026-07-28"],
|
||||
"capabilities": {} if self.outcome == "unsupported" else {"prompts": {}, "resources": {}},
|
||||
"ttlMs": 0,
|
||||
"cacheScope": "private",
|
||||
"resultType": "complete",
|
||||
},
|
||||
},
|
||||
)
|
||||
if payload.method == "initialize":
|
||||
return httpx2.Response(
|
||||
200,
|
||||
|
|
@ -15270,16 +15288,33 @@ class _DiscoveryUpstream:
|
|||
if self.outcome in ("paged", "paged_failure") and not (payload.params or {}).get("cursor")
|
||||
else {}
|
||||
)
|
||||
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {**result, **continuation}})
|
||||
return httpx2.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload.id,
|
||||
"result": {
|
||||
**result,
|
||||
**continuation,
|
||||
"ttlMs": self.ttl_ms,
|
||||
"cacheScope": "private",
|
||||
"resultType": "complete",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
@property
|
||||
def initializes(self) -> int:
|
||||
return sum(method == "initialize" for method, _auth in self.requests)
|
||||
return sum(method in ("initialize", "server/discover") for method, _auth in self.requests)
|
||||
|
||||
|
||||
def _discovery_server() -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http
|
||||
server_id="discovery",
|
||||
name="discovery",
|
||||
url="https://discovery.example/mcp",
|
||||
transport=MCPTransport.http,
|
||||
protocol_version="2026-07-28",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -15306,7 +15341,7 @@ async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None
|
|||
assert second[0].name == "example"
|
||||
assert second[0].description == "original"
|
||||
assert upstream.initializes == 1
|
||||
clock.now = 59.999
|
||||
clock.now = 59.9
|
||||
assert (await operation(server, None))[0].name == "discovery-example"
|
||||
assert upstream.initializes == 1
|
||||
clock.now = 60.001
|
||||
|
|
@ -15331,7 +15366,7 @@ async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: st
|
|||
with _mcp_upstream(upstream.respond):
|
||||
assert await operation(_discovery_server(), None) == []
|
||||
assert await operation(_discovery_server(), None) == []
|
||||
assert upstream.initializes == (2 if outcome == "failure" else 1)
|
||||
assert upstream.initializes == 2
|
||||
if outcome == "failure":
|
||||
upstream.outcome = "supported"
|
||||
assert (await operation(_discovery_server(), None))[0].name == "discovery-example"
|
||||
|
|
@ -15362,7 +15397,7 @@ async def test_discovery_cache_retries_failed_pagination_before_caching_complete
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_auth() -> None:
|
||||
async def test_discovery_cache_isolates_forwarded_credentials_and_static_auth_callers() -> None:
|
||||
import respx
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
|
|
@ -15373,7 +15408,7 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_
|
|||
with _mcp_upstream(upstream.respond):
|
||||
for user in (first_user, second_user):
|
||||
assert len(await manager.get_prompts_from_server(server, user)) == 1
|
||||
assert upstream.initializes == 1
|
||||
assert upstream.initializes == 2
|
||||
for credential in ("first-secret", "second-secret", "first-secret"):
|
||||
assert (
|
||||
len(
|
||||
|
|
@ -15383,7 +15418,7 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_
|
|||
)
|
||||
== 1
|
||||
)
|
||||
assert upstream.initializes == 3
|
||||
assert upstream.initializes == 4
|
||||
assert {auth for method, auth in upstream.requests if method == "prompts/list"} == {
|
||||
"",
|
||||
"first-secret",
|
||||
|
|
@ -15391,6 +15426,28 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
|
||||
@pytest.mark.parametrize("header", ("Authorization", "X-LiteLLM-API-Key"))
|
||||
@pytest.mark.parametrize("identified", (False, True))
|
||||
async def test_discovery_cache_isolates_keyless_admission_credentials(
|
||||
kind: str, header: str, identified: bool
|
||||
) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
server: Final = _discovery_server().model_copy(update={"static_headers": {"Authorization": "Bearer upstream"}})
|
||||
auth: Final = UserAPIKeyAuth(team_id="shared-team", user_id="known-user" if identified else None)
|
||||
operation: Final = {
|
||||
"prompts": manager.get_prompts_from_server,
|
||||
"resources": manager.get_resources_from_server,
|
||||
"templates": manager.get_resource_templates_from_server,
|
||||
}[kind]
|
||||
with _mcp_upstream(upstream.respond):
|
||||
for credential in ("Bearer first", "Bearer second", "Bearer first"):
|
||||
assert len(await operation(server, auth, raw_headers={header: credential})) == 1
|
||||
assert upstream.initializes == (1 if identified else 2)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> None:
|
||||
import respx
|
||||
|
|
@ -15526,35 +15583,35 @@ async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefi
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_retries_cancelled_fetches() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
|
||||
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
|
||||
|
||||
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
|
||||
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
|
||||
|
||||
async def cancelled() -> list[Prompt]:
|
||||
async def cancelled() -> ListPromptsResult:
|
||||
raise asyncio.CancelledError()
|
||||
|
||||
async def supported() -> list[Prompt]:
|
||||
return [Prompt(name="recovered")]
|
||||
async def supported() -> ListPromptsResult:
|
||||
return ListPromptsResult(prompts=[Prompt(name="recovered")], ttl_ms=60000)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await cache.get(("server", None), cancelled)
|
||||
assert [item.name for item in await cache.get(("server", None), supported)] == ["recovered"]
|
||||
assert [item.name for item in (await cache.get(("server", None), supported)).prompts] == ["recovered"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_cancels_fetch_when_last_waiter_leaves() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
|
||||
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
|
||||
|
||||
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
|
||||
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
|
||||
entered: Final = asyncio.Event()
|
||||
stopped: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
async def fetch() -> list[Prompt]:
|
||||
async def fetch() -> ListPromptsResult:
|
||||
entered.set()
|
||||
try:
|
||||
await release.wait()
|
||||
return [Prompt(name="result")]
|
||||
return ListPromptsResult(prompts=[Prompt(name="result")], ttl_ms=60000)
|
||||
finally:
|
||||
stopped.set()
|
||||
|
||||
|
|
@ -15572,16 +15629,16 @@ async def test_discovery_cache_cancels_fetch_when_last_waiter_leaves() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_bounds_detached_fetches_without_dropping_results() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache
|
||||
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
|
||||
|
||||
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
|
||||
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
|
||||
entered: Final[asyncio.Queue[None]] = asyncio.Queue()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
async def blocked() -> list[Prompt]:
|
||||
async def blocked() -> ListPromptsResult:
|
||||
await entered.put(None)
|
||||
await release.wait()
|
||||
return [Prompt(name="blocked")]
|
||||
return ListPromptsResult(prompts=[Prompt(name="blocked")], ttl_ms=60000)
|
||||
|
||||
tasks: Final = tuple(asyncio.create_task(cache.get((str(index), None), blocked)) for index in range(1024))
|
||||
try:
|
||||
|
|
@ -15589,16 +15646,16 @@ async def test_discovery_cache_bounds_detached_fetches_without_dropping_results(
|
|||
await asyncio.wait_for(entered.get(), timeout=5)
|
||||
active_tasks: Final = frozenset(asyncio.all_tasks())
|
||||
|
||||
async def overflow() -> list[Prompt]:
|
||||
async def overflow() -> ListPromptsResult:
|
||||
assert frozenset(asyncio.all_tasks()) <= active_tasks
|
||||
return [Prompt(name="overflow")]
|
||||
return ListPromptsResult(prompts=[Prompt(name="overflow")], ttl_ms=60000)
|
||||
|
||||
result: Final = await cache.get(("overflow", None), overflow)
|
||||
assert [item.name for item in result] == ["overflow"]
|
||||
assert [item.name for item in result.prompts] == ["overflow"]
|
||||
finally:
|
||||
release.set()
|
||||
outcomes: Final = await asyncio.gather(*tasks)
|
||||
assert all(result[0].name == "blocked" for result in outcomes)
|
||||
assert all(result.prompts[0].name == "blocked" for result in outcomes)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -15628,6 +15685,7 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N
|
|||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="authorization_code",
|
||||
protocol_version="2026-07-28",
|
||||
client_id="discovery-client",
|
||||
authorization_url="https://discovery.example/authorize",
|
||||
token_url="https://discovery.example/token",
|
||||
|
|
@ -15644,16 +15702,28 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N
|
|||
payload: Final = _JSONRPC_ADAPTER.validate_json(request.content)
|
||||
assert isinstance(payload, JSONRPCRequest)
|
||||
name: Final = {"Bearer token-a": "account-a", "Bearer token-b": "account-b"}[request.headers["authorization"]]
|
||||
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"prompts": [{"name": name}]}})
|
||||
return httpx2.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload.id,
|
||||
"result": {
|
||||
"prompts": [{"name": name}],
|
||||
"ttlMs": 60000,
|
||||
"cacheScope": "private",
|
||||
"resultType": "complete",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
with _mcp_upstream(respond):
|
||||
for manager in managers:
|
||||
for manager in (*managers, *managers):
|
||||
assert [item.name for item in await manager.get_prompts_from_server(server, user)] == [
|
||||
"discovery-account-a"
|
||||
]
|
||||
assert upstream.initializes == 2
|
||||
source.token = "token-b"
|
||||
for manager in managers:
|
||||
for manager in (*managers, *managers):
|
||||
assert [item.name for item in await manager.get_prompts_from_server(server, user)] == [
|
||||
"discovery-account-b"
|
||||
]
|
||||
|
|
@ -15699,69 +15769,71 @@ async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None
|
|||
assert len(await manager.get_prompts_from_server(server, user)) == 1
|
||||
assert len(await manager.get_prompts_from_server(server, user)) == 1
|
||||
assert store.calls == (("requesting-user", "discovery"), ("requesting-user", "discovery"))
|
||||
assert upstream.initializes == 1
|
||||
assert upstream.initializes == 2
|
||||
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
|
||||
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
|
||||
|
||||
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
|
||||
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
|
||||
|
||||
async def original() -> list[Prompt]:
|
||||
return [Prompt(name="original")]
|
||||
async def original() -> ListPromptsResult:
|
||||
return ListPromptsResult(prompts=[Prompt(name="original")], ttl_ms=60000)
|
||||
|
||||
async def refetched() -> list[Prompt]:
|
||||
return [Prompt(name="refetched")]
|
||||
async def refetched() -> ListPromptsResult:
|
||||
return ListPromptsResult(prompts=[Prompt(name="refetched")], ttl_ms=60000)
|
||||
|
||||
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"
|
||||
assert (await cache.get((f"server-{index:04}", None), original)).prompts[0].name == "original"
|
||||
assert (await cache.get(("server-1024", None), refetched)).prompts[0].name == "original"
|
||||
assert (await cache.get(("server-0000", None), refetched)).prompts[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
|
||||
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
|
||||
|
||||
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
|
||||
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
|
||||
entered: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
async def original() -> list[Prompt]:
|
||||
return [Prompt(name="original")]
|
||||
async def original() -> ListPromptsResult:
|
||||
return ListPromptsResult(prompts=[Prompt(name="original")], ttl_ms=60000)
|
||||
|
||||
async def blocked() -> list[Prompt]:
|
||||
async def blocked() -> ListPromptsResult:
|
||||
entered.set()
|
||||
await release.wait()
|
||||
return [Prompt(name="pending")]
|
||||
return ListPromptsResult(prompts=[Prompt(name="pending")], ttl_ms=60000)
|
||||
|
||||
async def refetched() -> list[Prompt]:
|
||||
return [Prompt(name="refetched")]
|
||||
async def refetched() -> ListPromptsResult:
|
||||
return ListPromptsResult(prompts=[Prompt(name="refetched")], ttl_ms=60000)
|
||||
|
||||
assert (await cache.get(("server", None), original))[0].name == "original"
|
||||
assert (await cache.get(("server-extra", None), original))[0].name == "original"
|
||||
assert (await cache.get(("server", None), original)).prompts[0].name == "original"
|
||||
assert (await cache.get(("server-extra", None), original)).prompts[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"
|
||||
assert (await asyncio.wait_for(task, timeout=5)).prompts[0].name == "pending"
|
||||
assert (await cache.get(("other", None), refetched)).prompts[0].name == "pending"
|
||||
assert (await cache.get(("server-extra", None), refetched)).prompts[0].name == "original"
|
||||
assert (await cache.get(("server", None), refetched)).prompts[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
|
||||
from litellm.proxy._experimental.mcp_server.catalog import _DiscoveryCache
|
||||
|
||||
cache: Final = _DiscoveryCache[Prompt](60, _DiscoveryClock(), TypeAdapter(tuple[Prompt, ...]))
|
||||
fetch: Final = AsyncMock(return_value=[Prompt(name="large", description=description)])
|
||||
cache: Final = _DiscoveryCache[ListPromptsResult](60, _DiscoveryClock(), TypeAdapter(ListPromptsResult))
|
||||
fetch: Final = AsyncMock(
|
||||
return_value=ListPromptsResult(prompts=[Prompt(name="large", description=description)], ttl_ms=60000)
|
||||
)
|
||||
for _ in range(2):
|
||||
result: Final = await cache.get(("server", None), fetch)
|
||||
assert result[0].description == description
|
||||
assert result.prompts[0].description == description
|
||||
assert fetch.await_count == 2
|
||||
|
||||
|
||||
|
|
@ -19005,3 +19077,34 @@ async def test_aggregate_publishes_complete_bare_routes_only_after_delivering_a_
|
|||
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
|
||||
await listing
|
||||
assert manager._get_mcp_server_from_tool_name("first") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
|
||||
async def test_discovery_does_not_retain_unknown_freshness(kind: str) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
upstream.ttl_ms = 0
|
||||
operation: Final = {
|
||||
"prompts": manager.get_prompts_from_server,
|
||||
"resources": manager.get_resources_from_server,
|
||||
"templates": manager.get_resource_templates_from_server,
|
||||
}[kind]
|
||||
with _mcp_upstream(upstream.respond):
|
||||
assert len(await operation(_discovery_server(), None)) == 1
|
||||
assert len(await operation(_discovery_server(), None)) == 1
|
||||
assert upstream.initializes == 2
|
||||
|
||||
|
||||
def test_discovery_keys_bind_static_auth_to_caller_and_configuration() -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = _discovery_server()
|
||||
first: Final = UserAPIKeyAuth(user_id="first", team_id="one")
|
||||
second: Final = UserAPIKeyAuth(user_id="second", team_id="two")
|
||||
updated: Final = server.model_copy(update={"url": "https://replacement.example/mcp"})
|
||||
keys: Final = (
|
||||
manager._discovery_key(server, first, None, None, None, None),
|
||||
manager._discovery_key(server, second, None, None, None, None),
|
||||
manager._discovery_key(updated, first, None, None, None, None),
|
||||
)
|
||||
assert len(set(keys)) == 3
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import json
|
||||
from typing import Final
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from mcp.types import CallToolResult, ImageContent, InputRequiredResult, TextContent, Tool
|
||||
|
|
@ -241,3 +241,30 @@ class TestToGatewayTool:
|
|||
assert renamed.input_schema == tool.input_schema and renamed.input_schema is not tool.input_schema
|
||||
assert renamed.meta == {"owner": "x"}
|
||||
assert renamed.description == "d"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("elapsed,expected", [(0, 1000), (0.0001, 999), (0.5, 500), (1, 0), (2, 0), (-1, 1000)])
|
||||
def test_freshness_aging_preserves_content_and_scope(elapsed: float, expected: int) -> None:
|
||||
from mcp.types import ListPromptsResult, Prompt
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import age_freshness
|
||||
|
||||
result: Final = ListPromptsResult(prompts=[Prompt(name="kept")], ttl_ms=1000, cache_scope="public")
|
||||
aged: Final = age_freshness(result, elapsed)
|
||||
assert aged.ttl_ms == expected
|
||||
assert aged.cache_scope == "public"
|
||||
assert aged.prompts == result.prompts
|
||||
assert result.ttl_ms == 1000
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"scopes,expected", [((), "private"), (("public", "public"), "public"), (("public", "private"), "private")]
|
||||
)
|
||||
def test_aggregate_freshness_never_broadens_sharing(
|
||||
scopes: tuple[Literal["private", "public"], ...], expected: str
|
||||
) -> None:
|
||||
from mcp.types import CacheableResult
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import aggregate_freshness
|
||||
|
||||
result: Final = aggregate_freshness(tuple(CacheableResult(ttl_ms=1000, cache_scope=scope) for scope in scopes))
|
||||
assert result.cache_scope == expected
|
||||
assert result.ttl_ms == (1000 if scopes else 0)
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ _IN_MEMORY_ONLY_CALLERS: Final = frozenset(
|
|||
"litellm/llms/vertex_ai/vertex_ai_non_gemini.py",
|
||||
"litellm/llms/watsonx/common_utils.py",
|
||||
"litellm/proxy/_experimental/mcp_server/byok_credential_cache.py",
|
||||
"litellm/proxy/_experimental/mcp_server/catalog.py",
|
||||
"litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py",
|
||||
"litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py",
|
||||
"litellm/proxy/_experimental/mcp_server/operations.py",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue