mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): cache upstream discovery lists
This commit is contained in:
parent
d51a7af655
commit
c038aaf622
5 changed files with 451 additions and 98 deletions
|
|
@ -2,5 +2,10 @@
|
|||
|
||||
LiteLLM MCP Client is a client that allows you to use MCP tools with LiteLLM.
|
||||
|
||||
## Gateway discovery caching
|
||||
|
||||
The MCP gateway caches each upstream server's prompt, resource, and resource-template lists for 60 seconds per worker. Set `LITELLM_MCP_DISCOVERY_CACHE_TTL` to a nonnegative number of seconds to change the lifetime, or `0` to disable caching. Invalid values use the 60-second default
|
||||
|
||||
Discovery results may remain unchanged until that lifetime expires. Server configuration updates invalidate the affected server's entries. Concurrent requests for the same list share one upstream fetch. Each list cache holds at most 1,024 entries per worker
|
||||
|
||||
User-dependent upstream authentication uses separate cache entries. Gateway access checks still run for every request. Successful empty lists and unsupported capabilities are cached; failed requests retain the existing empty-list response and are retried on the next request
|
||||
|
|
|
|||
|
|
@ -781,7 +781,7 @@ class MCPClient:
|
|||
# Return a default error result instead of raising
|
||||
return self.error_tool_result(e)
|
||||
|
||||
async def list_prompts(self) -> list[Prompt]:
|
||||
async def list_prompts(self, *, raise_on_error: bool = False) -> list[Prompt]:
|
||||
"""List available prompts from the server."""
|
||||
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
|
||||
|
||||
|
|
@ -811,6 +811,8 @@ class MCPClient:
|
|||
verbose_logger.warning("MCP client list_prompts was cancelled")
|
||||
raise
|
||||
except Exception as e:
|
||||
if raise_on_error:
|
||||
raise
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client list_prompts failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
|
||||
|
|
@ -869,7 +871,7 @@ class MCPClient:
|
|||
)
|
||||
raise
|
||||
|
||||
async def list_resources(self) -> list[Resource]:
|
||||
async def list_resources(self, *, raise_on_error: bool = False) -> list[Resource]:
|
||||
"""List available resources from the server."""
|
||||
verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio")
|
||||
|
||||
|
|
@ -899,6 +901,8 @@ class MCPClient:
|
|||
verbose_logger.warning("MCP client list_resources was cancelled")
|
||||
raise
|
||||
except Exception as e:
|
||||
if raise_on_error:
|
||||
raise
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client list_resources failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
|
||||
|
|
@ -916,7 +920,7 @@ class MCPClient:
|
|||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
||||
async def list_resource_templates(self) -> list[ResourceTemplate]:
|
||||
async def list_resource_templates(self, *, raise_on_error: bool = False) -> list[ResourceTemplate]:
|
||||
"""List available resource templates from the server."""
|
||||
verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio")
|
||||
|
||||
|
|
@ -949,6 +953,8 @@ class MCPClient:
|
|||
verbose_logger.warning("MCP client list_resource_templates was cancelled")
|
||||
raise
|
||||
except Exception as e:
|
||||
if raise_on_error:
|
||||
raise
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client list_resource_templates failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import asyncio
|
|||
import datetime
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
|
|
@ -26,8 +27,9 @@ from collections.abc import (
|
|||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass, replace
|
||||
from functools import lru_cache
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
|
||||
from urllib.parse import ParseResult, urlparse
|
||||
|
||||
import anyio
|
||||
|
|
@ -1677,6 +1679,91 @@ 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]
|
||||
|
||||
|
||||
@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({})
|
||||
|
||||
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})
|
||||
|
||||
@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 = self._entries.get(key)
|
||||
if entry is not None and entry.expires_at > self._clock():
|
||||
return tuple(item.model_copy(deep=True) for item in entry.items)
|
||||
pending: Final = self._pending.get(key)
|
||||
if pending is not None:
|
||||
return tuple(item.model_copy(deep=True) for item in await asyncio.shield(pending))
|
||||
task: Final = asyncio.create_task(self._fetch(key, fetch))
|
||||
self._pending = MappingProxyType({**self._pending, key: task})
|
||||
task.add_done_callback(self._observe_completion)
|
||||
return tuple(item.model_copy(deep=True) for item in await asyncio.shield(task))
|
||||
|
||||
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():
|
||||
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[-1023:],
|
||||
(
|
||||
key,
|
||||
_DiscoveryEntry(now + self._ttl, tuple(item.model_copy(deep=True) for item in items)),
|
||||
),
|
||||
)
|
||||
}
|
||||
)
|
||||
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}
|
||||
)
|
||||
|
||||
|
||||
def _mcp_discovery_cache_ttl() -> float:
|
||||
raw: Final = os.environ.get("LITELLM_MCP_DISCOVERY_CACHE_TTL", "60")
|
||||
try:
|
||||
ttl: Final = float(raw)
|
||||
except ValueError:
|
||||
verbose_logger.warning("Invalid LITELLM_MCP_DISCOVERY_CACHE_TTL; using 60 seconds")
|
||||
return 60.0
|
||||
if not math.isfinite(ttl) or ttl < 0:
|
||||
verbose_logger.warning("Invalid LITELLM_MCP_DISCOVERY_CACHE_TTL; using 60 seconds")
|
||||
return 60.0
|
||||
return ttl
|
||||
|
||||
|
||||
class MCPServerManager:
|
||||
_STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$")
|
||||
|
||||
|
|
@ -1793,6 +1880,7 @@ class MCPServerManager:
|
|||
cred_provider: UpstreamCredentialProvider | None = None,
|
||||
per_user_oauth_token_store: InvalidatableOAuthTokenStore | None = None,
|
||||
per_user_token_cache: MCPPerUserTokenCache | None = None,
|
||||
discovery_clock: Callable[[], float] = time.monotonic,
|
||||
):
|
||||
self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore(
|
||||
self.get_mcp_server_by_id
|
||||
|
|
@ -1802,6 +1890,10 @@ class MCPServerManager:
|
|||
oauth_token_store=self._per_user_oauth_token_store,
|
||||
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.registry: dict[str, MCPServer] = {}
|
||||
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
|
||||
self.config_mcp_servers: dict[str, MCPServer] = {}
|
||||
|
|
@ -2529,6 +2621,7 @@ class MCPServerManager:
|
|||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
|
||||
_warn_config_id_jag_server_outruns_sso(new_server)
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
self._set_oauth_discovery_deferred(
|
||||
server_id,
|
||||
|
|
@ -2730,6 +2823,7 @@ class MCPServerManager:
|
|||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server.server_id)
|
||||
prefix_root: Final = normalize_server_name(get_server_prefix(server))
|
||||
if server.spec_path and prefix_root:
|
||||
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
|
||||
|
|
@ -3106,6 +3200,7 @@ class MCPServerManager:
|
|||
# env_vars_are_encrypted=False.
|
||||
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -3142,6 +3237,7 @@ class MCPServerManager:
|
|||
previous_server=self.registry[mcp_server.server_id],
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -4375,6 +4471,38 @@ class MCPServerManager:
|
|||
)
|
||||
raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge)
|
||||
|
||||
def _invalidate_discovery_lists(self, server_id: str) -> None:
|
||||
self._prompt_discovery_cache.invalidate(server_id)
|
||||
self._resource_discovery_cache.invalidate(server_id)
|
||||
self._template_discovery_cache.invalidate(server_id)
|
||||
|
||||
def _discovery_key(
|
||||
self,
|
||||
server: MCPServer,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
extra_headers: dict[str, str] | None,
|
||||
stdio_env: dict[str, str] | None,
|
||||
subject_token: str | None,
|
||||
) -> _DiscoveryKey:
|
||||
per_user: Final = (
|
||||
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
|
||||
else None
|
||||
)
|
||||
material: Final = json.dumps(
|
||||
(identity, mcp_auth_header, extra_headers, stdio_env, subject_token), sort_keys=True, separators=(",", ":")
|
||||
)
|
||||
return server.server_id, hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
async def get_prompts_from_server(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -4384,47 +4512,36 @@ class MCPServerManager:
|
|||
add_prefix: bool = True,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
) -> list[Prompt]:
|
||||
"""
|
||||
Helper method to get prompts from a single MCP server with prefixed names.
|
||||
|
||||
Args:
|
||||
server (MCPServer): The server to query prompts from
|
||||
mcp_auth_header: Optional auth header for MCP server
|
||||
|
||||
Returns:
|
||||
List[Prompt]: List of prompts available on the server with prefixed names
|
||||
"""
|
||||
|
||||
verbose_logger.debug("Connecting to url: %s", server.url)
|
||||
verbose_logger.info("get_prompts_from_server for %s...", server.name)
|
||||
|
||||
client = None
|
||||
|
||||
try:
|
||||
if server.static_headers:
|
||||
if extra_headers is None:
|
||||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
headers: Final = (
|
||||
dict(
|
||||
chain(
|
||||
extra_headers.items() if extra_headers else (),
|
||||
server.static_headers.items() if server.static_headers else (),
|
||||
)
|
||||
)
|
||||
or None
|
||||
)
|
||||
stdio_env: Final = self._build_stdio_env(server, raw_headers)
|
||||
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
key: Final = self._discovery_key(
|
||||
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token
|
||||
)
|
||||
|
||||
prompts: Final = await client.list_prompts()
|
||||
async def fetch() -> list[Prompt]:
|
||||
client: Final = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
return await client.list_prompts(raise_on_error=True)
|
||||
|
||||
prefixed_or_original_prompts: Final = self._create_prefixed_prompts(prompts, server, add_prefix=add_prefix)
|
||||
|
||||
return prefixed_or_original_prompts
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, e)
|
||||
items: Final = await self._prompt_discovery_cache.get(key, fetch)
|
||||
return self._create_prefixed_prompts(items, server, add_prefix=add_prefix)
|
||||
except Exception as error:
|
||||
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error)
|
||||
return []
|
||||
|
||||
async def get_resources_from_server(
|
||||
|
|
@ -4436,38 +4553,36 @@ class MCPServerManager:
|
|||
add_prefix: bool = True,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
) -> list[Resource]:
|
||||
"""Fetch available resources from a single MCP server."""
|
||||
|
||||
verbose_logger.debug("Connecting to url: %s", server.url)
|
||||
verbose_logger.info("get_resources_from_server for %s...", server.name)
|
||||
|
||||
client = None
|
||||
|
||||
try:
|
||||
if server.static_headers:
|
||||
if extra_headers is None:
|
||||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
headers: Final = (
|
||||
dict(
|
||||
chain(
|
||||
extra_headers.items() if extra_headers else (),
|
||||
server.static_headers.items() if server.static_headers else (),
|
||||
)
|
||||
)
|
||||
or None
|
||||
)
|
||||
stdio_env: Final = self._build_stdio_env(server, raw_headers)
|
||||
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
key: Final = self._discovery_key(
|
||||
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token
|
||||
)
|
||||
|
||||
resources: Final = await client.list_resources()
|
||||
async def fetch() -> list[Resource]:
|
||||
client: Final = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
return await client.list_resources(raise_on_error=True)
|
||||
|
||||
prefixed_resources: Final = self._create_prefixed_resources(resources, server, add_prefix=add_prefix)
|
||||
|
||||
return prefixed_resources
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, e)
|
||||
items: Final = await self._resource_discovery_cache.get(key, fetch)
|
||||
return self._create_prefixed_resources(items, server, add_prefix=add_prefix)
|
||||
except Exception as error:
|
||||
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error)
|
||||
return []
|
||||
|
||||
async def get_resource_templates_from_server(
|
||||
|
|
@ -4479,40 +4594,36 @@ class MCPServerManager:
|
|||
add_prefix: bool = True,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
) -> list[ResourceTemplate]:
|
||||
"""Fetch available resource templates from a single MCP server."""
|
||||
|
||||
verbose_logger.debug("Connecting to url: %s", server.url)
|
||||
verbose_logger.info("get_resource_templates_from_server for %s...", server.name)
|
||||
|
||||
client = None
|
||||
|
||||
try:
|
||||
if server.static_headers:
|
||||
if extra_headers is None:
|
||||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
headers: Final = (
|
||||
dict(
|
||||
chain(
|
||||
extra_headers.items() if extra_headers else (),
|
||||
server.static_headers.items() if server.static_headers else (),
|
||||
)
|
||||
)
|
||||
or None
|
||||
)
|
||||
stdio_env: Final = self._build_stdio_env(server, raw_headers)
|
||||
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
key: Final = self._discovery_key(
|
||||
server, user_api_key_auth, mcp_auth_header, headers, stdio_env, subject_token
|
||||
)
|
||||
|
||||
resource_templates: Final = await client.list_resource_templates()
|
||||
async def fetch() -> list[ResourceTemplate]:
|
||||
client: Final = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
return await client.list_resource_templates(raise_on_error=True)
|
||||
|
||||
prefixed_templates: Final = self._create_prefixed_resource_templates(
|
||||
resource_templates, server, add_prefix=add_prefix
|
||||
)
|
||||
|
||||
return prefixed_templates
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to get resource templates from server %s: %s", server.name, e)
|
||||
items: Final = await self._template_discovery_cache.get(key, fetch)
|
||||
return self._create_prefixed_resource_templates(items, 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 []
|
||||
|
||||
async def read_resource_from_server(
|
||||
|
|
@ -5220,7 +5331,7 @@ class MCPServerManager:
|
|||
return prefixed_tools
|
||||
|
||||
def _create_prefixed_prompts(
|
||||
self, prompts: list[Prompt], server: MCPServer, add_prefix: bool = True
|
||||
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
|
||||
) -> list[Prompt]:
|
||||
"""
|
||||
Create prefixed prompts and update prompt mapping.
|
||||
|
|
@ -5247,7 +5358,7 @@ class MCPServerManager:
|
|||
return prefixed_prompts
|
||||
|
||||
def _create_prefixed_resources(
|
||||
self, resources: list[Resource], server: MCPServer, add_prefix: bool = True
|
||||
self, resources: Sequence[Resource], server: MCPServer, add_prefix: bool = True
|
||||
) -> list[Resource]:
|
||||
"""Prefix resource names and track origin server for read requests."""
|
||||
|
||||
|
|
@ -5264,7 +5375,7 @@ class MCPServerManager:
|
|||
|
||||
def _create_prefixed_resource_templates(
|
||||
self,
|
||||
resource_templates: list[ResourceTemplate],
|
||||
resource_templates: Sequence[ResourceTemplate],
|
||||
server: MCPServer,
|
||||
add_prefix: bool = True,
|
||||
) -> list[ResourceTemplate]:
|
||||
|
|
@ -6466,6 +6577,9 @@ class MCPServerManager:
|
|||
for registry_key in dropped_registry_keys:
|
||||
self._invalidate_oauth_discovery_state(previous_registry[registry_key].server_id)
|
||||
|
||||
for server_id in previous_registry.keys() | registered_registry.keys():
|
||||
if previous_registry.get(server_id) != registered_registry.get(server_id):
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self.registry = registered_registry
|
||||
# A discovery task may have published into ``previous_registry`` while
|
||||
# this replacement was being staged. Reconcile every published entry
|
||||
|
|
|
|||
|
|
@ -1704,8 +1704,9 @@ async def test_empty_http_event_stream_uses_the_existing_request_deadline() -> N
|
|||
"initialize_not_found",
|
||||
),
|
||||
)
|
||||
@pytest.mark.parametrize("raise_on_error", (False, True))
|
||||
async def test_optional_discovery_capabilities_and_errors(
|
||||
method: str, outcome: str, caplog: pytest.LogCaptureFixture
|
||||
method: str, outcome: str, caplog: pytest.LogCaptureFixture, raise_on_error: bool
|
||||
) -> None:
|
||||
import logging
|
||||
from unittest.mock import Mock
|
||||
|
|
@ -1783,7 +1784,11 @@ async def test_optional_discovery_capabilities_and_errors(
|
|||
"resources/list": client.list_resources,
|
||||
"resources/templates/list": client.list_resource_templates,
|
||||
}[method]
|
||||
result: Final = await operation()
|
||||
if raise_on_error and outcome in ("internal_error", "unauthorized", "timeout", "initialize_not_found"):
|
||||
with pytest.raises((McpError, httpx.HTTPError)):
|
||||
await operation(raise_on_error=True)
|
||||
return
|
||||
result: Final = await operation(raise_on_error=raise_on_error)
|
||||
|
||||
requests: Final = tuple(
|
||||
JSONRPCMessage.model_validate_json(call.args[0].content).root
|
||||
|
|
|
|||
|
|
@ -13020,3 +13020,226 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon
|
|||
assert cached.status == "healthy"
|
||||
assert len(attempts) == 2
|
||||
assert route.call_count == 1
|
||||
|
||||
|
||||
class _DiscoveryClock:
|
||||
def __init__(self) -> None:
|
||||
self.now = 0.0
|
||||
|
||||
def __call__(self) -> float:
|
||||
return self.now
|
||||
|
||||
|
||||
class _DiscoveryUpstream:
|
||||
def __init__(self) -> None:
|
||||
self.requests: tuple[tuple[str, str], ...] = ()
|
||||
self.outcome = "supported"
|
||||
self.entered = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
self.release.set()
|
||||
|
||||
async def respond(self, request: httpx.Request) -> httpx.Response:
|
||||
from mcp.types import JSONRPCMessage, JSONRPCRequest
|
||||
|
||||
if request.method == "DELETE":
|
||||
return httpx.Response(200)
|
||||
payload: Final = JSONRPCMessage.model_validate_json(request.content).root
|
||||
if not isinstance(payload, JSONRPCRequest):
|
||||
return httpx.Response(202)
|
||||
self.requests = (*self.requests, (payload.method, request.headers.get("authorization", "")))
|
||||
if payload.method == "initialize":
|
||||
return httpx.Response(200, json={
|
||||
"jsonrpc": "2.0", "id": payload.id,
|
||||
"result": {"protocolVersion": "2025-03-26", "serverInfo": {"name": "discovery", "version": "1"},
|
||||
"capabilities": {} if self.outcome == "unsupported" else {"prompts": {}, "resources": {}}},
|
||||
})
|
||||
self.entered.set()
|
||||
await self.release.wait()
|
||||
if self.outcome == "failure":
|
||||
return httpx.Response(503)
|
||||
if self.outcome == "cancelled":
|
||||
raise asyncio.CancelledError()
|
||||
if self.outcome == "rejected":
|
||||
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id,
|
||||
"error": {"code": -32601, "message": "Unsupported"}})
|
||||
result: Final = {
|
||||
"prompts/list": {"prompts": [{"name": "example", "description": "original"}]},
|
||||
"resources/list": {"resources": [{"name": "example", "uri": "test://example", "description": "original"}]},
|
||||
"resources/templates/list": {"resourceTemplates": [{"name": "example", "uriTemplate": "test://{name}", "description": "original"}]},
|
||||
"tools/list": {"tools": []},
|
||||
}[payload.method]
|
||||
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result})
|
||||
|
||||
@property
|
||||
def initializes(self) -> int:
|
||||
return sum(method == "initialize" 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)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
|
||||
async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None:
|
||||
import respx
|
||||
|
||||
clock: Final = _DiscoveryClock()
|
||||
manager: Final = MCPServerManager(discovery_clock=clock)
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server,
|
||||
"templates": manager.get_resource_templates_from_server}[kind]
|
||||
server: Final = _discovery_server()
|
||||
with respx.mock(base_url="https://discovery.example") as router:
|
||||
router.route().mock(side_effect=upstream.respond)
|
||||
first: Final = await operation(server, None)
|
||||
assert len(first) == 1
|
||||
assert first[0].name == "discovery-example"
|
||||
first[0].description = "caller changed it"
|
||||
second: Final = await operation(server, None, add_prefix=False)
|
||||
assert second[0].name == "example"
|
||||
assert second[0].description == "original"
|
||||
assert upstream.initializes == 1
|
||||
clock.now = 59.999
|
||||
assert (await operation(server, None))[0].name == "discovery-example"
|
||||
assert upstream.initializes == 1
|
||||
clock.now = 60.0
|
||||
assert (await operation(server, None))[0].name == "discovery-example"
|
||||
assert upstream.initializes == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ("prompts", "resources", "templates"))
|
||||
@pytest.mark.parametrize("outcome", ("unsupported", "rejected", "failure"))
|
||||
async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: str) -> None:
|
||||
import respx
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
upstream.outcome = outcome
|
||||
operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server,
|
||||
"templates": manager.get_resource_templates_from_server}[kind]
|
||||
with respx.mock(base_url="https://discovery.example") as router:
|
||||
router.route().mock(side_effect=upstream.respond)
|
||||
assert await operation(_discovery_server(), None) == []
|
||||
assert await operation(_discovery_server(), None) == []
|
||||
assert upstream.initializes == (2 if outcome == "failure" else 1)
|
||||
if outcome == "failure":
|
||||
upstream.outcome = "supported"
|
||||
assert (await operation(_discovery_server(), None))[0].name == "discovery-example"
|
||||
assert upstream.initializes == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_auth() -> None:
|
||||
import respx
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
server: Final = _discovery_server()
|
||||
first_user: Final = UserAPIKeyAuth(user_id="first")
|
||||
second_user: Final = UserAPIKeyAuth(user_id="second")
|
||||
with respx.mock(base_url="https://discovery.example") as router:
|
||||
router.route().mock(side_effect=upstream.respond)
|
||||
for user in (first_user, second_user):
|
||||
assert len(await manager.get_prompts_from_server(server, user)) == 1
|
||||
assert upstream.initializes == 1
|
||||
for credential in ("first-secret", "second-secret", "first-secret"):
|
||||
assert len(await manager.get_prompts_from_server(server, first_user, extra_headers={"Authorization": credential})) == 1
|
||||
assert upstream.initializes == 3
|
||||
assert {auth for method, auth in upstream.requests if method == "prompts/list"} == {"", "first-secret", "second-secret"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> None:
|
||||
import respx
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
upstream.release.clear()
|
||||
with respx.mock(base_url="https://discovery.example") as router:
|
||||
router.route().mock(side_effect=upstream.respond)
|
||||
tasks: Final = tuple(asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None)) for _ in range(10))
|
||||
await asyncio.wait_for(upstream.entered.wait(), timeout=5)
|
||||
tasks[0].cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await tasks[0]
|
||||
upstream.release.set()
|
||||
results: Final = await asyncio.wait_for(asyncio.gather(*tasks[1:]), timeout=5)
|
||||
assert all(result[0].name == "discovery-example" for result in results)
|
||||
assert upstream.initializes == 1
|
||||
assert results[0][0] is not results[1][0]
|
||||
assert (await manager.get_prompts_from_server(_discovery_server(), None))[0].name == "discovery-example"
|
||||
assert upstream.initializes == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old_results() -> None:
|
||||
import respx
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
upstream.release.clear()
|
||||
with respx.mock(base_url="https://discovery.example") as router:
|
||||
router.route().mock(side_effect=upstream.respond)
|
||||
task: Final = asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None))
|
||||
await asyncio.wait_for(upstream.entered.wait(), timeout=5)
|
||||
manager._invalidate_discovery_lists("discovery")
|
||||
upstream.release.set()
|
||||
assert (await task)[0].name == "discovery-example"
|
||||
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
|
||||
assert upstream.initializes == 2
|
||||
manager._invalidate_discovery_lists("discovery")
|
||||
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
|
||||
assert upstream.initializes == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
import respx
|
||||
|
||||
monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", "0")
|
||||
manager: Final = MCPServerManager()
|
||||
upstream: Final = _DiscoveryUpstream()
|
||||
with respx.mock(base_url="https://discovery.example") as router:
|
||||
router.route().mock(side_effect=upstream.respond)
|
||||
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
|
||||
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
|
||||
assert upstream.initializes == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)))
|
||||
def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl
|
||||
|
||||
monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", value)
|
||||
assert _mcp_discovery_cache_ttl() == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize("auth_type", (MCPAuth.oauth2, MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag))
|
||||
def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = _discovery_server().model_copy(update={"auth_type": auth_type})
|
||||
first: Final = manager._discovery_key(server, UserAPIKeyAuth(user_id="first"), None, None, None, None)
|
||||
second: Final = manager._discovery_key(server, UserAPIKeyAuth(user_id="second"), None, None, None, None)
|
||||
anonymous: Final = manager._discovery_key(server, None, None, None, None, None)
|
||||
assert len({first, second, anonymous}) == 3
|
||||
assert "first" not in str(first)
|
||||
assert "second" not in str(second)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
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())
|
||||
|
||||
async def cancelled() -> list[Prompt]:
|
||||
raise asyncio.CancelledError()
|
||||
|
||||
async def supported() -> list[Prompt]:
|
||||
return [Prompt(name="recovered")]
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await cache.get(("server", None), cancelled)
|
||||
assert [item.name for item in await cache.get(("server", None), supported)] == ["recovered"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue