diff --git a/litellm/llms/anthropic/pass_through/messages/mcp_handler.py b/litellm/llms/anthropic/pass_through/messages/mcp_handler.py index ab722a60ca5..7961a38dedf 100644 --- a/litellm/llms/anthropic/pass_through/messages/mcp_handler.py +++ b/litellm/llms/anthropic/pass_through/messages/mcp_handler.py @@ -147,6 +147,7 @@ async def anthropic_messages_with_mcp( tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, + served_tools=deduplicated_mcp_tools, tool_calls=list(tool_use_blocks), user_api_key_auth=context.user_api_key_auth, mcp_auth_header=context.mcp_auth_header, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a776470e2ab..49cd8cf6fb9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -28,7 +28,7 @@ from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache from itertools import chain, groupby -from types import MappingProxyType +from types import EllipsisType, MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse @@ -259,6 +259,20 @@ _user_env_vars_cache: Final[dict[tuple[str, str], tuple[dict[str, str], float]]] _USER_ENV_VARS_CACHE_TTL: Final = 60 # seconds _USER_ENV_VARS_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth +_ListedToolsByCaller: TypeAlias = Mapping[str | None, Mapping[str, MCPTool]] +_LISTED_TOOLS_CALLERS_PER_SERVER: Final = 256 + + +@dataclass(frozen=True, slots=True) +class ListedToolsCaller: + """Request inputs that select which upstream catalog a caller was shown by tools/list.""" + + user_api_key_auth: UserAPIKeyAuth | None = None + mcp_auth_header: str | dict[str, str] | None = None + raw_headers: Mapping[str, str] | None = None + oauth2_headers: Mapping[str, str] | None = None + + # Auth types whose upstream OAuth endpoints (protected-resource + authorization-server metadata) the # gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes. # OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the @@ -1172,6 +1186,57 @@ def _authorization_is_litellm_admission_credential( return bool(user_api_key_auth and user_api_key_auth.api_key and not admission_header) +def _server_auth_header_for( + server: MCPServer, + mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None, + mcp_auth_header: str | dict[str, str] | None, +) -> str | dict[str, str] | None: + """Server-specific ``x-mcp--authorization`` header, else the deprecated global one.""" + server_specific: Final = ( + lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers, + alias=server.alias, + server_name=server.server_name, + access_groups=server.access_groups, + ) + if mcp_server_auth_headers + else None + ) + return mcp_auth_header if server_specific is None else server_specific + + +def listed_tools_caller_for( + server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | dict[str, str] | None, + mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None, + raw_headers: Mapping[str, str] | None, + oauth2_headers: Mapping[str, str] | None, +) -> ListedToolsCaller: + """The caller a tools/call must look its listed entry up under: the same inputs tools/list keyed by.""" + return ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=_server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header), + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ) + + +def _admission_identity( + auth: UserAPIKeyAuth, raw_headers: Mapping[str, str] | None +) -> tuple[str | None, str | None, str | None, str | None, str | None]: + """The admission identity the served catalog is shaped for: the hashed key, user, team and + organization, plus the admission credential (``x-litellm-api-key``, else ``Authorization``) of a + caller admitted with neither a key nor a user.""" + keyless: Final = auth.api_key is None and auth.user_id is None + credential: Final = ( + _raw_header_value(raw_headers, "x-litellm-api-key") or _raw_header_value(raw_headers, "authorization") + if keyless + else None + ) + return auth.api_key, auth.user_id, auth.team_id, auth.org_id, credential + + def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str: """Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection. @@ -1297,6 +1362,16 @@ async def _resolve_byok_mcp_auth_header( return mcp_auth_header +def _catalog_auth_header( + mcp_auth_header: str | dict[str, str] | None, + catalog_auth_header: str | dict[str, str] | None | EllipsisType, +) -> str | dict[str, str] | None: + """The header the client supplied, which keys the caller's catalog slot on both tools/list and + tools/call. A caller that already swapped a stored BYOK credential into ``mcp_auth_header`` passes + the client's value explicitly, since the stored credential must never be read to find the slot.""" + return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header + + def _client_forwarded_authorization_headers( mcp_server: MCPServer, oauth2_headers: dict[str, str] | None, @@ -1925,6 +2000,8 @@ class MCPServerManager: "gmail_send_email": "zapier_mcp_server", } """ + self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list + self._listed_tools_generations: dict[str, int] = {} # mutable-ok: bumped per server save self._upstream_initialize_instructions_by_server_id: dict[str, str] = {} # Per-server monotonic timestamp of last upstream prefetch attempt (success, # empty result, or failure). Used to throttle re-probes for servers that do @@ -3242,6 +3319,8 @@ class MCPServerManager: self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) + if new_server.spec_path: + self._drop_listed_tools(mcp_server.server_id) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Added MCP Server: %s", new_server.name) @@ -3279,6 +3358,8 @@ class MCPServerManager: self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) + if new_server.spec_path: + self._drop_listed_tools(mcp_server.server_id) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Updated MCP Server: %s", new_server.name) @@ -3754,25 +3835,14 @@ class MCPServerManager: verbose_logger.warning("MCP Server %s not found", server_id) return [] - # Get server-specific auth header if available - server_auth_header: str | dict[str, str] | None = None - if mcp_server_auth_headers: - server_auth_header = lookup_mcp_server_auth_in_headers( - mcp_server_auth_headers, - alias=server.alias, - server_name=server.server_name, - access_groups=server.access_groups, - ) - - # Fall back to deprecated mcp_auth_header if no server-specific header found - if server_auth_header is None: - server_auth_header = mcp_auth_header + server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header) try: tools: Final = await self._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, user_api_key_auth=user_api_key_auth, + record_listing=True, ) return tools except Exception as e: @@ -3852,7 +3922,7 @@ class MCPServerManager: def _build_stdio_env( self, server: MCPServer, - raw_headers: dict[str, str] | None = None, + raw_headers: Mapping[str, str] | None = None, ) -> dict[str, str] | None: """Resolve stdio env values, supporting header-driven placeholders.""" @@ -4147,13 +4217,20 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None = None, client_ip: str | None = None, proxy_logging_obj: ProxyLogging | None = None, - ) -> list[MCPTool]: + *, + catalog_auth_header: str | dict[str, str] | None | EllipsisType = ..., + record_listing: bool = False, + ) -> Sequence[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. Args: server (MCPServer): The server to query tools from mcp_auth_header: Optional auth header for MCP server + catalog_auth_header: The header the client supplied, keying the caller's catalog slot; + defaults to ``mcp_auth_header`` + record_listing: Record the served catalog into the caller's listed-tools slot; only a + listing actually served to the caller sets it Returns: List[MCPTool]: List of tools available on the server with prefixed names @@ -4169,6 +4246,13 @@ class MCPServerManager: verbose_logger.info("_get_tools_from_server for %s...", server.name) client = None + listed_caller: Final = ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=_catalog_auth_header(mcp_auth_header, catalog_auth_header), + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ) + listed_generation: Final = self._listed_tools_generations.get(server.server_id, 0) try: # Tool *listing* must not be blocked by missing per-user env vars — @@ -4266,8 +4350,12 @@ class MCPServerManager: # applied (e.g. "test_petstore-getinventory"). Do NOT pass them # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". + unprefixed_tools: Final = guarded_openapi + self._record_listed_tools( + server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing + ) if not add_prefix: - return list(guarded_openapi) + return unprefixed_tools return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] else: tools = await self._fetch_tools_with_timeout(client, server.name) @@ -4281,7 +4369,10 @@ class MCPServerManager: raw_headers=raw_headers, ) prefixed_or_original_tools: Final = self._create_prefixed_tools( - list(guarded_tools), server, add_prefix=add_prefix + guarded_tools, server, add_prefix=add_prefix + ) + self._record_listed_tools( + server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing ) return prefixed_or_original_tools @@ -4332,8 +4423,99 @@ class MCPServerManager: ) self._invalidate_discovery_lists(server_id) + self._drop_listed_tools(server_id) invalidate_oauth_metadata_cache(server_id) + def _drop_listed_tools(self, server_id: str) -> None: + self._listed_tools_by_server_id.pop(server_id, None) + self._listed_tools_generations[server_id] = self._listed_tools_generations.get(server_id, 0) + 1 + + def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None: + """Key the listed-tool cache by every request input that can change the served catalog. + + The catalog is guardrail-shaped for the caller's admission identity (default-on guardrails, + key or team selections and opt-outs), so every admitted caller gets its own slot, keyed by + ``_admission_identity``: the hashed key, user, team and organization, plus the admission + credential of a caller admitted with neither a key nor a user (a team-only JWT). Forwarded + headers, header-driven stdio env, the caller bearer on every server whose egress forwards it + (``_consumes_caller_authorization``) or exchanges it as the OBO subject, and the + server-specific auth header also reach upstream and split the slot further. Only unkeyed + listings with none of those share the ``None`` slot. + """ + if caller is None: + return None + auth: Final = caller.user_api_key_auth + forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None + header_env: Final = self._build_stdio_env(server, caller.raw_headers) + stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env + caller_bearer: Final = ( + self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) + if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange + else None + ) + identity: Final = None if auth is None else _admission_identity(auth, caller.raw_headers) + if not (identity or caller.mcp_auth_header or forwarded or stdio_env or caller_bearer): + return None + material: Final = json.dumps( + (identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer), + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha256(material.encode()).hexdigest() + + @staticmethod + def _forwarded_header_values( + server: MCPServer, raw_headers: Mapping[str, str] | None + ) -> tuple[tuple[str, str], ...]: + if not raw_headers or not server.extra_headers: + return () + forwarded_names: Final = frozenset(name.lower() for name in server.extra_headers) + return tuple( + sorted((name.lower(), value) for name, value in raw_headers.items() if name.lower() in forwarded_names) + ) + + def listed_tools_generation(self, server_id: str) -> int: + return self._listed_tools_generations.get(server_id, 0) + + def record_listed_tools( + self, + server: MCPServer, + tools: Sequence[MCPTool], + caller: ListedToolsCaller | None, + generation: int, + *, + record_listing: bool = True, + ) -> None: + self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing) + + def _record_listed_tools( + self, + server: MCPServer, + tools: Sequence[MCPTool], + caller: ListedToolsCaller | None, + generation: int | None = None, + *, + record_listing: bool = True, + ) -> None: + """Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation + read before the listing's upstream fetch; the record is skipped when it no longer matches.""" + if not record_listing: + return + if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0): + return + identity: Final = self._listed_tools_identity(server, caller) + listing: Final = MappingProxyType({tool.name: tool for tool in tools}) + existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})) + shared: Final = existing.get(None) + callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity)) + evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0) + entries: Final = ( + *(() if shared is None else ((None, shared),)), + *callers[evicted:], + (identity, listing), + ) + self._listed_tools_by_server_id[server.server_id] = MappingProxyType(dict(entries)) + def _discovery_key( self, server: MCPServer, @@ -4343,9 +4525,11 @@ class MCPServerManager: stdio_env: dict[str, str] | None, subject_token: str | None, credential_fingerprint: str | None = None, + per_caller: bool = False, ) -> _DiscoveryKey: per_user: Final = ( - server.requires_per_user_auth + 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) @@ -5253,7 +5437,12 @@ class MCPServerManager: {seen: kept for seen, kept in self._catalog_alert_signatures.items() if seen != key} ) - def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]: + def _create_prefixed_tools( + self, + tools: Sequence[MCPTool], + server: MCPServer, + add_prefix: bool = True, + ) -> list[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -5281,6 +5470,13 @@ class MCPServerManager: verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name) return prefixed_tools + def get_listed_tool(self, server: MCPServer, name: str, caller: ListedToolsCaller | None = None) -> MCPTool | None: + identity: Final = self._listed_tools_identity(server, caller) + listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity) + if not listed: + return None + return listed.get(name) + def _create_prefixed_prompts( self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True ) -> list[Prompt]: @@ -5514,6 +5710,7 @@ class MCPServerManager: raw_headers: dict[str, str] | None = None, litellm_logging_obj: "LiteLLMLoggingObj | None" = None, guardrail_context: Mapping[str, object] | None = None, + tool: MCPTool | None = None, ) -> dict[str, Any]: """ Run pre-call checks and guardrail hooks for an MCP tool call. @@ -5527,6 +5724,9 @@ class MCPServerManager: ``pre_mcp_call`` evaluation (or a block) on the spend-log row the Guardrails Monitor counts. It stays optional so callers that do no logging are unchanged. + ``tool`` is the upstream tool definition when one was listed, so guardrails + can see its description and input schema, not just the name and arguments. + Returns a dict that may contain: - "arguments": hook-modified tool arguments (only if changed) - "extra_headers": headers injected by pre_mcp_call guardrail hooks @@ -5590,6 +5790,8 @@ class MCPServerManager: "user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None), "incoming_bearer_token": incoming_bearer_token, "headers": logging_safe_mcp_headers(raw_headers), + "tool_description": tool.description if tool is not None else None, + "tool_input_schema": tool.input_schema if tool is not None else None, } # Create MCP request object for processing @@ -5805,21 +6007,7 @@ class MCPServerManager: GuardrailRaisedException: If guardrails block the call HTTPException: If an HTTP error occurs """ - # Get server-specific auth header if available (case-insensitive) - # FIX: Added case-insensitive matching to handle auth header keys that may not match - # the exact case of server alias/name (e.g., '1litellmagcgateway' vs '1LiteLLMAGCGateway') - server_auth_header: dict[str, str] | str | None = None - if mcp_server_auth_headers: - server_auth_header = lookup_mcp_server_auth_in_headers( - mcp_server_auth_headers, - alias=mcp_server.alias, - server_name=mcp_server.server_name, - access_groups=mcp_server.access_groups, - ) - - # Fall back to deprecated mcp_auth_header if no server-specific header found - if server_auth_header is None: - server_auth_header = mcp_auth_header + server_auth_header: Final = _server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header) # Extract subject token for OAuth2 Token Exchange (OBO) and ID-JAG flows subject_token: str | None = None @@ -6242,6 +6430,9 @@ class MCPServerManager: guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + *, + catalog_auth_header: str | None | EllipsisType = ..., + listed_tool: MCPTool | None | EllipsisType = ..., ) -> CallToolResult | InputRequiredResult: """ Call a tool with the given name and arguments @@ -6253,6 +6444,8 @@ class MCPServerManager: user_api_key_auth: User authentication mcp_auth_header: MCP auth header (deprecated) mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value} + catalog_auth_header: The header the client supplied, keying the caller's catalog slot; + defaults to ``mcp_auth_header`` as received, before BYOK resolution proxy_logging_obj: Optional ProxyLogging object for hook integration litellm_logging_obj: Optional request logger the guardrail hooks record their evaluations onto, so MCP guardrail activity reaches the @@ -6264,6 +6457,7 @@ class MCPServerManager: """ start_time: Final = datetime.datetime.now() mcp_server: Final = self._resolve_mcp_server_for_tool_call(server_name, name) + client_auth_header: Final = _catalog_auth_header(mcp_auth_header, catalog_auth_header) # Resolved before any hook runs so a missing BYOK credential (401) never # leaves during-hook side effects (audit logging, rate-limit bookkeeping) @@ -6273,6 +6467,9 @@ class MCPServerManager: user_api_key_auth, mcp_auth_header, ) + listed_caller: Final = listed_tools_caller_for( + mcp_server, user_api_key_auth, client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers + ) ######################################################### # Pre MCP Tool Call Hook @@ -6289,6 +6486,7 @@ class MCPServerManager: raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + tool=self.get_listed_tool(mcp_server, name, listed_caller) if listed_tool is ... else listed_tool, ) if "arguments" in hook_result: arguments = hook_result["arguments"] diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index f98a8c0ea5a..d83a4a72e2f 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -81,12 +81,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( outcome_wire_value, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + ListedToolsCaller, MCPServerManager, _caller_authorization_fans_out, _client_forwarded_authorization_headers, _resolve_openapi_tool_auth, _should_strip_caller_authorization, global_mcp_server_manager, + listed_tools_caller_for, ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( _redact_mcp_resource_url, @@ -954,6 +956,8 @@ async def _get_tools_from_mcp_servers( request_tags: list[str] | None = None, client_ip: str | None = None, mcp_proxy_mode: bool = False, + *, + record_listing: bool = False, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -964,6 +968,8 @@ async def _get_tools_from_mcp_servers( mcp_servers: Optional list of server names/aliases to filter by mcp_server_auth_headers: Optional dict of server-specific auth headers oauth2_headers: Optional dict of oauth2 headers + record_listing: Record each served catalog into the caller's listed-tools slot; only a + listing actually served to the caller sets it Returns: AggregateToolListing: Combined tools from filtered servers plus each server's @@ -1111,12 +1117,14 @@ async def _get_tools_from_mcp_servers( prefetched_creds=_prefetched_oauth_creds, ) + catalog_auth_header: Final = server_auth_header if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None: server_auth_header = await _get_byok_credential(server, user_api_key_auth) try: from litellm.proxy.proxy_server import proxy_logging_obj + listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) tools: Final = await global_mcp_server_manager._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, @@ -1127,6 +1135,8 @@ async def _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, + catalog_auth_header=catalog_auth_header, + record_listing=False, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1135,6 +1145,21 @@ async def _get_tools_from_mcp_servers( server_id=server.server_id, user_api_key_auth=user_api_key_auth, ) + global_mcp_server_manager.record_listed_tools( + server, + [ + tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)}) + for tool in filtered_tools + ], + ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=catalog_auth_header, + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ), + listed_generation, + record_listing=record_listing, + ) if mcp_proxy_mode: from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity @@ -1458,6 +1483,8 @@ async def _list_mcp_tools( list_tools_log_source: str | None = None, client_ip: str | None = None, mcp_proxy_mode: bool = False, + *, + record_listing: bool = False, ) -> AggregateToolListing: """ List all available MCP tools. @@ -1468,6 +1495,8 @@ async def _list_mcp_tools( mcp_servers: Optional list of server names/aliases to filter by mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value} client_ip: Client IP for IP-based server access control + record_listing: Record each served catalog into the caller's listed-tools slot; only a + listing actually served to the caller sets it Returns: AggregateToolListing: Combined tools from all accessible servers plus each server's @@ -1486,6 +1515,7 @@ async def _list_mcp_tools( list_tools_log_source=list_tools_log_source, client_ip=client_ip, mcp_proxy_mode=mcp_proxy_mode, + record_listing=record_listing, ) verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools)) return listing @@ -1798,6 +1828,7 @@ async def _list_tools_before_first_call( oauth2_headers=oauth2_headers, raw_headers=raw_headers, client_ip=client_ip, + record_listing=False, ) except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e) @@ -2001,6 +2032,7 @@ async def _execute_mcp_tool( if mcp_server is None: mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + client_auth_header: Final = mcp_auth_header if mcp_server: standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get("mcp_server_cost_info") if litellm_logging_obj: @@ -2071,6 +2103,18 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + tool=global_mcp_server_manager.get_listed_tool( + mcp_server, + original_tool_name, + listed_tools_caller_for( + mcp_server, + user_api_key_auth, + client_auth_header, + mcp_server_auth_headers, + raw_headers, + oauth2_headers, + ), + ), ) # `pre_call_tool_check` may return guardrail-modified # arguments; honor them on the local path too. @@ -2120,6 +2164,7 @@ async def _execute_mcp_tool( arguments=arguments, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, + catalog_auth_header=client_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, @@ -2139,7 +2184,8 @@ async def _execute_mcp_tool( # not in the registry either, `_handle_local_mcp_tool` below reports # 404 and nothing runs, so demanding a server here would turn every # unknown tool name into a misleading 503. - if global_mcp_tool_registry.get_tool(original_tool_name) is not None: + registered_local_tool: Final = global_mcp_tool_registry.get_tool(original_tool_name) + if registered_local_tool is not None: # `mcp_server` is None here because the tool name is not in the # tool -> server mapping, but the name still carries a prefix # that the server-level check above compared against the @@ -2181,6 +2227,18 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + tool=global_mcp_server_manager.get_listed_tool( + prefix_server, + original_tool_name, + listed_tools_caller_for( + prefix_server, + user_api_key_auth, + client_auth_header, + mcp_server_auth_headers, + raw_headers, + oauth2_headers, + ), + ), ) if "arguments" in hook_result: arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args @@ -2583,8 +2641,12 @@ async def _handle_managed_mcp_tool( guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + *, + catalog_auth_header: str | None, ) -> CallToolResult | InputRequiredResult: - """Handle tool execution for managed server tools""" + """Handle tool execution for managed server tools. ``catalog_auth_header`` is the header the client + supplied, which keys the caller's catalog slot; ``mcp_auth_header`` may already be the resolved + BYOK credential.""" # Import here to avoid circular import from litellm.proxy.proxy_server import proxy_logging_obj @@ -2594,6 +2656,7 @@ async def _handle_managed_mcp_tool( arguments=arguments, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, + catalog_auth_header=catalog_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, @@ -2712,6 +2775,7 @@ async def _execute_handle_list_tools( log_list_tools_to_spendlogs=log_list_tools_to_spendlogs, list_tools_log_source="mcp_protocol", client_ip=_client_ip, + record_listing=True, ) verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools)) if not listing.outcomes: diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7f936ec9269..845dcaa1f16 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -223,6 +223,7 @@ if MCP_AVAILABLE: ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, + ListedToolsCaller, global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( @@ -701,6 +702,8 @@ if MCP_AVAILABLE: extra_headers: dict[str, str] | None, client_ip: str | None, proxy_logging_obj: "ProxyLogging | None", + *, + record_listing: bool, ) -> list[MCPTool]: return await global_mcp_server_manager._get_tools_from_server( server=server, @@ -711,11 +714,12 @@ if MCP_AVAILABLE: client_ip=client_ip, user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, + record_listing=record_listing, ) async def _get_tools_for_single_server( - server, - server_auth_header, + server: MCPServer, + server_auth_header: dict[str, str] | str | None, raw_headers: dict[str, str] | None = None, user_api_key_auth: UserAPIKeyAuth | None = None, extra_headers: dict[str, str] | None = None, @@ -731,31 +735,44 @@ if MCP_AVAILABLE: """ from litellm.proxy.proxy_server import proxy_logging_obj - tools = await _list_server_tools( - server, server_auth_header, raw_headers, user_api_key_auth, extra_headers, client_ip, proxy_logging_obj + listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) + tools: Final = await _list_server_tools( + server, + server_auth_header, + raw_headers, + user_api_key_auth, + extra_headers, + client_ip, + proxy_logging_obj, + record_listing=False, ) - if not apply_tool_filters: - return _create_tool_response_objects(tools, server) - - # Always apply allowed_tools/disallowed_tools so the blacklist is - # enforced even when no allowlist is set (matches the SSE/HTTP path). - tools = filter_tools_by_allowed_tools(tools, server) - - # Filter by the key's effective tool permissions through the same - # function the MCP protocol path uses (direct grants, toolset grants, - # and team/agent/org ceilings), so REST listing cannot drift from it. - # Entries here are tool names on one server, written bare by every - # writer, and dispatch compares them bare; matching a wider set of - # spellings would advertise a tool that tools/call then refuses - if user_api_key_auth: - tools = await filter_tools_by_key_team_permissions( - tools=tools, + server_filtered: Final = filter_tools_by_allowed_tools(tools, server) if apply_tool_filters else tools + served_tools: Final = ( + await filter_tools_by_key_team_permissions( + tools=server_filtered, server_id=server.server_id, user_api_key_auth=user_api_key_auth, ) + if apply_tool_filters and user_api_key_auth + else server_filtered + ) + if apply_tool_filters: + # Only a listing shaped for the caller's runtime view may set their + # listed-tools slot; the admin-only unfiltered configuration view + # must not warm it. + global_mcp_server_manager.record_listed_tools( + server, + served_tools, + ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + raw_headers=raw_headers, + ), + listed_generation, + ) - return _create_tool_response_objects(tools, server) + return _create_tool_response_objects(served_tools, server) async def fetch_pinnable_tool_catalog( server: MCPServer, request: Request, user_api_key_dict: UserAPIKeyAuth @@ -774,6 +791,7 @@ if MCP_AVAILABLE: await _get_user_oauth_extra_headers(server, user_api_key_dict), IPAddressUtils.get_mcp_client_ip(request), None, + record_listing=False, ) scan: Final = await scan_tool_descriptions( apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index 6621d94ed96..c683e1f2c5d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn import httpx from fastapi import HTTPException -from pydantic import TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger @@ -61,6 +61,7 @@ _GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset( _INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027" _AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...]) _MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool") +_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object]) _OBO_CACHE_MAX_ENTRIES: Final = 1000 _DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0 _TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0 @@ -82,6 +83,13 @@ def _parse_aadsts_codes(raw: object) -> tuple[int, ...]: return () +def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None: + try: + return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw) + except ValidationError: + return None + + def entra_assertion(value: object) -> str | None: """``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion. A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``.""" @@ -100,6 +108,14 @@ class _EvaluateResponse(TypedDict, total=False): correlationId: ReadOnly[str] +class _ToolReference(BaseModel): + model_config = ConfigDict(frozen=True) + + name: str + description: str | None = None + input_schema: Mapping[str, object] | None = Field(default=None, serialization_alias="inputSchema") + + class _UnavailableDetail(TypedDict): error: ReadOnly[str] message: ReadOnly[str] @@ -392,8 +408,14 @@ class Agent365Guardrail(CustomGuardrail): arguments: Final = data.get("mcp_arguments") server_name: Final = str(data.get("mcp_server_name") or "litellm") agent_id: Final = user_api_key_dict.key_alias + description: Final = data.get("mcp_tool_description") + tool_reference: Final = _ToolReference( + name=tool_name, + description=description if isinstance(description, str) and description else None, + input_schema=_parse_tool_input_schema(data.get("mcp_input_schema")), + ) payload: Final[dict[str, object]] = { # mutable-ok: JSON body with optional fields added below - "tool": {"name": tool_name}, + "tool": tool_reference.model_dump(by_alias=True, exclude_none=True), "serverName": server_name, "conversationId": self._resolve_conversation_id(data), } diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index be4e46f247e..e1504172d9e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1018,6 +1018,7 @@ if MCP_AVAILABLE: mcp_auth_header=None, mcp_servers=None, mcp_server_auth_headers=None, + record_listing=True, ) tools: Final = listing.tools dumped_tools: Final = [tool.model_dump(by_alias=True) for tool in tools] diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8af90e19ecb..1c738aefb97 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1151,6 +1151,8 @@ def _overrides_moderation_hook(callback: CustomLogger) -> bool: _LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...]) +_MCP_TOOL_DESCRIPTION: Final[TypeAdapter[str | None]] = TypeAdapter(str | None) +_MCP_TOOL_INPUT_SCHEMA: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None) @dataclass(frozen=True, slots=True) @@ -1475,9 +1477,14 @@ class ProxyLogging: TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({})) ) - mcp_tool_description: Final = kwargs.get("mcp_tool_description") - mcp_input_schema: Final = kwargs.get("mcp_input_schema") - description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else "" + mcp_tool_description: Final = request_obj.tool_description or kwargs.get("mcp_tool_description") + mcp_input_schema: Final = ( + request_obj.tool_input_schema + if request_obj.tool_input_schema is not None + else kwargs.get("mcp_input_schema") + ) + listing_description: Final = kwargs.get("mcp_tool_description") + description_line: Final = f"\nDescription: {listing_description}" if listing_description else "" tool_call_content: Final = ( f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}" ) @@ -1735,6 +1742,8 @@ class ProxyLogging: tool_name=kwargs.get("name", ""), arguments=kwargs.get("arguments", {}), server_name=kwargs.get("server_name"), + tool_description=_MCP_TOOL_DESCRIPTION.validate_python(kwargs.get("tool_description")), + tool_input_schema=_MCP_TOOL_INPUT_SCHEMA.validate_python(kwargs.get("tool_input_schema")), user_api_key_auth=user_api_key_auth_dict, hidden_params=HiddenParams(), ) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 14c12fc571d..5e2793ed40d 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -286,6 +286,7 @@ async def aresponses_api_with_mcp( call_params=call_params, previous_response_id=previous_response_id, tool_server_map=tool_server_map, + served_tools=original_mcp_tools, **kwargs, ) await mcp_streaming_response._create_initial_response_iterator() @@ -339,6 +340,7 @@ async def aresponses_api_with_mcp( tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, + served_tools=original_mcp_tools, tool_calls=tool_calls, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -395,6 +397,7 @@ async def aresponses_api_with_mcp( final_response = MCPEnhancedStreamingIterator( tool_server_map=tool_server_map, + served_tools=original_mcp_tools, base_iterator=final_response, mcp_events=tool_execution_events, user_api_key_auth=user_api_key_auth, diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index b17e0befba6..8aa4181cc8a 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -435,6 +435,7 @@ async def acompletion_with_mcp( # Execute tool calls self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=self.tool_server_map, + served_tools=deduplicated_mcp_tools, tool_calls=self.tool_calls, user_api_key_auth=self.user_api_key_auth, mcp_auth_header=self.mcp_auth_header, @@ -609,6 +610,7 @@ async def acompletion_with_mcp( tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, tool_calls=tool_calls, + served_tools=deduplicated_mcp_tools, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index fb7882adcb4..1d6c0695621 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -696,6 +696,7 @@ class LiteLLM_Proxy_MCP_Handler: litellm_trace_id: str | None = None, request_tags: list[str] | None = None, guardrail_context: Mapping[str, object] | None = None, + served_tools: Sequence[MCPTool] | None = None, ) -> list[MCPToolResult]: """Execute tool calls and return results.""" from fastapi import HTTPException @@ -860,6 +861,11 @@ class LiteLLM_Proxy_MCP_Handler: proxy_logging_obj=proxy_logging_obj, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + listed_tool=( + next((tool for tool in served_tools if tool.name == tool_name), None) + if served_tools is not None + else ... + ), ) if proxy_logging_obj: @@ -1152,6 +1158,7 @@ class LiteLLM_Proxy_MCP_Handler: call_params: Mapping[str, object], previous_response_id: str | None, tool_server_map: dict[str, str], + served_tools: Sequence[MCPTool] | None = None, **kwargs, ) -> Any: """ @@ -1181,6 +1188,7 @@ class LiteLLM_Proxy_MCP_Handler: base_iterator=None, # Will be created internally mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events tool_server_map=tool_server_map, + served_tools=served_tools, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, user_api_key_auth=kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"), diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index c0bb92cfb2a..b1f12233f33 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -281,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None, user_api_key_auth: "UserAPIKeyAuth | None" = None, original_request_params: dict[str, Any] | None = None, + served_tools: Sequence[MCPTool] | None = None, ): # MCP setup self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or [] @@ -300,6 +301,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.mcp_discovery_generated = True # Events are already generated self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility self.tool_server_map = tool_server_map + self.served_tools = tuple(served_tools) if served_tools is not None else None # Iterator references self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = ( @@ -796,6 +798,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # Execute the tools tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=self.tool_server_map, + served_tools=self.served_tools, tool_calls=tool_calls, user_api_key_auth=self.user_api_key_auth, mcp_auth_header=self.mcp_auth_header, diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index da7401e2a2e..ce86e79a6f0 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -429,6 +429,8 @@ class MCPPreCallRequestObject(BaseModel): tool_name: str arguments: dict[str, Any] server_name: str | None = None + tool_description: str | None = None + tool_input_schema: Mapping[str, object] | None = None user_api_key_auth: dict[str, Any] | None = None hidden_params: HiddenParams = HiddenParams() @@ -452,6 +454,8 @@ class MCPDuringCallRequestObject(BaseModel): tool_name: str arguments: dict[str, Any] server_name: str | None = None + tool_description: str | None = None + tool_input_schema: Mapping[str, object] | None = None start_time: float | None = None hidden_params: HiddenParams = HiddenParams() diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index a3693433de4..d39d7e4029e 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -216,6 +216,16 @@ JsonRpc = Mapping[str, object] class ScriptedTool: name: str respond: Callable[[JsonRpc], Reply | JsonRpc] + description: str | Callable[[Mapping[str, str]], str] | None = None + input_schema: JsonRpc = field(default_factory=lambda: {"type": "object"}) + + def listing(self, headers: Mapping[str, str]) -> JsonRpc: + described: Final = self.description(headers) if callable(self.description) else self.description + return { + "name": self.name, + "inputSchema": self.input_schema, + **({} if described is None else {"description": described}), + } def jsonrpc_reply(identity: object, result: JsonRpc) -> Reply: @@ -253,9 +263,7 @@ def scripted_peer(*tools: ScriptedTool) -> Iterator[McpPeer]: }, ) if method == "tools/list": - return jsonrpc_reply( - identity, {"tools": [{"name": name, "inputSchema": {"type": "object"}} for name in by_name]} - ) + return jsonrpc_reply(identity, {"tools": [tool.listing(request.headers) for tool in by_name.values()]}) if method != "tools/call": return jsonrpc_error(identity, -32601, f"unsupported method {method}") tool: Final = by_name.get(body["params"]["name"]) diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index c807852ef6b..10bb7787cc9 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -133,6 +133,7 @@ def wire_server( class OwnedHTTPServer(ThreadingHTTPServer): daemon_threads = False + request_queue_size = 128 def server_bind(self) -> None: super().server_bind() diff --git a/tests/integration/mcp/test_mcp_access_matrix.py b/tests/integration/mcp/test_mcp_access_matrix.py index e458911cc68..bf8adbc7cd7 100644 --- a/tests/integration/mcp/test_mcp_access_matrix.py +++ b/tests/integration/mcp/test_mcp_access_matrix.py @@ -1,20 +1,29 @@ +import re +import textwrap import uuid +from collections.abc import Iterator +from pathlib import Path from typing import Final import pytest -from integration._support.client import Gateway +from integration._support.client import Gateway, gateway_from_environment from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, EntryPoint, McpCaller, McpPeer, + Outcome, PeerKind, + ScriptedTool, peer_of, register_mcp, + scripted_peer, + text_result, tool_calls, ) from integration._support.mcp_grants import SUBJECTS, Subject, grant +from integration._support.process import owned_proxy CALLABLE: Final = {"add": {"a": 1, "b": 2}, "multiply": {"a": 2, "b": 3}} RESULTS: Final = {"add": "3", "multiply": "6"} @@ -135,3 +144,87 @@ def test_same_tool_name_on_two_servers_routes_by_prefix(gateway: Gateway) -> Non assert outcome.ok and outcome.text == "10", outcome.raw assert tool_calls(first.drain()) == () assert [call["body"]["params"]["name"] for call in tool_calls(second.drain())] == ["add"] + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +@pytest.mark.parametrize("subject", SUBJECTS) +def test_each_subjects_call_is_evaluated_only_against_the_catalog_its_own_listing_served( + echo_rig: Gateway, subject: Subject +) -> None: + described: Final = "Adds under grant " + uuid.uuid4().hex[:8] + tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + group: Final = "grp" + uuid.uuid4().hex[:8] + alias: Final = "cat" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, mcp_access_groups=[group]) + caller: Final = grant( + scenario, subject, (identity,), (identity,), access_group=group, allowed_tools={identity: ("add",)} + ) + reach: Final = McpCaller(echo_rig, caller.key, "mcp", alias, caller.headers) + assert reach.initialize().ok + cold: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE})) + listed: Final = reach.list_tools() + assert listed.ok and f"{alias}-add" in listed.tools, listed.raw + warm: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE})) + assert (cold, warm) == (_UNLISTED, described), (cold, warm) + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" + + +def test_end_users_of_one_key_share_its_catalog_slot_because_the_identity_excludes_the_end_user( + echo_rig: Gateway, +) -> None: + described: Final = "Adds for end users " + uuid.uuid4().hex[:8] + tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "eu" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + granted: Final = grant(scenario, "end_user", (identity,), (identity,)) + first: Final = McpCaller(echo_rig, granted.key, "mcp", alias, granted.headers) + second: Final = McpCaller( + echo_rig, granted.key, "mcp", alias, {"x-litellm-end-user-id": "integration-" + uuid.uuid4().hex[:10]} + ) + assert second.initialize().ok + assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == _UNLISTED + listed: Final = first.list_tools() + assert listed.ok and f"{alias}-add" in listed.tools, listed.raw + assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == described, ( + "the end-user header is intentionally not part of the catalog identity: one key, one slot" + ) + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" diff --git a/tests/integration/mcp/test_mcp_accounting_guardrails.py b/tests/integration/mcp/test_mcp_accounting_guardrails.py index 323afad40db..a276718c280 100644 --- a/tests/integration/mcp/test_mcp_accounting_guardrails.py +++ b/tests/integration/mcp/test_mcp_accounting_guardrails.py @@ -1,11 +1,23 @@ +import json import uuid -from collections.abc import Iterator +from collections.abc import Generator, Iterator, Mapping from contextlib import contextmanager +from dataclasses import dataclass from hashlib import sha256 +from pathlib import Path from typing import Final +import httpx import pytest -from integration._support.client import Gateway, JsonValue, Scenario, eventually +import yaml +from integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, +) from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, @@ -13,10 +25,16 @@ from integration._support.mcp import ( McpCaller, McpPeer, Outcome, + ScriptedTool, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, ) +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue DEFAULT_COST: Final = 0.25 ADD_COST: Final = 0.5 @@ -191,3 +209,472 @@ def test_guardrail_removal_stops_blocking_without_restart(gateway: Gateway) -> N lambda calls: len(calls) >= 1, seconds=40, ) + + +MASK_ME: Final = "mask-integration-secret" +MASKED: Final = "[MASKED]" +COUNT_MISMATCH: Final = "count-mismatch-marker" +LOOKUP_DESCRIPTION: Final = "Look up one record" +LOOKUP_SCHEMA: Final = { + "type": "object", + "properties": {"record": {"type": "string", "description": "record identifier"}}, +} +SPEND_ROW: Final = 'SELECT status, metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s' +_RECORDER_CODE: Final = """\ +import os + +import httpx +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger + +SINK = "{sink}/native" + + +def _record(stage, data, call_type): + logging_obj = data.get("litellm_logging_obj") + return {{ + "stage": stage, + "pid": os.getpid(), + "call_type": call_type, + "litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id, + "messages": data.get("messages"), + "mcp_tool_name": data.get("mcp_tool_name"), + "mcp_arguments": data.get("mcp_arguments"), + "mcp_tool_description": data.get("mcp_tool_description"), + "mcp_input_schema": data.get("mcp_input_schema"), + }} + + +async def _post(record): + async with httpx.AsyncClient(timeout=5) as client: + await client.post(SINK, json=record) + + +class HookRecorder(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + await _post(_record("pre", data, call_type)) + + async def async_moderation_hook(self, data, user_api_key_dict, call_type): + await _post(_record("during", data, call_type)) + + async def async_post_mcp_tool_call_hook(self, kwargs, response_obj, start_time, end_time): + await _post( + {{ + "stage": "post", + "pid": os.getpid(), + "litellm_call_id": kwargs.get("litellm_call_id"), + "tool": kwargs.get("mcp_tool_call_metadata"), + "content": [item.model_dump() for item in response_obj.mcp_tool_call_response], + }} + ) + + +class SinkGuardrail(CustomGuardrail): + def __init__(self, api_base, **kwargs): + super().__init__(**kwargs) + self.api_base = api_base + + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + payload = {{ + "pid": os.getpid(), + "input_type": input_type, + "litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id, + "texts": inputs.get("texts"), + "tools": inputs.get("tools"), + "structured_messages": inputs.get("structured_messages"), + "mcp_tool_name": request_data.get("mcp_tool_name"), + }} + async with httpx.AsyncClient(timeout=5) as client: + verdict = (await client.post(self.api_base, json=payload)).json() + return {{**inputs, "texts": verdict["texts"]}} + + +recorder = HookRecorder() +""" + + +@dataclass(frozen=True, slots=True) +class Sunk: + target: str + body: dict[str, JsonValue] + + +@dataclass(frozen=True, slots=True) +class HooksRig: + gateway: Gateway + sibling: Gateway + sink: Wire + guardrail: str + + def sunk(self) -> tuple[Sunk, ...]: + return tuple(Sunk(request.target, JSON_OBJECT.validate_json(request.body)) for request in self.sink.drain()) + + +def _guardrail_sink(request: Request) -> Reply: + if not request.target.startswith("/guardrail"): + return Reply() + texts: Final = JSON_OBJECT.validate_json(request.body).get("texts") + assert isinstance(texts, list), texts + masked: Final = [str(text).replace(MASK_ME, MASKED) for text in texts] + extra: Final = ["extra"] if any(COUNT_MISMATCH in text for text in masked) else [] + return Reply(body=json.dumps({"texts": [*masked, *extra]}).encode()) + + +@pytest.fixture(scope="module") +def hooks_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[HooksRig]: + directory: Final = tmp_path_factory.mktemp("guardrail-payloads") + guardrail: Final = "sink" + uuid.uuid4().hex[:8] + with wire_server(_guardrail_sink) as sink: + (directory / "hook_recorder.py").write_text(_RECORDER_CODE.format(sink=sink.url)) + config: Final = JSON_OBJECT.validate_python( + yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + ) + config["guardrails"] = [ + { + "guardrail_name": guardrail, + "litellm_params": { + "guardrail": "hook_recorder.SinkGuardrail", + "mode": ["pre_mcp_call", "post_mcp_call"], + "default_on": True, + "api_base": f"{sink.url}/guardrail", + }, + } + ] + config["litellm_settings"] = { + **object_value(config["litellm_settings"]), + "callbacks": ["hook_recorder.recorder"], + } + path: Final = directory / "config.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, directory, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate, + owned_proxy(gateway, directory, {}, config=path) as sibling, + ): + yield HooksRig(candidate, sibling, sink, guardrail) + + +def _worker(gateway: Gateway) -> int: + response: Final = gateway.client.get("/debug/memory/summary", headers={"x-litellm-api-key": gateway.key}) + assert response.status_code == 200, response.text + worker: Final = JSON_OBJECT.validate_json(response.content)["worker_pid"] + assert isinstance(worker, int), response.text + return worker + + +@contextmanager +def _pinned(gateway: Gateway) -> Generator[tuple[Gateway, int], None, None]: + limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=120) + with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False, limits=limits) as client: + pinned: Final = Gateway(client, gateway.key, gateway.upstream_url) + yield pinned, _worker(pinned) + + +def _lookup_tool(result: str = "found") -> ScriptedTool: + return ScriptedTool( + "lookup", lambda _: text_result(result), description=LOOKUP_DESCRIPTION, input_schema=LOOKUP_SCHEMA + ) + + +def _generic(sunk: tuple[Sunk, ...], input_type: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(item.body for item in sunk if item.target == "/guardrail" and item.body["input_type"] == input_type) + + +def _native(sunk: tuple[Sunk, ...], stage: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(item.body for item in sunk if item.target == "/native" and item.body["stage"] == stage) + + +def _only(records: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]: + assert len(records) == 1, records + return records[0] + + +def _scan(sunk: tuple[Sunk, ...], call_id: JsonValue) -> dict[str, JsonValue]: + return _only(tuple(record for record in _generic(sunk, "request") if record["litellm_call_id"] == call_id)) + + +def _texts(content: JsonValue) -> list[JsonValue]: + assert isinstance(content, list), content + return [object_value(item)["text"] for item in content] + + +def _has_lookup(listing: Outcome) -> bool: + return any(tool.endswith("lookup") for tool in listing.tools) + + +def _spend_row(call_id: JsonValue) -> dict[str, JsonValue]: + assert isinstance(call_id, str), call_id + rows: Final = eventually(lambda: read_rows(SPEND_ROW, (call_id,)), lambda found: len(found) == 1, seconds=70) + return rows[0] + + +def _synthetic_message(name: str, arguments: Mapping[str, str]) -> list[dict[str, str]]: + return [{"role": "user", "content": f"Tool: {name}\nArguments: {dict(arguments)}"}] + + +def test_generic_sink_and_native_hooks_receive_listed_metadata_on_typed_keys_with_the_message_bytes_unchanged( + hooks_rig: HooksRig, +) -> None: + with ( + scripted_peer(_lookup_tool()) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "payload" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + hooks_rig.sunk() + listed: Final = eventually(caller.list_tools, _has_lookup, seconds=30) + name: Final = next(tool for tool in listed.tools if tool.endswith("lookup")) + scans: Final = _generic(hooks_rig.sunk(), "request") + assert scans and all(scan["texts"] == [LOOKUP_DESCRIPTION, "record identifier"] for scan in scans), scans + arguments: Final = {"record": "r-1"} + outcome: Final = caller.call(name, arguments) + assert outcome.text == "found", outcome.raw + assert _worker(pinned) == worker + sunk: Final = hooks_rig.sunk() + pre: Final = _only(_native(sunk, "pre")) + call_id: Final = pre["litellm_call_id"] + generic: Final = _scan(sunk, call_id) + assert generic["texts"] == [LOOKUP_DESCRIPTION, "record identifier", "r-1"], generic + assert generic["tools"] == [ + { + "type": "function", + "function": { + "name": "lookup", + "description": LOOKUP_DESCRIPTION, + "parameters": {**LOOKUP_SCHEMA, "additionalProperties": False}, + "strict": False, + }, + } + ], generic + response: Final = _only(_generic(sunk, "response")) + assert (response["texts"], response["pid"]) == (["found"], worker), response + during: Final = _only(_native(sunk, "during")) + post: Final = _only(_native(sunk, "post")) + assert all(record["pid"] == worker for record in (pre, during, post)), sunk + assert all(record["litellm_call_id"] == call_id for record in (pre, during, post)), sunk + assert pre["call_type"] == "call_mcp_tool" and during["call_type"] == "call_mcp_tool", sunk + assert pre["messages"] == _synthetic_message("lookup", arguments), pre + assert during["messages"] == _synthetic_message("lookup", arguments), during + assert (pre["mcp_tool_name"], pre["mcp_arguments"]) == ("lookup", arguments), pre + assert pre["mcp_tool_description"] == LOOKUP_DESCRIPTION, pre + assert pre["mcp_input_schema"] == LOOKUP_SCHEMA, pre + assert (during["mcp_tool_description"], during["mcp_input_schema"]) == (None, None), during + assert _texts(post["content"]) == ["found"], post + row: Final = _spend_row(call_id) + assert row["status"] == "success", row + assert _tool_metadata(row)["name"] == "lookup", row + metadata: Final = row["metadata"] + assert isinstance(metadata, dict) and metadata["applied_guardrails"] == [hooks_rig.guardrail], metadata + + +def test_pre_call_mask_reaches_the_peer_and_post_call_mask_reaches_the_caller_on_one_call_id( + hooks_rig: HooksRig, +) -> None: + with ( + scripted_peer(_lookup_tool(f"found {MASK_ME}")) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "mask" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + name: Final = next( + tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool + ) + peer.drain() + hooks_rig.sunk() + outcome: Final = caller.call(name, {"record": MASK_ME}) + assert outcome.text == f"found {MASKED}", outcome.raw + assert _worker(pinned) == worker + reached: Final = tool_calls(peer.drain()) + assert len(reached) == 1, reached + params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"]) + assert params["arguments"] == {"record": MASKED}, params + sunk: Final = hooks_rig.sunk() + generic: Final = _scan(sunk, _only(_native(sunk, "pre"))["litellm_call_id"]) + scanned: Final = generic["texts"] + assert isinstance(scanned, list) and scanned[-1] == MASK_ME and MASKED not in scanned, generic + assert _only(_generic(sunk, "response"))["texts"] == [f"found {MASK_ME}"], sunk + during: Final = _only(_native(sunk, "during")) + assert during["mcp_arguments"] == {"record": MASKED}, during + assert during["messages"] == _synthetic_message("lookup", {"record": MASKED}), during + assert _only(_native(sunk, "post"))["litellm_call_id"] == generic["litellm_call_id"], sunk + assert _spend_row(generic["litellm_call_id"])["status"] == "success" + + +def test_call_time_description_and_schema_come_from_the_catalog_of_the_worker_that_served_the_listing( + hooks_rig: HooksRig, +) -> None: + with ( + scripted_peer(_lookup_tool()) as peer, + hooks_rig.gateway.scenario() as scenario, + _pinned(hooks_rig.sibling) as (second, second_worker), + ): + alias: Final = "local" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + other: Final = McpCaller(second, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(other.initialize, lambda outcome: outcome.ok, seconds=60).ok + with _pinned(hooks_rig.gateway) as (first, first_worker): + assert first_worker != second_worker + lister: Final = McpCaller(first, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(lister.initialize, lambda outcome: outcome.ok, seconds=30).ok + name: Final = next( + tool for tool in eventually(lister.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool + ) + hooks_rig.sunk() + assert other.call(name, {"record": "r-2"}).text == "found" + assert _worker(second) == second_worker + elsewhere: Final = hooks_rig.sunk() + unlisted: Final = _only(_native(elsewhere, "pre")) + assert unlisted["pid"] == second_worker, unlisted + assert (unlisted["mcp_tool_description"], unlisted["mcp_input_schema"]) == (None, None), unlisted + assert _scan(elsewhere, unlisted["litellm_call_id"])["texts"] == ["r-2"], elsewhere + assert lister.call(name, {"record": "r-3"}).text == "found" + assert _worker(first) == first_worker + at_lister: Final = hooks_rig.sunk() + listed: Final = _only(_native(at_lister, "pre")) + assert listed["pid"] == first_worker, listed + assert (listed["mcp_tool_description"], listed["mcp_input_schema"]) == (LOOKUP_DESCRIPTION, LOOKUP_SCHEMA) + assert _scan(at_lister, listed["litellm_call_id"])["texts"] == [ + LOOKUP_DESCRIPTION, + "record identifier", + "r-3", + ] + assert _has_lookup(other.list_tools()) + hooks_rig.sunk() + assert other.call(name, {"record": "r-4"}).text == "found" + assert _worker(second) == second_worker + populated: Final = _only(_native(hooks_rig.sunk(), "pre")) + assert populated["pid"] == second_worker, populated + assert populated["mcp_tool_description"] == LOOKUP_DESCRIPTION, populated + + +@pytest.mark.parametrize("entry", ("mcp", "rest")) +@pytest.mark.parametrize("listed", (False, True)) +@pytest.mark.parametrize("record", ("r-1", MASK_ME)) +def test_long_descriptions_do_not_refuse_small_tpm_calls_or_change_masked_message_bytes( + hooks_rig: HooksRig, entry: EntryPoint, listed: bool, record: str +) -> None: + description: Final = "Gateway tool metadata. " * 300 + tool: Final = ScriptedTool( + "lookup", lambda _: text_result("found"), description=description, input_schema=LOOKUP_SCHEMA + ) + with ( + scripted_peer(tool) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "quota" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}, tpm_limit=64) + caller: Final = McpCaller(pinned, key, entry, headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + if listed: + listing: Final = eventually(lambda: caller.list_tools(identity), _has_lookup, seconds=30) + assert _has_lookup(listing), listing.raw + hooks_rig.sunk() + peer.drain() + arguments: Final = {"record": record} + outcome: Final = caller.call(f"{alias}-lookup", arguments, identity) + assert outcome.ok and outcome.text == "found", outcome.raw + assert _worker(pinned) == worker + reached: Final = tool_calls(peer.drain()) + assert len(reached) == 1, reached + params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"]) + masked_arguments: Final = {"record": record.replace(MASK_ME, MASKED)} + assert params["arguments"] == masked_arguments, params + sunk: Final = hooks_rig.sunk() + pre: Final = _only(tuple(hook for hook in _native(sunk, "pre") if hook["mcp_tool_name"] == "lookup")) + during: Final = _only(tuple(hook for hook in _native(sunk, "during") if hook["mcp_tool_name"] == "lookup")) + assert pre["messages"] == _synthetic_message("lookup", arguments), pre + assert during["messages"] == _synthetic_message("lookup", masked_arguments), during + assert _spend_row(pre["litellm_call_id"])["status"] == "success" + + +def _nested_schema(levels: int) -> dict[str, object]: + if levels == 0: + return {"type": "string", "description": "deepest leaf"} + return {"type": "object", "properties": {"a": _nested_schema(levels - 1)}} + + +def test_a_schema_past_the_scan_depth_is_not_published_while_one_at_the_limit_is_scanned_on_the_call( + hooks_rig: HooksRig, +) -> None: + shallow: Final = ScriptedTool("shallow", lambda _: text_result("found"), input_schema=_nested_schema(49)) + deep: Final = ScriptedTool("deep", lambda _: text_result("found"), input_schema=_nested_schema(50)) + with ( + scripted_peer(shallow, deep) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "depth" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + listing: Final = eventually(caller.list_tools, lambda outcome: len(outcome.tools) > 0, seconds=30) + assert listing.tools == (f"{alias}-shallow",), listing.raw + assert _worker(pinned) == worker + hooks_rig.sunk() + assert caller.call(f"{alias}-shallow", {"record": "r-1"}).text == "found" + scanned: Final = _only(_generic(hooks_rig.sunk(), "request")) + assert scanned["texts"] == ["deepest leaf", "r-1"], scanned + unpublished: Final = caller.call(f"{alias}-deep", {"record": "r-2"}) + assert _worker(pinned) == worker + assert unpublished.error is not None and "Tool 'deep' not found" in unpublished.raw, unpublished.raw + reached: Final = tool_calls(peer.drain()) + assert [object_value(JSON_OBJECT.validate_python(call["body"])["params"])["name"] for call in reached] == [ + "shallow" + ] + relisted: Final = _generic(hooks_rig.sunk(), "request") + assert all((scan["mcp_tool_name"], scan["litellm_call_id"]) == ("shallow", None) for scan in relisted), relisted + rows: Final = _rows(key, 3) + assert [(row["call_type"], row["status"]) for row in rows] == [ + ("list_mcp_tools", "success"), + ("call_mcp_tool", "success"), + ("call_mcp_tool", "failure"), + ], rows + assert [_tool_metadata(row)["name"] for row in rows[1:]] == ["shallow", "deep"], rows + failed: Final = object_value(JSON_OBJECT.validate_python(rows[2]["metadata"])["error_information"]) + assert failed["error_message"] == "404: Tool 'deep' not found", failed + + +def test_an_adapter_returning_the_wrong_number_of_texts_fails_closed_before_the_peer(hooks_rig: HooksRig) -> None: + with ( + scripted_peer(_lookup_tool()) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "count" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + name: Final = next( + tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool + ) + assert _worker(pinned) == worker + hooks_rig.sunk() + blocked: Final = caller.call(name, {"record": COUNT_MISMATCH}) + assert _worker(pinned) == worker + assert blocked.error is not None, blocked.raw + assert ( + "guardrail returned 4 texts for 3 MCP tool strings, so the redaction cannot be mapped back" in blocked.raw + ) + assert tool_calls(peer.drain()) == (), "the blocked call reached the peer" + sunk: Final = hooks_rig.sunk() + assert _only(_generic(sunk, "request"))["texts"] == [LOOKUP_DESCRIPTION, "record identifier", COUNT_MISMATCH] + assert _native(sunk, "post") == (), sunk + rows: Final = _rows(key, 2) + assert [(row["call_type"], row["status"]) for row in rows] == [ + ("list_mcp_tools", "success"), + ("call_mcp_tool", "failure"), + ], rows + assert _tool_metadata(rows[1])["arguments"] == {"record": COUNT_MISMATCH}, rows[1] diff --git a/tests/integration/mcp/test_mcp_credentials.py b/tests/integration/mcp/test_mcp_credentials.py index 95a5e46646b..16dcaab274a 100644 --- a/tests/integration/mcp/test_mcp_credentials.py +++ b/tests/integration/mcp/test_mcp_credentials.py @@ -1,21 +1,32 @@ import base64 +import re +import textwrap import uuid +from collections.abc import Iterator +from pathlib import Path from typing import Final import pytest -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, EntryPoint, McpCaller, McpPeer, + Outcome, + ScriptedTool, call_tool, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, tool_names, ) +from integration._support.oauth_server import oauth_server +from integration._support.process import owned_proxy +from pydantic import TypeAdapter ADD: Final = {"a": 2, "b": 3} STATIC_MODES: Final = ( @@ -192,3 +203,288 @@ def test_byok_server_uses_the_calling_users_stored_credential_and_fails_closed_w assert removed.status_code in (200, 204), removed.text eventually(lambda: call_tool(gateway, owner_key, identity, name, ADD), lambda value: value.status_code == 401) assert tool_calls(peer.drain()) == () + + +def _listings(peer: McpPeer) -> tuple[dict[str, object], ...]: + return tuple( + item + for item in peer.drain() + if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list" + ) + + +def test_oauth2_byok_listing_sends_the_minted_token_not_the_users_stored_secret(gateway: Gateway) -> None: + with mcp_peer() as peer, oauth_server() as auth, gateway.scenario() as scenario: + alias: Final = "cc" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, + peer, + alias, + auth_type="oauth2", + oauth2_flow="client_credentials", + is_byok=True, + token_url=auth.issuer + "/token", + credentials={"client_id": "cc-client", "client_secret": "cc-secret-" + uuid.uuid4().hex}, + ) + owner: Final = scenario.user() + owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]}) + secret: Final = "byok-" + uuid.uuid4().hex + stored: Final = gateway.client.post( + f"/v1/mcp/server/{identity}/user-credential", + json={"credential": secret}, + headers={"x-litellm-api-key": owner_key}, + ) + assert stored.status_code in (200, 201), stored.text + scenario.cleanups.callback( + gateway.client.delete, + f"/v1/mcp/server/{identity}/user-credential", + headers={"x-litellm-api-key": owner_key}, + ) + peer.drain() + auth.drain() + response: Final = gateway.client.get( + "/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key} + ) + assert response.status_code == 200, response.text + assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text + assert [request["grant_type"] for request in auth.token_requests()] == ["client_credentials"] + listings: Final = _listings(peer) + assert len(listings) == 1, listings + sent: Final = _header(listings[0], b"authorization") + assert sent is not None and auth.is_live(sent.decode().removeprefix("Bearer ")), sent + assert secret.encode() not in sent, "stored BYOK secret replaced the minted token on tools/list" + + +@pytest.mark.parametrize(("auth_type", "header", "shape"), STATIC_MODES[:2]) +def test_byok_rest_listing_sends_the_servers_static_credential_not_the_users_stored_secret( + gateway: Gateway, auth_type: str, header: bytes, shape: str +) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "byok" + uuid.uuid4().hex[:8] + static: Final = "static-" + uuid.uuid4().hex + identity: Final = register_mcp( + scenario, peer, alias, auth_type=auth_type, is_byok=True, credentials={"auth_value": static} + ) + owner: Final = scenario.user() + owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]}) + secret: Final = "byok-" + uuid.uuid4().hex + stored: Final = gateway.client.post( + f"/v1/mcp/server/{identity}/user-credential", + json={"credential": secret}, + headers={"x-litellm-api-key": owner_key}, + ) + assert stored.status_code in (200, 201), stored.text + scenario.cleanups.callback( + gateway.client.delete, + f"/v1/mcp/server/{identity}/user-credential", + headers={"x-litellm-api-key": owner_key}, + ) + peer.drain() + response: Final = gateway.client.get( + "/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key} + ) + assert response.status_code == 200, response.text + assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text + listings: Final = _listings(peer) + assert len(listings) == 1, listings + assert _header(listings[0], header) == shape.format(secret=static, basic="").encode(), listings[0]["headers"] + peer.drain() + called: Final = call_tool(gateway, owner_key, identity, f"{alias}-add", ADD) + assert called.status_code == 200, called.text + assert _header(_one_call(peer), header) == shape.format(secret=secret, basic="").encode() + + +def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_user(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "byok" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token", is_byok=True) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + peer.drain() + response: Final = gateway.client.get( + "/mcp-rest/tools/list", + params={"server_id": identity}, + headers={"x-litellm-api-key": key, "x-mcp-auth": "Bearer hdr"}, + ) + assert response.status_code == 200, response.text + names: Final = {tool["name"] for tool in response.json()["tools"]} + assert "add" in names, names + listings: Final = _listings(peer) + assert len(listings) == 1, listings + assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr" + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +_PROBE_ARGUMENTS: Final = {"probe": _PROBE} +_HEADERS: Final = TypeAdapter(dict[str, str]) + + +def _store_byok_credential(scenario: Scenario, identity: str, key: str, secret: str) -> None: + stored: Final = scenario.gateway.client.post( + f"/v1/mcp/server/{identity}/user-credential", json={"credential": secret}, headers={"x-litellm-api-key": key} + ) + assert stored.status_code in (200, 201), stored.text + scenario.cleanups.callback( + scenario.gateway.client.delete, f"/v1/mcp/server/{identity}/user-credential", headers={"x-litellm-api-key": key} + ) + + +def test_rotating_the_credential_drops_the_callers_listing_until_it_lists_again(echo_rig: Gateway) -> None: + with mcp_peer() as peer, echo_rig.scenario() as scenario: + first: Final = "cred-" + uuid.uuid4().hex + second: Final = "cred-" + uuid.uuid4().hex + alias: Final = "rot" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, auth_type="bearer_token", credentials={"auth_value": first} + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(echo_rig, key, "mcp", alias) + assert caller.list_tools().ok + assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers" + rotated: Final = echo_rig.request( + "PUT", "/v1/mcp/server", {"server_id": identity, "credentials": {"auth_value": second}} + ) + assert rotated.status_code == 202, rotated.text + eventually( + lambda: _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)), lambda seen: seen == _UNLISTED + ) + peer.drain() + assert caller.list_tools().ok + assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers" + relisted: Final = _listings(peer) + assert len(relisted) == 1, relisted + assert _header(relisted[0], b"authorization") == f"Bearer {second}".encode(), relisted[0]["headers"] + + +def test_byok_callers_are_evaluated_against_their_own_listing_and_the_stored_secret_never_keys_the_slot( + echo_rig: Gateway, +) -> None: + with mcp_peer() as peer, echo_rig.scenario() as scenario: + alias: Final = "byok" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, auth_type="api_key", is_byok=True) + owner_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]}) + stranger_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]}) + owner_secret: Final = "byok-" + uuid.uuid4().hex + replacement: Final = "byok-" + uuid.uuid4().hex + _store_byok_credential(scenario, identity, owner_key, owner_secret) + _store_byok_credential(scenario, identity, stranger_key, "byok-" + uuid.uuid4().hex) + owner: Final = McpCaller(echo_rig, owner_key, "mcp", alias) + stranger: Final = McpCaller(echo_rig, stranger_key, "mcp", alias) + peer.drain() + assert owner.list_tools().ok + listings: Final = _listings(peer) + assert [_header(item, b"x-api-key") for item in listings] == [owner_secret.encode()], listings + own: Final = _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS)) + other: Final = _echoed_description(stranger.call(f"{alias}-add", _PROBE_ARGUMENTS)) + assert (own, other) == ("Add two integers", _UNLISTED), (own, other) + _store_byok_credential(scenario, identity, owner_key, replacement) + assert _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers", ( + "the slot is keyed by the client-supplied header, never by the stored credential" + ) + sent: Final = eventually( + lambda: (owner.call(f"{alias}-add", ADD).ok, tool_calls(peer.drain())), + lambda value: any(_header(call, b"x-api-key") == replacement.encode() for call in value[1]), + ) + assert sent[0], sent + + +def test_callers_with_different_server_scoped_auth_headers_are_evaluated_against_their_own_listings( + echo_rig: Gateway, +) -> None: + tool: Final = ScriptedTool( + "add", + lambda _: text_result("3"), + description=lambda headers: "Adds for " + headers.get("authorization", "nobody"), + ) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "scoped" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + acme_token: Final = "acme-" + uuid.uuid4().hex + globex_token: Final = "globex-" + uuid.uuid4().hex + acme: Final = McpCaller(echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {acme_token}"}) + globex: Final = McpCaller( + echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {globex_token}"} + ) + assert acme.list_tools().ok and globex.list_tools().ok + seen: Final = ( + _echoed_description(acme.call(f"{alias}-add", _PROBE_ARGUMENTS)), + _echoed_description(globex.call(f"{alias}-add", _PROBE_ARGUMENTS)), + ) + assert seen == (f"Adds for Bearer {acme_token}", f"Adds for Bearer {globex_token}"), seen + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" + + +def test_deprecated_string_x_mcp_auth_callers_on_a_user_less_key_own_separate_listings(echo_rig: Gateway) -> None: + tool: Final = ScriptedTool( + "add", + lambda _: text_result("3"), + description=lambda headers: "Adds for " + headers.get("authorization", "nobody"), + ) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "legacy" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token") + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + first_token: Final = "first-" + uuid.uuid4().hex + second_token: Final = "second-" + uuid.uuid4().hex + first: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {first_token}"}) + second: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {second_token}"}) + assert _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED + peer.drain() + assert first.list_tools().ok + listings: Final = _listings(peer) + assert len(listings) == 1, listings + listed_with: Final = _HEADERS.validate_python(listings[0]["headers"]) + assert listed_with.get("authorization") == f"Bearer {first_token}", listed_with + warm: Final = eventually( + lambda: _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)), + lambda seen: seen != _UNLISTED, + ) + assert warm == f"Adds for Bearer {first_token}", warm + assert _echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED + assert second.list_tools().ok + seen: Final = eventually( + lambda: ( + _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)), + _echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)), + ), + lambda pair: _UNLISTED not in pair, + ) + assert seen == (f"Adds for Bearer {first_token}", f"Adds for Bearer {second_token}"), seen + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" diff --git a/tests/integration/mcp/test_mcp_lifecycle.py b/tests/integration/mcp/test_mcp_lifecycle.py index fa253f03520..eed3cbf658e 100644 --- a/tests/integration/mcp/test_mcp_lifecycle.py +++ b/tests/integration/mcp/test_mcp_lifecycle.py @@ -1,7 +1,9 @@ import functools import json import uuid +from collections.abc import Mapping from contextlib import ExitStack +from hashlib import sha256 from pathlib import Path from typing import Final @@ -14,15 +16,27 @@ from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests from integration._support.mcp import ( + JsonRpc, McpCaller, Outcome, + ScriptedTool, call_tool, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, tool_names, ) from integration._support.process import owned_proxy +from pydantic import TypeAdapter + +_SPEND_NONCES: Final = ( + "SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce" + ' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s' +) +_OBJECTS: Final = TypeAdapter(Mapping[str, object]) +_STRINGS: Final = TypeAdapter(Mapping[str, str]) @pytest.mark.covers("mcp.call_tool.saved_headers.reach_actual_transport") @@ -463,9 +477,7 @@ def _update_tool_permissions( assert updated.status_code == 200, updated.text -def _listing_on_both( - gateway: Gateway, peer: Gateway, key: str, expected: set[str] -) -> None: +def _listing_on_both(gateway: Gateway, peer: Gateway, key: str, expected: set[str]) -> None: for worker in (gateway, peer): listing: Final = eventually( functools.partial(_granted_view, worker, key), @@ -525,3 +537,56 @@ def test_key_update_tool_permission_widen_narrow_and_clear_apply_on_both_workers _listing_on_both(gateway, peer, key, all_tools) nulled: Final = _multiply_outcome_on_both(gateway, peer, key, alias) assert [call.text for call in nulled] == ["6", "6"], [call.raw for call in nulled] + + +def _nonce_echo(params: JsonRpc) -> JsonRpc: + return text_result(_STRINGS.validate_python(params["arguments"])["nonce"]) + + +def _call_params(call: Mapping[str, object]) -> Mapping[str, object]: + return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"]) + + +def _listed(caller: McpCaller, name: str) -> None: + listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45) + assert listing.error is None, (caller.gateway.client.base_url, listing.raw) + + +def test_tool_calls_on_both_workers_stay_base_compatible_after_each_worker_lists( + gateway: Gateway, peer: Gateway +) -> None: + schema: Final = {"type": "object", "properties": {"nonce": {"type": "string"}}} + tool: Final = ScriptedTool("echo", _nonce_echo, description="Echo the nonce back", input_schema=schema) + with scripted_peer(tool) as upstream, gateway.scenario() as scenario: + alias: Final = "compat" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, upstream, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = f"{alias}-echo" + first: Final = McpCaller(gateway, key, "mcp", alias) + second: Final = McpCaller(peer, key, "mcp", alias) + _listed(first, name) + _listed(second, name) + upstream.drain() + nonces: Final = (uuid.uuid4().hex, uuid.uuid4().hex) + outcomes: Final = ( + first.call(name, {"nonce": nonces[0]}), + second.call(name, {"nonce": nonces[1]}), + first.call(name, {"nonce": nonces[0]}), + ) + assert [outcome.text for outcome in outcomes] == [nonces[0], nonces[1], nonces[0]], [o.raw for o in outcomes] + params: Final = [_call_params(call) for call in tool_calls(upstream.drain())] + assert [set(entry) - {"_meta"} for entry in params] == [{"name", "arguments"}] * 3, params + assert [entry["name"] for entry in params] == ["echo"] * 3, params + assert [entry["arguments"] for entry in params] == [ + {"nonce": nonces[0]}, + {"nonce": nonces[1]}, + {"nonce": nonces[0]}, + ], params + rows: Final = eventually( + lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")), + lambda found: len(found) >= 3, + seconds=70, + ) + assert sorted((str(row["status"]), str(row["nonce"])) for row in rows) == sorted( + ("success", nonce) for nonce in (nonces[0], nonces[1], nonces[0]) + ), rows diff --git a/tests/integration/mcp/test_mcp_listed_tool_metadata.py b/tests/integration/mcp/test_mcp_listed_tool_metadata.py new file mode 100644 index 00000000000..d09864da4b1 --- /dev/null +++ b/tests/integration/mcp/test_mcp_listed_tool_metadata.py @@ -0,0 +1,531 @@ +"""pre_mcp_call guardrails are handed the tool entry ``tools/list`` served to the caller. + +One owned proxy carries a default-on ``custom_code`` pre_mcp_call guardrail. At listing time it masks +``SECRET`` out of every scanned text. At call time, when an argument carries the probe marker, it +blocks and echoes the description and parameters it was handed, which is the only way to observe from outside +what metadata the gateway attached to the hook +""" + +import json +import threading +import uuid +from collections.abc import Callable, Generator, Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack, contextmanager +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.mcp import ( + EntryPoint, + JsonRpc, + McpCaller, + ScriptedTool, + listed_tools, + openapi_peer, + register_mcp, + scripted_peer, + text_result, + tool_calls, +) +from integration._support.process import owned_proxy +from pydantic import TypeAdapter + +_ECHO: Final = "catalog-echo:" +_PROBE: Final = "catalog-probe" +_CALLERS_PER_SERVER: Final = 256 +_PIN_SECONDS: Final = 120 +_COLD: Final[tuple[str, JsonRpc]] = ("", {"type": "object", "properties": {}, "additionalProperties": False}) +_PID: Final = TypeAdapter(int) +_HEADERS: Final = TypeAdapter(dict[bytes, bytes]) +_SCHEMA: Final = TypeAdapter(dict[str, object]) +_LOOKUP_SCHEMA: Final = { + "type": "object", + "properties": {"probe": {"type": "string", "description": "a probe marker"}}, + "additionalProperties": False, +} +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + ' texts = list(inputs.get("texts") or [])\n' + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' if "{_PROBE}" in texts:\n' + f' return block("{_ECHO}" + json_stringify(' + '{"description": function.get("description"), "parameters": function.get("parameters")}))\n' + ' masked = [text.replace("SECRET", "[MASKED]") for text in texts]\n' + " if masked != texts:\n" + " return modify(texts=masked)\n" + " return allow()\n" +) + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("listed-tool-metadata") + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8], + "litellm_params": { + "guardrail": "custom_code", + "mode": "pre_mcp_call", + "default_on": True, + "custom_code": _GUARDRAIL_CODE, + }, + } + ] + config["general_settings"] = {**config["general_settings"], "proxy_config_reload_interval_seconds": 1} + path: Final = directory / "config.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, directory, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate, + ExitStack() as stack, + ): + _two_workers(stack, candidate) + yield candidate + + +def _strings(value: object) -> Iterator[str]: + if isinstance(value, str): + yield value + return + children: Final = value.values() if isinstance(value, Mapping) else value if isinstance(value, list) else () + for child in children: + yield from _strings(child) + + +def _decoded(raw: str) -> object: + data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:")) + return json.loads(data[-1] if data else raw) + + +def _echoed(raw: str) -> tuple[str | None, Mapping[str, object] | None]: + """The (description, parameters) the guardrail was handed, recovered from its block reason.""" + carrier: Final = next((text for text in _strings(_decoded(raw)) if _ECHO in text), None) + assert carrier is not None, raw + echoed, _ = json.JSONDecoder().raw_decode(carrier.split(_ECHO, 1)[1]) + assert isinstance(echoed, dict), carrier + return echoed.get("description"), echoed.get("parameters") + + +def _probe(caller: McpCaller, name: str, server_id: str) -> tuple[str | None, Mapping[str, object] | None]: + outcome: Final = caller.call(name, {"probe": _PROBE}, server_id=server_id) + assert outcome.error is not None, outcome.raw + return _echoed(outcome.raw) + + +def _worker(gateway: Gateway) -> int: + response: Final = gateway.request("GET", "/debug/memory/summary") + assert response.status_code == 200, response.text + return _PID.validate_python(response.json()["worker_pid"]) + + +@contextmanager +def _pinned(rig: Gateway) -> Generator[Gateway, None, None]: + """A single keep-alive connection, so every request on it is served by the worker that accepted it.""" + limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1) + with httpx.Client(base_url=rig.client.base_url, timeout=15, trust_env=False, limits=limits) as client: + yield Gateway(client, rig.key, rig.upstream_url) + + +def _connection(stack: ExitStack, rig: Gateway, wanted: Callable[[int], bool]) -> tuple[Gateway, int]: + """A pinned connection to a worker ``wanted`` accepts; a worker still starting up accepts nothing yet.""" + + def attempt() -> tuple[Gateway, int] | None: + with ExitStack() as candidate: + gateway: Final = candidate.enter_context(_pinned(rig)) + pid: Final = _worker(gateway) + if not wanted(pid): + return None + stack.enter_context(candidate.pop_all()) + return gateway, pid + + found: Final = eventually(attempt, lambda pair: pair is not None, seconds=_PIN_SECONDS) + assert found is not None + return found + + +def _two_workers(stack: ExitStack, rig: Gateway) -> tuple[tuple[Gateway, int], tuple[Gateway, int]]: + """Connects while the first worker is busy answering, so the idle worker wins the accept race.""" + first: Final = _connection(stack, rig, lambda _: True) + stop: Final = threading.Event() + + def keep_busy() -> None: + while not stop.is_set(): + _worker(first[0]) + + with ThreadPoolExecutor(max_workers=1) as pool: + busy: Final = pool.submit(keep_busy) + try: + other: Final = _connection(stack, rig, lambda pid: pid != first[1]) + finally: + stop.set() + busy.result() + return first, other + + +def _served_name(gateway: Gateway, key: str, identity: str, tool: str) -> str: + """The prefixed name this worker lists for ``tool`` once its registry reload carries the server.""" + listing: Final = eventually( + lambda: McpCaller(gateway, key, "rest").list_tools(server_id=identity), + lambda value: value.ok and any(full.endswith(tool) for full in value.tools), + ) + return next(full for full in listing.tools if full.endswith(tool)) + + +def _settled_probe( + gateway: Gateway, key: str, name: str, identity: str +) -> tuple[str | None, Mapping[str, object] | None]: + """The hook echo for a direct call, once this worker's registry reload carries the server.""" + outcome: Final = eventually( + lambda: McpCaller(gateway, key, "rest").call(name, {"probe": _PROBE}, server_id=identity), + lambda value: value.error is not None and _ECHO in value.raw, + ) + return _echoed(outcome.raw) + + +def _forwarded_tenants(observed: tuple[dict[str, object], ...]) -> frozenset[bytes]: + return frozenset(_HEADERS.validate_python(item["headers"]).get(b"x-tenant", b"") for item in observed) - {b""} + + +def _lookup_tool(description: str | Callable[[Mapping[str, str]], str] = "Look up one record") -> ScriptedTool: + return ScriptedTool("lookup", lambda _: text_result("found"), description=description, input_schema=_LOOKUP_SCHEMA) + + +@pytest.mark.parametrize("entry", ["rest", "mcp"]) +def test_pre_call_hook_receives_the_description_and_input_schema_the_caller_was_listed( + rig: Gateway, entry: EntryPoint +) -> None: + schema: Final = {"type": "object", "properties": {"probe": {"type": "string", "description": "a probe marker"}}} + tool: Final = ScriptedTool( + "lookup", lambda _: text_result("found"), description="Look up one record", input_schema=schema + ) + with scripted_peer(tool) as peer, rig.scenario() as scenario: + alias: Final = "meta" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(rig, key, entry, headers={"x-mcp-servers": alias}) + assert caller.initialize().ok + listed: Final = caller.list_tools(server_id=identity) + assert listed.ok, listed.raw + name: Final = next(full for full in listed.tools if full.endswith("lookup")) + description, parameters = _probe(caller, name, identity) + assert description == "Look up one record", (description, parameters) + assert parameters is not None and parameters.get("properties") == schema["properties"], parameters + + +def test_each_caller_is_evaluated_against_the_catalog_its_own_forwarded_headers_produced(rig: Gateway) -> None: + tool: Final = ScriptedTool( + "report", + lambda _: text_result("ok"), + description=lambda headers: f"Report for tenant {headers.get('x-tenant', 'nobody')}", + ) + with scripted_peer(tool) as peer, rig.scenario() as scenario: + alias: Final = "tenant" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + acme: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "acme"}) + globex: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "globex"}) + acme_listing: Final = acme.list_tools() + globex_listing: Final = globex.list_tools() + assert acme_listing.ok and globex_listing.ok, (acme_listing.raw, globex_listing.raw) + name: Final = next(full for full in acme_listing.tools if full.endswith("report")) + acme_seen, _ = _probe(acme, name, identity) + globex_seen, _ = _probe(globex, name, identity) + assert (acme_seen, globex_seen) == ("Report for tenant acme", "Report for tenant globex"), ( + "each caller's tools/call must be evaluated against the catalog its own headers listed" + ) + + +def test_call_is_evaluated_against_the_masked_description_the_listing_served(rig: Gateway) -> None: + tool: Final = ScriptedTool("read_note", lambda _: text_result("note"), description="Read a note") + with scripted_peer(tool) as peer, rig.scenario() as scenario: + alias: Final = "note" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, tool_name_to_description={"read_note": "Read a SECRET note"} + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(rig, key, identity) + name: Final = next(full for full in served if full.endswith("read_note")) + assert served[name]["description"] == "Read a [MASKED] note", served[name] + seen, _ = _probe(McpCaller(rig, key, "rest"), name, identity) + assert seen == "Read a [MASKED] note", "the admin override must not restore wording the listing masked" + + +def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_served(rig: Gateway) -> None: + with openapi_peer() as peer, rig.scenario() as scenario: + alias: Final = "pets" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"} + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(rig, key, identity) + name: Final = next(full for full in served if full.endswith("getpet")) + assert served[name]["description"] == "Fetch one [MASKED] pet", served[name] + seen, parameters = _probe(McpCaller(rig, key, "rest"), name, identity) + assert seen == "Fetch one [MASKED] pet", "the OpenAPI call path must hand hooks the entry the listing served" + assert parameters is not None and "petId" in parameters.get("properties", {}), parameters + assert not [call for call in peer.drain() if call["path"].startswith("/pets")], "blocked before upstream" + + +def test_openapi_call_is_evaluated_against_the_entry_this_key_was_listed_not_the_last_listing( + rig: Gateway, +) -> None: + with openapi_peer() as peer, rig.scenario() as scenario: + alias: Final = "pets" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"} + ) + guarded: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + opted_out: Final = scenario.key( + object_permission={"mcp_servers": [identity]}, metadata={"disable_global_guardrails": True} + ) + guarded_served: Final = listed_tools(rig, guarded, identity) + opted_out_served: Final = listed_tools(rig, opted_out, identity) + name: Final = next(full for full in guarded_served if full.endswith("getpet")) + assert (guarded_served[name]["description"], opted_out_served[name]["description"]) == ( + "Fetch one [MASKED] pet", + "Fetch one SECRET pet", + ), (guarded_served[name], opted_out_served[name]) + seen, _ = _probe(McpCaller(rig, guarded, "rest"), name, identity) + assert seen == "Fetch one [MASKED] pet", ( + "the guarded key must be evaluated against its own listing, not the opted-out key's later one" + ) + + +def test_direct_call_without_a_listing_hands_the_hook_no_metadata_on_either_worker(rig: Gateway) -> None: + with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack: + (first, first_pid), (other, other_pid) = _two_workers(stack, rig) + with first.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "cold" + uuid.uuid4().hex[:8]) + observer: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = _served_name(first, observer, identity, "lookup") + seen: Final = tuple(_settled_probe(gateway, key, name, identity) for gateway in (first, other)) + assert seen == (_COLD, _COLD), (seen, first_pid, other_pid) + assert (_worker(first), _worker(other)) == (first_pid, other_pid) + assert tool_calls(peer.drain()) == (), "the probe is blocked at the hook, before the upstream" + + +def test_warm_metadata_is_local_to_the_worker_that_served_the_listing(rig: Gateway) -> None: + with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack: + (first, first_pid), (other, other_pid) = _two_workers(stack, rig) + with first.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "local" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = _served_name(first, key, identity, "lookup") + warm: Final = ("Look up one record", _LOOKUP_SCHEMA) + assert _probe(McpCaller(first, key, "rest"), name, identity) == warm, first_pid + assert _settled_probe(other, key, name, identity) == _COLD, ( + "a worker that never served this caller a listing has no catalog for it", + other_pid, + ) + assert _served_name(other, key, identity, "lookup") == name + assert _probe(McpCaller(other, key, "rest"), name, identity) == warm, other_pid + assert (_worker(first), _worker(other)) == (first_pid, other_pid) + assert tool_calls(peer.drain()) == () + + +def test_listed_tools_without_a_description_still_hand_the_hook_the_schema_the_listing_served(rig: Gateway) -> None: + open_schema: Final[JsonRpc] = {"type": "object", "properties": {}, "additionalProperties": True} + undescribed: Final = ScriptedTool("undescribed", lambda _: text_result("ok"), input_schema=open_schema) + blank: Final = ScriptedTool("blank", lambda _: text_result("ok"), description="", input_schema=open_schema) + with scripted_peer(undescribed, blank) as peer, _pinned(rig) as worker, worker.scenario() as scenario: + pid: Final = _worker(worker) + identity: Final = register_mcp(scenario, peer, "bare" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(worker, key, identity) + names: Final = tuple(next(full for full in served if full.endswith(tool)) for tool in ("undescribed", "blank")) + assert tuple((served[name].get("description") or "", served[name]["inputSchema"]) for name in names) == ( + ("", open_schema), + ("", open_schema), + ), served + seen: Final = tuple(_probe(McpCaller(worker, key, "rest"), name, identity) for name in names) + assert seen == (("", open_schema), ("", open_schema)), (seen, _COLD) + assert _worker(worker) == pid + assert tool_calls(peer.drain()) == () + + +def test_hook_receives_the_nested_schema_with_the_leaves_the_listing_masked(rig: Gateway) -> None: + schema: Final = { + "type": "object", + "required": ["filter"], + "additionalProperties": False, + "properties": { + "filter": { + "type": "object", + "properties": { + "path": {"type": "string", "description": "SECRET path"}, + "tags": {"type": "array", "items": {"type": "string", "description": "one SECRET tag"}}, + }, + } + }, + } + masked: Final = _SCHEMA.validate_python(json.loads(json.dumps(schema).replace("SECRET", "[MASKED]"))) + tool: Final = ScriptedTool( + "search", lambda _: text_result("hit"), description="Search SECRET records", input_schema=schema + ) + with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario: + pid: Final = _worker(worker) + identity: Final = register_mcp(scenario, peer, "nested" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(worker, key, identity) + name: Final = next(full for full in served if full.endswith("search")) + assert (served[name]["description"], served[name]["inputSchema"]) == ("Search [MASKED] records", masked), served + assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Search [MASKED] records", masked), ( + "the hook must be handed the nested schema exactly as the listing served it" + ) + assert _worker(worker) == pid + assert tool_calls(peer.drain()) == () + + +def test_a_server_definition_update_drops_the_catalog_on_every_worker_until_the_caller_lists_again( + rig: Gateway, +) -> None: + with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack: + (first, first_pid), (other, other_pid) = _two_workers(stack, rig) + with first.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "upd" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = _served_name(first, key, identity, "lookup") + assert _served_name(other, key, identity, "lookup") == name + before: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other)) + assert before == (("Look up one record", _LOOKUP_SCHEMA),) * 2, before + updated: Final = first.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}}, + ) + assert updated.status_code == 202, updated.text + assert _probe(McpCaller(first, key, "rest"), name, identity) == _COLD, ( + "the worker that applied the update must drop its catalog at once", + first_pid, + ) + assert ( + eventually(lambda: _probe(McpCaller(other, key, "rest"), name, identity), lambda seen: seen == _COLD) + == _COLD + ), other_pid + relisted: Final = tuple( + listed_tools(gateway, key, identity)[name]["description"] for gateway in (first, other) + ) + assert relisted == ("Audited lookup", "Audited lookup"), relisted + after: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other)) + assert after == (("Audited lookup", _LOOKUP_SCHEMA),) * 2, after + assert (_worker(first), _worker(other)) == (first_pid, other_pid) + assert tool_calls(peer.drain()) == () + + +def test_a_listing_in_flight_across_a_server_update_does_not_resurrect_the_old_catalog(rig: Gateway) -> None: + started: Final = threading.Event() + release: Final = threading.Event() + + def describe(_: Mapping[str, str]) -> str: + started.set() + assert release.wait(20), "the listing was never released" + return "Look up one record" + + with ( + scripted_peer(_lookup_tool(describe)) as peer, + ExitStack() as stack, + ThreadPoolExecutor(max_workers=1) as pool, + ): + worker, pid = _connection(stack, rig, lambda _: True) + sibling, _ = _connection(stack, rig, lambda candidate: candidate == pid) + with worker.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "stale" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + release.set() + name: Final = _served_name(worker, key, identity, "lookup") + assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Look up one record", _LOOKUP_SCHEMA) + started.clear() + release.clear() + pending: Final = pool.submit(McpCaller(worker, key, "rest").list_tools, identity) + assert started.wait(10), "the upstream never saw the in-flight listing" + updated: Final = sibling.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}}, + ) + assert updated.status_code == 202, updated.text + release.set() + stale: Final = pending.result(timeout=20) + assert stale.ok, stale.raw + assert _probe(McpCaller(worker, key, "rest"), name, identity) == _COLD, ( + "a listing fetched before the update must not be recorded after it" + ) + assert listed_tools(worker, key, identity)[name]["description"] == "Audited lookup" + assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Audited lookup", _LOOKUP_SCHEMA) + assert (_worker(worker), _worker(sibling)) == (pid, pid) + assert tool_calls(peer.drain()) == () + + +def test_a_server_keeps_the_newest_256_caller_catalogs_and_evicts_the_oldest(rig: Gateway) -> None: + tool: Final = _lookup_tool(lambda headers: f"Lookup for tenant {headers.get('x-tenant', 'nobody')}") + with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario: + pid: Final = _worker(worker) + alias: Final = "cap" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + tenants: Final = tuple(f"t{index}" for index in range(_CALLERS_PER_SERVER + 1)) + callers: Final = { + tenant: McpCaller(worker, key, "server_mcp", alias, headers={"x-tenant": tenant}) for tenant in tenants + } + listings: Final = tuple(callers[tenant].list_tools() for tenant in tenants) + assert all(listing.ok for listing in listings), [listing.raw for listing in listings if not listing.ok] + name: Final = next(full for full in listings[0].tools if full.endswith("lookup")) + + def seen(tenant: str) -> str | None: + return _probe(callers[tenant], name, identity)[0] + + assert (seen("t0"), seen("t1"), seen(tenants[-1])) == ( + "", + "Lookup for tenant t1", + f"Lookup for tenant {tenants[-1]}", + ), "the oldest of 257 callers is evicted, the newest 256 keep their own catalog" + assert callers["t0"].list_tools().ok + assert (seen("t0"), seen("t1")) == ("Lookup for tenant t0", ""), "relisting makes t0 newest and evicts t1" + assert _worker(worker) == pid + observed: Final = peer.drain() + assert _forwarded_tenants(observed) == frozenset(tenant.encode() for tenant in tenants), len(observed) + assert tool_calls(observed) == () + + +def test_an_admin_include_disabled_tools_listing_does_not_warm_the_runtime_catalog(rig: Gateway) -> None: + """``include_disabled_tools=true`` is the admin-only configuration view, not a listing the caller + runs against: recording it would warm tools/call metadata no runtime listing ever served.""" + tool: Final = _lookup_tool() + with scripted_peer(tool) as peer, rig.scenario() as scenario: + alias: Final = "adminview" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + admin: Final = rig.key + caller: Final = McpCaller(rig, admin, "rest") + name: Final = eventually( + lambda: caller.call(f"{alias}-lookup", {"probe": _PROBE}, server_id=identity), + lambda value: value.error is not None and _ECHO in value.raw, + ) + assert _echoed(name.raw) == _COLD, "before any listing the call is cold" + + view: Final = eventually( + lambda: rig.client.get( + "/mcp-rest/tools/list", + params={"server_id": identity, "include_disabled_tools": "true"}, + headers={"x-litellm-api-key": admin}, + ), + lambda response: response.status_code == 200 + and any(entry["name"].endswith("lookup") for entry in response.json()["tools"]), + ) + assert view.status_code == 200, view.text + + after_view: Final = _probe(caller, f"{alias}-lookup", identity) + assert after_view == _COLD, ( + "the admin-only include_disabled_tools view must not record the caller's listed-tools slot" + ) + + runtime: Final = caller.list_tools(server_id=identity) + assert runtime.ok, runtime.raw + assert _probe(caller, f"{alias}-lookup", identity)[0] == "Look up one record", ( + "a genuine runtime listing still warms the slot" + ) diff --git a/tests/integration/mcp/test_mcp_llm_endpoints.py b/tests/integration/mcp/test_mcp_llm_endpoints.py index 01bf006a03e..af3bff1fd8b 100644 --- a/tests/integration/mcp/test_mcp_llm_endpoints.py +++ b/tests/integration/mcp/test_mcp_llm_endpoints.py @@ -1,16 +1,35 @@ +import asyncio import json import uuid -from collections.abc import Callable, Iterator, Mapping, Sequence +from collections.abc import Callable, Generator, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from dataclasses import dataclass -from typing import Final, Literal +from hashlib import sha256 +from pathlib import Path +from typing import Final, Literal, TypeVar import httpx import pytest -from integration._support.client import Gateway, Scenario -from integration._support.mcp import McpPeer, mcp_peer, register_mcp, tool_calls +import yaml +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.mcp import ( + McpPeer, + ScriptedTool, + mcp_peer, + register_mcp, + scripted_peer, + text_result, + tool_calls, +) from integration._support.mcp_grants import create_toolset +from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, Wire, wire_server +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ChatCompletionMessageParam +from openai.types.responses.tool_param import Mcp +from pydantic import JsonValue, TypeAdapter Surface = Literal["chat", "responses", "messages", "messages_bridge"] SURFACES: Final[tuple[Surface, ...]] = ("chat", "responses", "messages", "messages_bridge") @@ -18,6 +37,9 @@ ADD: Final = {"a": 2, "b": 3} ANSWER: Final = "the sum is 5" GATEWAY_REF: Final = {"type": "mcp", "server_url": "litellm_proxy", "server_label": "litellm"} AUTO: Final = {**GATEWAY_REF, "require_approval": "never"} +OUTAGE: Final = "bridge-outage" +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) def _json(body: Mapping[str, object]) -> Reply: @@ -44,25 +66,94 @@ def _has_tool_result(body: Mapping[str, object]) -> bool: return False -def _model_double(tool: str) -> Callable[[Request], Reply]: - arguments: Final = json.dumps(ADD) +@dataclass(frozen=True, slots=True) +class Turn: + tool: str + arguments: str + answer: str + +def _fixed_turn(tool: str) -> Callable[[Mapping[str, JsonValue]], Turn]: + return lambda _: Turn(tool, json.dumps(ADD), ANSWER) + + +def _echoing_turn(body: Mapping[str, JsonValue]) -> Turn: + names: Final = _tool_names(body) + return Turn(names[0] if names else "", json.dumps({"query": _prompt(body)}), _tool_result_text(body) or "") + + +def _prompt(body: Mapping[str, JsonValue]) -> str: + inputs: Final = body.get("input") + if isinstance(inputs, str): + return inputs + items: Final = inputs if isinstance(inputs, list) else body.get("messages") + first: Final = items[0] if isinstance(items, list) and items else None + return str(first["content"]) if isinstance(first, dict) and isinstance(first.get("content"), str) else "" + + +def _tool_result_text(body: Mapping[str, JsonValue]) -> str | None: + inputs: Final = body.get("input") + messages: Final = body.get("messages") + items: Final = inputs if isinstance(inputs, list) else messages if isinstance(messages, list) else () + for item in items: + if not isinstance(item, dict): + continue + if item.get("type") == "function_call_output": + return str(item["output"]) + if item.get("role") == "tool": + return str(item["content"]) + if (block := _tool_result_block(item)) is not None: + return block + return None + + +def _tool_result_block(item: Mapping[str, JsonValue]) -> str | None: + content: Final = item.get("content") + for block in content if isinstance(content, list) else (): + if isinstance(block, dict) and block.get("type") == "tool_result": + return str(block["content"]) + return None + + +def _responses_stream(response: Mapping[str, JsonValue], item: Mapping[str, JsonValue]) -> Reply: + events: Final = ( + {"type": "response.created", "sequence_number": 0, "response": {**response, "status": "in_progress"}}, + {"type": "response.in_progress", "sequence_number": 1, "response": {**response, "status": "in_progress"}}, + {"type": "response.output_item.added", "sequence_number": 2, "output_index": 0, "item": item}, + {"type": "response.output_item.done", "sequence_number": 3, "output_index": 0, "item": item}, + {"type": "response.completed", "sequence_number": 4, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _model_double(plan: Callable[[Mapping[str, JsonValue]], Turn]) -> Callable[[Request], Reply]: def respond(request: Request) -> Reply: if request.method == "GET" and request.target.endswith("/models"): return _json({"object": "list", "data": []}) - body: Final = json.loads(request.body) - assert isinstance(body, dict), request.body + body: Final = JSON_OBJECT.validate_json(request.body) + if OUTAGE in _prompt(body): + outage: Final = {"error": {"message": "scripted provider outage", "type": "server_error", "code": None}} + return Reply(status=500, body=json.dumps(outage).encode()) + turn: Final = plan(body) done: Final = _has_tool_result(body) + identity: Final = uuid.uuid4().hex[:12] usage: Final = {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} if request.target.endswith("/chat/completions"): message: Final = ( - {"role": "assistant", "content": ANSWER} + {"role": "assistant", "content": turn.answer} if done else { "role": "assistant", "content": None, "tool_calls": [ - {"id": "call_1", "type": "function", "function": {"name": tool, "arguments": arguments}} + { + "id": "call_1", + "type": "function", + "function": {"name": turn.tool, "arguments": turn.arguments}, + } ], } ) @@ -74,7 +165,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: else message ) chunk: Final = { - "id": "chatcmpl-1", + "id": f"chatcmpl-{identity}", "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini", @@ -89,7 +180,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: ) return _json( { - "id": "chatcmpl-1", + "id": f"chatcmpl-{identity}", "object": "chat.completion", "created": 1, "model": "gpt-4o-mini", @@ -99,13 +190,13 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: ) if request.target.endswith("/messages"): content: Final = ( - [{"type": "text", "text": ANSWER}] + [{"type": "text", "text": turn.answer}] if done - else [{"type": "tool_use", "id": "toolu_1", "name": tool, "input": ADD}] + else [{"type": "tool_use", "id": "toolu_1", "name": turn.tool, "input": json.loads(turn.arguments)}] ) return _json( { - "id": "msg_1", + "id": f"msg_{identity}", "type": "message", "role": "assistant", "model": "claude", @@ -116,39 +207,34 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: } ) assert request.target.endswith("/responses"), request.target - output: Final = ( - [ - { - "type": "message", - "id": "msg_1", - "role": "assistant", - "status": "completed", - "content": [{"type": "output_text", "text": ANSWER, "annotations": []}], - } - ] - if done - else [ - { - "type": "function_call", - "id": "fc_1", - "call_id": "call_1", - "name": tool, - "arguments": arguments, - "status": "completed", - } - ] - ) - return _json( + item: Final[dict[str, JsonValue]] = ( { - "id": "resp_1", - "object": "response", - "created_at": 1, + "type": "message", + "id": f"msg_{identity}", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": turn.answer, "annotations": []}], + } + if done + else { + "type": "function_call", + "id": f"fc_{identity}", + "call_id": f"call_{identity}", + "name": turn.tool, + "arguments": turn.arguments, "status": "completed", - "model": "gpt-4o-mini", - "output": output, - "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, } ) + response: Final[dict[str, JsonValue]] = { + "id": f"resp_{identity}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [item], + "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, + } + return _responses_stream(response, item) if body.get("stream") is True else _json(response) return respond @@ -240,7 +326,7 @@ def _rig(gateway: Gateway, surface: Surface) -> Iterator[Rig]: alias: Final = "llm" + uuid.uuid4().hex[:8] with ( mcp_peer() as peer, - wire_server(_model_double(f"{alias}-add")) as wire, + wire_server(_model_double(_fixed_turn(f"{alias}-add"))) as wire, gateway.scenario() as scenario, ): server_id: Final = register_mcp(scenario, peer, alias) @@ -403,3 +489,460 @@ def test_toolset_gateway_url_gives_a_key_of_an_ungranted_team_no_tools_and_never assert _peer_add_calls(rig.peer) == (), "denied caller reached the peer" assert all(rig.tool not in names for names in rig.upstream_tools()), rig.upstream_tools() assert response.status_code in (200, 400, 401, 403), response.text + + +Bridge = Literal["chat", "responses", "messages"] +BRIDGES: Final[tuple[Bridge, ...]] = ("chat", "responses") +Client = Literal["sync", "async"] +HOOK_ECHO: Final = "bridge-echo:" +HOOK_PROBE: Final = "bridge-probe" +HOOK_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + ' texts = list(inputs.get("texts") or [])\n' + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + " for text in texts:\n" + f' if "{HOOK_PROBE}" in text:\n' + f' return block("{HOOK_ECHO}" + json_stringify(' + '{"description": function.get("description"), "parameters": function.get("parameters")}))\n' + " return allow()\n" +) +LOOKUP: Final[tuple[str, dict[str, JsonValue]]] = ( + "Look up one record", + {"type": "object", "properties": {"query": {"type": "string"}}}, +) +REPORT: Final[tuple[str, dict[str, JsonValue]]] = ( + "Write one report", + {"type": "object", "properties": {"query": {"type": "string"}, "format": {}}}, +) +COLD: Final[tuple[str, dict[str, JsonValue]]] = ( + "", + {"type": "object", "properties": {}, "additionalProperties": False}, +) +RELOAD_FAST: Final = {"PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": "5"} +CALL_ID: Final = "x-litellm-call-id" +T = TypeVar("T") +Definition = tuple[str, str, JsonValue] +Echo = tuple[JsonValue, JsonValue] + + +def _served(listing: tuple[str, Mapping[str, JsonValue]]) -> Echo: + return listing[0], {**listing[1], "additionalProperties": False} + + +@dataclass(frozen=True, slots=True) +class Hooked: + proxy: Gateway + sink: Wire + + +@pytest.fixture(scope="module") +def hooked(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Hooked]: + directory: Final = tmp_path_factory.mktemp("bridge-hooks") + base: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + with wire_server(lambda _: _json({"flagged": False, "session_id": "scripted"})) as sink: + echo: Final = {"guardrail": "custom_code", "mode": "pre_mcp_call", "default_on": True, "custom_code": HOOK_CODE} + pillar: Final = { + "guardrail": "pillar", + "mode": ["pre_call", "pre_mcp_call"], + "default_on": True, + "api_key": "sk-pillar-" + uuid.uuid4().hex, + "api_base": sink.url, + "on_flagged_action": "monitor", + } + guardrails: Final = [ + {"guardrail_name": "bridge-echo-" + uuid.uuid4().hex[:8], "litellm_params": echo}, + {"guardrail_name": "bridge-sink-" + uuid.uuid4().hex[:8], "litellm_params": pillar}, + ] + path: Final = directory / "config.yaml" + path.write_text(yaml.safe_dump({**base, "guardrails": guardrails})) + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, directory, RELOAD_FAST, config=path, workers=2) as proxy, + ): + yield Hooked(proxy, sink) + + +@dataclass(frozen=True, slots=True) +class BridgeRig: + hooked: Hooked + scenario: Scenario + peer: McpPeer + wire: Wire + alias: str + server_id: str + model: str + bridge: Bridge + + def tool(self, name: str) -> str: + return f"{self.alias}-{name}" + + def names(self) -> frozenset[str]: + return frozenset(("lookup", "report")) + + def mcp(self, name: str) -> Mcp: + return {**AUTO_MCP, "allowed_tools": [self.tool(name)]} + + def url(self, path: str) -> str: + return str(self.hooked.proxy.client.base_url).rstrip("/") + path + + def post(self, key: str, prompt: str, tools: Sequence[Mcp], **extra: object) -> httpx.Response: + headers: Final = {"Authorization": f"Bearer {key}"} + if self.bridge == "chat": + body: Final = {"model": self.model, "messages": [{"role": "user", "content": prompt}], "tools": list(tools)} + return httpx.post(self.url("/v1/chat/completions"), headers=headers, json={**body, **extra}, timeout=90) + if self.bridge == "responses": + body_r: Final = {"model": self.model, "input": prompt, "tools": list(tools), **extra} + return httpx.post(self.url("/v1/responses"), headers=headers, json=body_r, timeout=90) + body_m: Final = { + "model": self.model, + "max_tokens": 64, + "messages": [{"role": "user", "content": prompt}], + "tools": list(tools), + **extra, + } + return httpx.post(self.url("/v1/messages"), headers=headers, json=body_m, timeout=90) + + def upstream_by_prompt(self) -> Mapping[str, tuple[tuple[Definition, ...], ...]]: + bodies: Final = tuple( + JSON_OBJECT.validate_json(request.body) for request in self.wire.drain() if request.method == "POST" + ) + prompts: Final = frozenset(_prompt(body) for body in bodies) + return {prompt: tuple(_definitions(body) for body in bodies if _prompt(body) == prompt) for prompt in prompts} + + def peer_calls(self) -> tuple[tuple[str, JsonValue], ...]: + return tuple(_peer_call(call) for call in tool_calls(self.peer.drain())) + + def hook_messages(self, marker: str) -> tuple[str, ...]: + posted: Final = tuple( + JSON_OBJECT.validate_json(request.body) for request in self.hooked.sink.drain() if request.method == "POST" + ) + contents: Final = tuple(content for payload in posted for content in _contents(payload)) + return tuple(content for content in contents if _synthetic(content, marker)) + + +AUTO_MCP: Final[Mcp] = { + "type": "mcp", + "server_label": "litellm", + "server_url": "litellm_proxy", + "require_approval": "never", +} + + +def _peer_call(call: Mapping[str, object]) -> tuple[str, JsonValue]: + params: Final = object_value(object_value(JSON_VALUE.validate_python(call["body"]))["params"]) + return str(params["name"]), params["arguments"] + + +def _contents(payload: Mapping[str, JsonValue]) -> Iterator[str]: + messages: Final = payload.get("messages") + for message in messages if isinstance(messages, list) else (): + if isinstance(message, dict): + yield str(message.get("content")) + + +def _function(tool: JsonValue) -> dict[str, JsonValue] | None: + if not isinstance(tool, dict): + return None + function: Final = tool.get("function", tool) + return function if isinstance(function, dict) else None + + +def _definitions(body: Mapping[str, JsonValue]) -> tuple[Definition, ...]: + tools: Final = body.get("tools") + functions: Final = tuple(_function(tool) for tool in tools) if isinstance(tools, list) else () + return tuple( + (str(function["name"]), str(function.get("description", "")), function.get("parameters")) + for function in functions + if function is not None + ) + + +def _uniform(upstream: Mapping[str, tuple[tuple[Definition, ...], ...]]) -> Mapping[str, tuple[Definition, ...] | None]: + return { + prompt: rounds[0] if all(definitions == rounds[0] for definitions in rounds) else None + for prompt, rounds in upstream.items() + } + + +def _strings(value: JsonValue) -> Iterator[str]: + if isinstance(value, str): + yield value + return + children: Final = value.values() if isinstance(value, dict) else value if isinstance(value, list) else () + for child in children: + yield from _strings(child) + + +def _echoed(value: JsonValue) -> Echo: + carrier: Final = next((text for text in _strings(value) if HOOK_ECHO in text), None) + assert carrier is not None, value + payload: Final = carrier.split(HOOK_ECHO, 1)[1] + end: Final = json.JSONDecoder().raw_decode(payload)[1] + echoed: Final = JSON_OBJECT.validate_json(payload[:end]) + return echoed.get("description"), echoed.get("parameters") + + +@contextmanager +def _bridge_rig(hooked: Hooked, bridge: Bridge) -> Generator[BridgeRig, None, None]: + alias: Final = "brg" + uuid.uuid4().hex[:8] + lookup: Final = ScriptedTool( + "lookup", + lambda params: text_result("found:" + json.dumps(JSON_OBJECT.validate_python(params)["arguments"])), + description=LOOKUP[0], + input_schema=LOOKUP[1], + ) + report: Final = ScriptedTool( + "report", lambda _: text_result("reported"), description=REPORT[0], input_schema=REPORT[1] + ) + with ( + scripted_peer(lookup, report) as peer, + wire_server(_model_double(_echoing_turn)) as wire, + hooked.proxy.scenario() as scenario, + ): + server_id: Final = register_mcp(scenario, peer, alias) + model: Final = scenario.model(model=_upstream_model(bridge), api_base=wire.url + "/v1") + rig: Final = BridgeRig(hooked, scenario, peer, wire, alias, server_id, model, bridge) + eventually( + lambda: tuple(_on_worker(hooked.proxy, lambda client: _master_listing(rig, client)) for _ in range(6)), + lambda seen: len({pid for pid, _ in seen}) >= 2 and all(names == rig.names() for _, names in seen), + seconds=45, + ) + peer.drain() + hooked.sink.drain() + yield rig + + +def _bridge_key(rig: BridgeRig) -> str: + return rig.scenario.key(object_permission={"mcp_servers": [rig.server_id]}) + + +@dataclass(frozen=True, slots=True) +class Seen: + response_id: str + call_id: str + text: str + + +def _ask_sync(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen: + sdk: Final = OpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90) + if rig.bridge == "responses": + if stream: + raw_events: Final = sdk.responses.with_raw_response.create( + model=rig.model, input=prompt, tools=tools, stream=True + ) + completed: Final = next( + event.response for event in raw_events.parse() if event.type == "response.completed" + ) + return Seen(completed.id, raw_events.headers[CALL_ID], completed.output_text) + raw_response: Final = sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools) + response: Final = raw_response.parse() + return Seen(response.id, raw_response.headers[CALL_ID], response.output_text) + messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}] + extra: Final = {"tools": list(tools)} + if stream: + raw_chunks: Final = sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, stream=True, extra_body=extra + ) + parts: Final = tuple( + (chunk.id, chunk.choices[0].delta.content or "") for chunk in raw_chunks.parse() if chunk.choices + ) + return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts)) + raw_completion: Final = sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, extra_body=extra + ) + completion: Final = raw_completion.parse() + return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "") + + +async def _ask_async(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen: + sdk: Final = AsyncOpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90) + if rig.bridge == "responses": + if stream: + raw_events: Final = await sdk.responses.with_raw_response.create( + model=rig.model, input=prompt, tools=tools, stream=True + ) + completed: Final = [ + event.response async for event in raw_events.parse() if event.type == "response.completed" + ] + return Seen(completed[0].id, raw_events.headers[CALL_ID], completed[0].output_text) + raw_response: Final = await sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools) + response: Final = raw_response.parse() + return Seen(response.id, raw_response.headers[CALL_ID], response.output_text) + messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}] + extra: Final = {"tools": list(tools)} + if stream: + raw_chunks: Final = await sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, stream=True, extra_body=extra + ) + parts: Final = [ + (chunk.id, chunk.choices[0].delta.content or "") async for chunk in raw_chunks.parse() if chunk.choices + ] + return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts)) + raw_completion: Final = await sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, extra_body=extra + ) + completion: Final = raw_completion.parse() + return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "") + + +def _ask(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool, client: Client) -> Seen: + if client == "async": + return asyncio.run(_ask_async(rig, key, prompt, tools, stream)) + return _ask_sync(rig, key, prompt, tools, stream) + + +def _synthetic(content: str, marker: str) -> bool: + return content.startswith("Tool: lookup\n") and marker in content + + +def _on_worker(gateway: Gateway, act: Callable[[httpx.Client], T]) -> tuple[int, T]: + with httpx.Client(base_url=str(gateway.client.base_url), timeout=30) as client: + summary: Final = client.get("/debug/memory/summary", headers={"Authorization": f"Bearer {gateway.key}"}) + assert summary.status_code == 200, summary.text + pid: Final = JSON_OBJECT.validate_json(summary.content)["worker_pid"] + assert isinstance(pid, int), summary.text + return pid, act(client) + + +def _both_workers(gateway: Gateway) -> frozenset[int]: + return eventually( + lambda: frozenset(_on_worker(gateway, lambda _: None)[0] for _ in range(6)), lambda pids: len(pids) >= 2 + ) + + +def _master_listing(rig: BridgeRig, client: httpx.Client) -> frozenset[str]: + headers: Final = {"x-litellm-api-key": rig.hooked.proxy.key} + response: Final = client.get("/mcp-rest/tools/list", headers=headers, params={"server_id": rig.server_id}) + tools: Final = JSON_OBJECT.validate_json(response.content).get("tools") if response.status_code == 200 else None + return frozenset(str(object_value(tool)["name"]) for tool in tools) if isinstance(tools, list) else frozenset() + + +def _direct_probe(rig: BridgeRig, key: str, name: str, client: httpx.Client) -> Echo: + body: Final = {"server_id": rig.server_id, "name": rig.tool(name), "arguments": {"query": HOOK_PROBE}} + response: Final = client.post("/mcp-rest/tools/call", headers={"x-litellm-api-key": key}, json=body) + return _echoed(JSON_VALUE.validate_json(response.content)) + + +def _spend_row(key: str, call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT api_key, call_type, status, cache_hit FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,) + ), + lambda found: len(found) >= 1, + seconds=70, + ) + assert rows[0]["api_key"] == sha256(key.encode()).hexdigest(), rows + return rows[0] + + +@pytest.mark.parametrize("client", ("sync", "async")) +@pytest.mark.parametrize("stream", (False, True), ids=("plain", "stream")) +@pytest.mark.parametrize("bridge", BRIDGES) +def test_bridge_hook_sees_the_definition_of_the_tool_filtered_for_that_request( + hooked: Hooked, bridge: Bridge, stream: bool, client: Client +) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + marker: Final = "m" + uuid.uuid4().hex + probe: Final = f"{marker} {HOOK_PROBE}" + found: Final = _ask(rig, key, marker, [rig.mcp("lookup")], stream, client) + assert found.text == "found:" + json.dumps({"query": marker}), found + blocked: Final = _ask(rig, key, probe, [rig.mcp("lookup")], stream, client) + assert _echoed(blocked.text) == _served(LOOKUP), blocked + expected: Final = ((rig.tool("lookup"), *_served(LOOKUP)),) + upstream: Final = rig.upstream_by_prompt() + assert set(upstream) == {marker, probe} and all( + definitions == expected for definitions in upstream[marker] + upstream[probe] + ), upstream + assert rig.peer_calls() == (("lookup", {"query": marker}),) + assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",) + assert _spend_row(key, found.call_id)["status"] == "success" + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_direct_call_after_bridge_only_discovery_stays_cold(hooked: Hooked, bridge: Bridge) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + workers: Final = _both_workers(rig.hooked.proxy) + bridged: Final = _ask(rig, key, HOOK_PROBE, [rig.mcp("lookup")], False, "sync") + assert _echoed(bridged.text) == _served(LOOKUP), bridged + direct: Final = eventually( + lambda: tuple( + _on_worker(rig.hooked.proxy, lambda client: _direct_probe(rig, key, "lookup", client)) for _ in range(6) + ), + lambda seen: frozenset(pid for pid, _ in seen) == workers, + ) + assert all(echo == COLD for _, echo in direct) and frozenset(pid for pid, _ in direct) == workers, ( + direct, + workers, + ) + assert rig.peer_calls() == () + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_concurrent_requests_of_one_key_with_different_allowed_tools_each_see_their_own_definition( + hooked: Hooked, bridge: Bridge +) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + prompts: Final = {name: f"{name} {uuid.uuid4().hex} {HOOK_PROBE}" for name in ("lookup", "report")} + + def ask(name: str) -> Seen: + return _ask(rig, key, prompts[name], [rig.mcp(name)], False, "sync") + + with ThreadPoolExecutor(2) as pool: + lookup, report = pool.map(ask, ("lookup", "report")) + assert (_echoed(lookup.text), _echoed(report.text)) == (_served(LOOKUP), _served(REPORT)), (lookup, report) + upstream: Final = rig.upstream_by_prompt() + assert _uniform(upstream) == { + prompts["lookup"]: ((rig.tool("lookup"), *_served(LOOKUP)),), + prompts["report"]: ((rig.tool("report"), *_served(REPORT)),), + }, upstream + assert rig.peer_calls() == () + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_provider_outage_reaches_the_caller_and_never_the_peer_or_the_hooks(hooked: Hooked, bridge: Bridge) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + prompt: Final = f"{OUTAGE} {uuid.uuid4().hex}" + response: Final = rig.post(key, prompt, [rig.mcp("lookup")]) + assert response.status_code == 500, response.text + upstream: Final = rig.upstream_by_prompt() + assert set(upstream) == {prompt} and all( + definitions == ((rig.tool("lookup"), *_served(LOOKUP)),) for definitions in upstream[prompt] + ), upstream + assert rig.peer_calls() == () + assert rig.hook_messages(prompt) == () + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_identical_nonstream_repeat_is_a_cache_hit_without_new_model_peer_or_hook_traffic( + hooked: Hooked, bridge: Bridge +) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + marker: Final = "m" + uuid.uuid4().hex + first: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync") + assert first.text == "found:" + json.dumps({"query": marker}), first + assert set(rig.upstream_by_prompt()) == {marker} and rig.peer_calls() == (("lookup", {"query": marker}),) + assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",) + assert _spend_row(key, first.call_id)["cache_hit"] != "True" + repeat: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync") + assert repeat.text == first.text, (first, repeat) + assert (rig.upstream_by_prompt(), rig.peer_calls(), rig.hook_messages(marker)) == ({}, (), ()), repeat + assert _spend_row(key, repeat.call_id)["cache_hit"] == "True" + + +def test_messages_bridge_hook_keeps_the_base_shape_without_request_local_metadata(hooked: Hooked) -> None: + with _bridge_rig(hooked, "messages") as rig: + key: Final = _bridge_key(rig) + marker: Final = "m" + uuid.uuid4().hex + probe: Final = f"{marker} {HOOK_PROBE}" + found: Final = rig.post(key, marker, [rig.mcp("lookup")]) + assert found.status_code == 200, found.text + blocked: Final = rig.post(key, probe, [rig.mcp("lookup")]) + assert blocked.status_code == 200, blocked.text + assert _echoed(JSON_VALUE.validate_json(blocked.content)) == COLD, blocked.text + assert rig.peer_calls() == (("lookup", {"query": marker}),) + assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",) diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 60563e7aacd..e95cdb902e8 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -1,16 +1,20 @@ import base64 import hashlib +import re import secrets +import textwrap import time import uuid +from collections.abc import Iterator from dataclasses import dataclass +from pathlib import Path from typing import Final from urllib.parse import parse_qs, urlsplit import httpx import jwt import pytest -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, eventually, gateway_from_environment from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, @@ -27,6 +31,7 @@ from integration._support.mcp import ( ) from integration._support.mcp_grants import create_toolset from integration._support.oauth_server import AuthorizationServer, oauth_server +from integration._support.process import owned_proxy ADD: Final = {"a": 2, "b": 3} CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb" @@ -503,3 +508,70 @@ def test_resource_scoped_session_bearer_opens_a_team_toolset_inside_its_server_a refused: Final = _toolset_rpc(gateway, bearer, outside_name, "tools/list", {}) assert refused.status == 403, refused.raw assert tool_calls(peer.drain()) == () + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +def test_token_exchange_callers_with_different_subject_tokens_own_separate_listings(echo_rig: Gateway) -> None: + with mcp_peer() as peer, oauth_server() as auth, echo_rig.scenario() as scenario: + alias: Final = "te" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, + peer, + alias, + auth_type="oauth2_token_exchange", + token_exchange_endpoint=auth.issuer + "/token", + credentials={"client_id": "te-client", "client_secret": "te-secret"}, + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + first_subject: Final = "subject-" + uuid.uuid4().hex + first: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": f"Bearer {first_subject}"}) + second: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": "Bearer subject-" + uuid.uuid4().hex}) + auth.drain() + assert first.list_tools().ok + assert [request["subject_token"] for request in auth.token_requests()] == [first_subject] + probe: Final = {"probe": _PROBE} + own: Final = _echoed_description(first.call(f"{alias}-add", probe)) + other: Final = _echoed_description(second.call(f"{alias}-add", probe)) + assert (own, other) == ("Add two integers", _UNLISTED), ( + "the caller bearer is part of the identity on a token-exchange server: one subject, one slot" + ) + assert second.list_tools().ok + assert _echoed_description(second.call(f"{alias}-add", probe)) == "Add two integers" + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" diff --git a/tests/integration/mcp/test_mcp_resilience.py b/tests/integration/mcp/test_mcp_resilience.py index 8efb54a18fd..d821181962f 100644 --- a/tests/integration/mcp/test_mcp_resilience.py +++ b/tests/integration/mcp/test_mcp_resilience.py @@ -1,13 +1,23 @@ +import itertools import uuid +from collections.abc import Mapping +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path from typing import Final import pytest +import yaml from integration._support.client import Gateway, eventually +from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, EntryPoint, + JsonRpc, McpCaller, Outcome, + ScriptedTool, disconnecting_tool, echo_tool, listed_tools, @@ -15,8 +25,67 @@ from integration._support.mcp import ( register_mcp, scripted_peer, slow_tool, + text_result, tool_calls, ) +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply +from pydantic import BaseModel, TypeAdapter + +_ECHO: Final = "catalog-echo:" +_PROBE: Final = "catalog-probe" +_DESCRIPTION: Final = "Look up one record" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + ' texts = list(inputs.get("texts") or [])\n' + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' if "{_PROBE}" in texts:\n' + f' return block("{_ECHO}" + json_stringify({{"description": function.get("description")}}))\n' + " return allow()\n" +) +_BURST: Final = 20 +_OUTAGE: Final = 6 +_SPEND_NONCES: Final = ( + "SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce" + ' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s' +) +_OBJECTS: Final = TypeAdapter(Mapping[str, object]) +_STRINGS: Final = TypeAdapter(Mapping[str, str]) +_ECHOED: Final = TypeAdapter(Mapping[str, str | None]) + + +class _Content(BaseModel): + text: str + + +class _Result(BaseModel): + content: tuple[_Content, ...] + + +class _RpcReply(BaseModel): + id: int + result: _Result + + +class _SessionsReport(BaseModel): + worker_pid: int + + +@pytest.fixture(scope="module") +def echo_config(tmp_path_factory: pytest.TempPathFactory) -> Path: + base: Final = _OBJECTS.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + guardrail: Final = { + "guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8], + "litellm_params": { + "guardrail": "custom_code", + "mode": "pre_mcp_call", + "default_on": True, + "custom_code": _GUARDRAIL_CODE, + }, + } + path: Final = tmp_path_factory.mktemp("failure-recovery") / "config.yaml" + path.write_text(yaml.safe_dump({**base, "guardrails": [guardrail]})) + return path def _call(caller: McpCaller, name: str, arguments: dict[str, object], entry: EntryPoint, identity: str) -> Outcome: @@ -134,3 +203,151 @@ def test_peer_restart_on_the_same_url_is_picked_up_without_gateway_restart(gatew ) assert back.text == '{"a": 1, "b": 1}', back.raw assert len(tool_calls(replacement.drain())) >= 1 + + +def _rpc_reply(raw: str) -> _RpcReply: + data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:")) + return _RpcReply.model_validate_json(data[-1] if data else raw) + + +def _call_params(call: Mapping[str, object]) -> Mapping[str, object]: + return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"]) + + +def _call_nonce(call: Mapping[str, object]) -> str: + return _STRINGS.validate_python(_call_params(call)["arguments"])["nonce"] + + +def _listed(caller: McpCaller, name: str) -> None: + listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45) + assert listing.error is None, (caller.gateway.client.base_url, listing.raw) + + +@dataclass(frozen=True, slots=True) +class _Worker: + caller: McpCaller + pid: int + + +def _worker(proxy: Gateway, key: str, alias: str) -> _Worker: + sessions: Final = proxy.client.get("/v1/mcp/sessions", headers={"x-litellm-api-key": proxy.key}) + assert sessions.status_code == 200, sessions.text + return _Worker(McpCaller(proxy, key, "mcp", alias), _SessionsReport.model_validate_json(sessions.text).worker_pid) + + +def _served(worker: _Worker, name: str, nonce: str) -> None: + served: Final = worker.caller.call(name, {"nonce": nonce}) + assert served.text == "found", (worker.pid, served.raw) + + +def _probed_description(worker: _Worker, name: str) -> str | None: + """The description the pre_mcp_call guardrail on that worker was handed, recovered from its block reason.""" + blocked: Final = worker.caller.call(name, {"nonce": _PROBE}) + assert blocked.error is not None, (worker.pid, blocked.raw) + carrier: Final = next((item.text for item in _rpc_reply(blocked.raw).result.content if _ECHO in item.text), None) + assert carrier is not None, (worker.pid, blocked.raw) + return _ECHOED.validate_json(carrier.split(_ECHO, 1)[1])["description"] + + +def _catalog_is_cold(worker: _Worker, name: str) -> bool: + return not _probed_description(worker, name) + + +@pytest.mark.timeout(600) +def test_worker_restart_cools_its_listed_catalog_while_the_sibling_worker_keeps_serving( + gateway: Gateway, echo_config: Path, tmp_path: Path +) -> None: + tool: Final = ScriptedTool("lookup", lambda _: text_result("found"), description=_DESCRIPTION) + with scripted_peer(tool) as peer, gateway.scenario() as scenario: + alias: Final = "cold" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = f"{alias}-lookup" + with owned_proxy_process(gateway, tmp_path / "sibling", {}, config=echo_config) as sibling_proxy: + sibling: Final = _worker(sibling_proxy.gateway, key, alias) + with owned_proxy_process(gateway, tmp_path / "first", {}, config=echo_config) as first_proxy: + first: Final = _worker(first_proxy.gateway, key, alias) + assert first.pid != sibling.pid + _served(first, name, "first-unlisted") + _served(sibling, name, "sibling-unlisted") + assert _catalog_is_cold(first, name) and _catalog_is_cold(sibling, name) + _listed(first.caller, name) + assert _probed_description(first, name) == _DESCRIPTION + assert _catalog_is_cold(sibling, name), "a listing on one worker warmed its sibling" + _listed(sibling.caller, name) + assert _probed_description(sibling, name) == _DESCRIPTION + _served(sibling, name, "sibling-alone") + with owned_proxy_process(gateway, tmp_path / "restarted", {}, config=echo_config) as restarted_proxy: + restarted: Final = _worker(restarted_proxy.gateway, key, alias) + assert restarted.pid not in (first.pid, sibling.pid) + assert _catalog_is_cold(restarted, name), "a restarted worker kept the old process's catalog" + assert _probed_description(sibling, name) == _DESCRIPTION + _served(restarted, name, "restarted-unlisted") + _listed(restarted.caller, name) + assert _probed_description(restarted, name) == _DESCRIPTION + calls: Final = tool_calls(peer.drain()) + assert [_call_nonce(call) for call in calls] == [ + "first-unlisted", + "sibling-unlisted", + "sibling-alone", + "restarted-unlisted", + ], calls + assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in calls), calls + + +def _outage_echo(name: str, failures: int) -> ScriptedTool: + attempts: Final = itertools.count(1) + + def respond(params: JsonRpc) -> Reply | JsonRpc: + if next(attempts) <= failures: + return Reply(status=503, body=b'{"error": "scripted outage"}') + return text_result(_STRINGS.validate_python(params["arguments"])["nonce"]) + + return ScriptedTool(name, respond) + + +def _echo_call(caller: McpCaller, name: str, nonce: str) -> Outcome: + return caller.call(name, {"nonce": nonce}) + + +def _logged_nonces(key: str, count: int) -> tuple[tuple[str, str], ...]: + rows: Final = eventually( + lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")), + lambda found: len(found) >= count, + seconds=70, + ) + return tuple(sorted((str(row["status"]), str(row["nonce"])) for row in rows)) + + +def test_peer_outage_during_a_bounded_burst_fails_exactly_the_outage_calls_and_lands_each_call_once( + gateway: Gateway, peer: Gateway +) -> None: + with scripted_peer(_outage_echo("echo", _OUTAGE)) as upstream, gateway.scenario() as scenario: + alias: Final = "burst" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, upstream, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = f"{alias}-echo" + callers: Final = (McpCaller(gateway, key, "mcp", alias), McpCaller(peer, key, "mcp", alias)) + for caller in callers: + _listed(caller, name) + upstream.drain() + nonces: Final = tuple(uuid.uuid4().hex for _ in range(_BURST)) + with ThreadPoolExecutor(max_workers=_BURST) as pool: + outcomes: Final = tuple(pool.map(_echo_call, itertools.cycle(callers), itertools.repeat(name), nonces)) + raws: Final = [outcome.raw for outcome in outcomes] + failed: Final = tuple(nonce for nonce, outcome in zip(nonces, outcomes) if outcome.error is not None) + assert len(failed) == _OUTAGE, raws + assert all(outcome.error is not None or outcome.text == nonce for nonce, outcome in zip(nonces, outcomes)), raws + assert all(_rpc_reply(outcome.raw).id == 1 for outcome in outcomes), raws + burst_calls: Final = tool_calls(upstream.drain()) + assert sorted(_call_nonce(call) for call in burst_calls) == sorted(nonces), burst_calls + assert all(_call_params(call)["arguments"] == {"nonce": _call_nonce(call)} for call in burst_calls), burst_calls + assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in burst_calls), burst_calls + recovered: Final = tuple( + _echo_call(caller, name, nonce) for caller, nonce in zip(itertools.cycle(callers), failed) + ) + assert [outcome.text for outcome in recovered] == list(failed), [outcome.raw for outcome in recovered] + assert sorted(_call_nonce(call) for call in tool_calls(upstream.drain())) == sorted(failed) + assert _logged_nonces(key, _BURST + len(failed)) == tuple( + sorted([("success", nonce) for nonce in nonces] + [("failure", nonce) for nonce in failed]) + ) diff --git a/tests/integration/mcp/test_mcp_toolsets.py b/tests/integration/mcp/test_mcp_toolsets.py index 3dd309db665..b97fa6e72cb 100644 --- a/tests/integration/mcp/test_mcp_toolsets.py +++ b/tests/integration/mcp/test_mcp_toolsets.py @@ -1,5 +1,8 @@ +import re import secrets +import textwrap import uuid +from collections.abc import Iterator from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Final @@ -7,14 +10,17 @@ from typing import Final import httpx import pytest import yaml -from integration._support.client import Gateway, Scenario, object_value +from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value from integration._support.mcp import ( INITIALIZE, Outcome, + ScriptedTool, _outcome_from_rest, _outcome_from_rpc, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, ) from integration._support.mcp_grants import create_toolset @@ -507,3 +513,63 @@ def test_a_member_of_two_teams_sees_the_union_and_each_route_stays_narrowed_to_i crossed: Final = _route_call(gateway, headers, first_name, f"{alias}-multiply") assert not crossed.ok, crossed.raw assert tool_calls(peer.drain()) == () + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +def test_a_team_keys_toolset_route_listing_feeds_its_own_calls_but_not_a_team_mates(echo_rig: Gateway) -> None: + described: Final = "Adds for the team " + uuid.uuid4().hex[:8] + tool: Final = ScriptedTool("add", lambda _: text_result("9"), description=described) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "lit6029echo" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + key: Final = scenario.key(team_id=team_id) + team_mate: Final = scenario.key(team_id=team_id) + _assert_team_grants_only(echo_rig, team_id, key, granted_id) + listed: Final = _toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/list", {}) + assert listed.ok and listed.tools == (f"{alias}-add",), listed.raw + probe: Final[dict[str, object]] = {"name": f"{alias}-add", "arguments": {"probe": _PROBE}} + own: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/call", probe)) + mate: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(team_mate), granted_name, "tools/call", probe)) + assert (own, mate) == (described, _UNLISTED), ( + "the slot is keyed by the hashed key, so a team-mate that never listed is handed nothing" + ) + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py b/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py index 37db61031c9..fc9d77a10a2 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py @@ -226,6 +226,76 @@ async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials( assert execution["guardrail_context"] == {"metadata": {"guardrails": ("block-all",)}} +@pytest.mark.asyncio +async def test_anthropic_messages_with_mcp_hands_execution_the_requests_served_tools(): + """ + Regression test: /v1/messages auto-execution must carry this request's + resolved tool definitions into execution, matching the Responses and chat + completions bridges. + + Given: A request whose MCP reference resolves to a definition carrying a + description and input schema + When: The model asks for that tool and the gateway executes it + Then: _execute_tool_calls receives the definition under served_tools, so + pre_mcp_call hooks can judge the call on what the model was shown + + Dropping it does not fail loudly; the call still runs, but the hook sees + only the name and arguments, leaving /v1/messages permanently colder than + the other two bridges even though all three resolve the same definitions. + """ + from mcp.types import Tool + + from litellm.llms.anthropic.pass_through.messages import mcp_handler + from litellm.responses.mcp.request_context import MCPRequestContext + + served = [ + Tool( + name="read_wiki_structure", + description="Read the structure of a wiki", + inputSchema={"type": "object", "properties": {"repoName": {"type": "string"}}}, + ) + ] + + process = AsyncMock(return_value=(served, {"read_wiki_structure": "deepwiki"})) + execute = AsyncMock( + return_value=[{"tool_call_id": "toolu_1", "result": "ok", "name": "read_wiki_structure"}] + ) + responses = [ + { + "stop_reason": "tool_use", + "content": [{"type": "tool_use", "id": "toolu_1", "name": "read_wiki_structure", "input": {}}], + }, + {"stop_reason": "end_turn", "content": [{"type": "text", "text": "done"}]}, + ] + + with ( + patch.object(MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")), + patch.object( + import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, + "_process_mcp_tools_without_openai_transform", + new=process, + ), + patch.object( + import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, + "_execute_tool_calls", + new=execute, + ), + patch("litellm.anthropic_messages", new=AsyncMock(side_effect=responses)), + ): + await mcp_handler.anthropic_messages_with_mcp( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="claude-sonnet-4-5", + tools=[MCP_REFERENCE], + ) + + execution = execute.call_args.kwargs + assert execution.get("served_tools") == served, ( + "The request's resolved tool definitions must reach execution so pre_mcp_call " + "hooks see the listed description and input schema" + ) + + @pytest.mark.asyncio async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped(): """ diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py index e1e4cd3d161..8e87837611a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py @@ -48,6 +48,7 @@ def _bare_manager() -> MOD.MCPServerManager: reaches the guardrail hooks; they have their own coverage elsewhere. """ mgr = MOD.MCPServerManager.__new__(MOD.MCPServerManager) + mgr._listed_tools_by_server_id = {} mgr.check_allowed_or_banned_tools = lambda name, server: True mgr.validate_allowed_params = lambda tool_name, arguments, server: None diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py index ed5d67164bd..e52a86d76af 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py @@ -1,17 +1,29 @@ from litellm.proxy._experimental.mcp_server import operations as mcp_operations +import asyncio import json from datetime import datetime +from unittest.mock import AsyncMock, patch import pytest from fastapi import HTTPException from mcp.shared.exceptions import MCPError +from mcp.types import CallToolResult, TextContent +from mcp.types import Tool as MCPTool from pydantic import AnyUrl import litellm from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._experimental.mcp_server import server from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller +from litellm.proxy._experimental.mcp_server.tool_search import ( + handle_mcp_proxy_tool, + mcp_proxy_tool_id, + with_mcp_proxy_identity, +) from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth +from litellm.types.mcp import MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer AUTH = UserAPIKeyAuth(api_key="key") @@ -130,3 +142,46 @@ async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.Monke assert hook_payload["arguments"] == arguments assert "raw_headers" not in hook_payload assert "raw-scope-secret" not in recorder.events[1][1] + + +@pytest.mark.asyncio +async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_listed_tool() -> None: + """/mcp/proxy tools/list serves only the meta-tools, so the catalog call_tool reads to resolve its + tool_id was never served: it must not fill the caller's listed-tools slot, and the pre-call hook + must see no listed tool for the call.""" + manager = mcp_operations.global_mcp_server_manager + server = MCPServer(server_id="proxy-meta", name="proxy-meta", transport=MCPTransport.http, url="http://meta") + auth = UserAPIKeyAuth(api_key="sk-proxy-meta", user_id="proxy-caller") + upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})] + served_as = with_mcp_proxy_identity(MCPTool(name="proxy-meta-echo", inputSchema={}), server.server_id) + pre_call_tool_check = AsyncMock(return_value={}) + + async def call_regular_mcp_tool(*, tasks: list[asyncio.Task[object]], **_: object) -> CallToolResult: + await asyncio.gather(*tasks) + return CallToolResult(content=[TextContent(type="text", text="echoed")]) + + with ( + patch.dict(manager.registry, {server.server_id: server}), + patch.dict(manager.tool_name_to_mcp_server_name_mapping), + patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), + patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), + patch.object(manager, "pre_call_tool_check", pre_call_tool_check), + patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool), + ): + try: + result = await handle_mcp_proxy_tool( + name="call_tool", + arguments={"tool_id": mcp_proxy_tool_id(served_as), "arguments": {}}, + user_api_key_dict=auth, + ) + listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=auth)) + finally: + manager._drop_listed_tools(server.server_id) + + assert result.is_error is False + assert result.content[0].text == "echoed" + pre_call_tool_check.assert_awaited_once() + assert pre_call_tool_check.await_args.kwargs["name"] == "echo" + assert pre_call_tool_check.await_args.kwargs["tool"] is None + assert listed is None diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index d3679506a2f..316988ef175 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -963,6 +963,8 @@ async def test_get_tools_from_mcp_servers(): user_api_key_auth=None, oauth2_headers=None, proxy_logging_obj=None, + catalog_auth_header=None, + record_listing=True, ): if server.server_id == "server1_id": return [mock_tool_1] @@ -1998,6 +2000,7 @@ async def test_get_tools_for_single_server(): client_ip=None, user_api_key_auth=None, proxy_logging_obj=ANY, + record_listing=False, ) # Verify the result diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cf6cf93f35b..cb017afbea5 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6,6 +6,7 @@ import json import logging import os import sys +import time from collections.abc import AsyncIterator from datetime import datetime from pathlib import Path @@ -41,8 +42,10 @@ from mcp.types import Tool as MCPTool from pydantic import AnyUrl, TypeAdapter from litellm.constants import MCP_METADATA_TIMEOUT +from litellm.proxy._experimental.mcp_server import discoverable_endpoints from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + ListedToolsCaller, MCPServerManager, _deserialize_json_dict, _flow_endpoints_missing, @@ -53,6 +56,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _obo_retry_applies, _resolve_openapi_tool_auth, _should_strip_caller_authorization, + listed_tools_caller_for, ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -116,7 +120,6 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte assert sampling.await_args.kwargs["raw_headers"] == {"x-test-caller": "sampling-caller"} - @pytest.mark.asyncio async def test_sampling_callback_keeps_creation_context_after_caller_switch(): from mcp.server.auth.middleware.auth_context import auth_context_var @@ -218,8 +221,6 @@ def _reload_mcp_manager_module(): return reloaded - - @pytest.fixture(autouse=True) def enable_eager_mcp_oauth_discovery(monkeypatch): monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1") @@ -1576,7 +1577,9 @@ class TestMCPServerManager: assert not any("oauth2_id_jag" in message for message in caplog.messages) @pytest.mark.asyncio - async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, config_only_mcp_manager_factory, monkeypatch, caplog): + async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso( + self, config_only_mcp_manager_factory, monkeypatch, caplog + ): self._clear_sso_env(monkeypatch) monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid") manager = config_only_mcp_manager_factory() @@ -4862,7 +4865,9 @@ class TestMCPServerManager: @pytest.mark.parametrize("auth_type", [MCPAuth.none, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.oauth2]) @pytest.mark.parametrize("is_byok", [False, True]) @pytest.mark.parametrize("scheme", ["http", "https"]) - async def test_openapi_health_loads_spec_without_mcp_handshake(self, respx_mock, monkeypatch, auth_type, is_byok, scheme): + async def test_openapi_health_loads_spec_without_mcp_handshake( + self, respx_mock, monkeypatch, auth_type, is_byok, scheme + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4912,14 +4917,28 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("failure", "expected_status", "expected_error"), [ - (httpx.Response(401, text="secret response content"), "unhealthy", "OpenAPI specification request failed (HTTP 401)"), + ( + httpx.Response(401, text="secret response content"), + "unhealthy", + "OpenAPI specification request failed (HTTP 401)", + ), (httpx.Response(404), "unhealthy", "OpenAPI specification request failed (HTTP 404)"), (httpx.Response(500), "unhealthy", "OpenAPI specification request failed (HTTP 500)"), - (httpx.ConnectError("secret network details"), "unhealthy", "OpenAPI specification could not be loaded (ConnectError)"), - (httpx.Response(200, text="secret invalid JSON body"), "unhealthy", "OpenAPI specification could not be loaded (JSONDecodeError)"), + ( + httpx.ConnectError("secret network details"), + "unhealthy", + "OpenAPI specification could not be loaded (ConnectError)", + ), + ( + httpx.Response(200, text="secret invalid JSON body"), + "unhealthy", + "OpenAPI specification could not be loaded (JSONDecodeError)", + ), ], ) - async def test_openapi_health_reports_safe_failures(self, respx_mock, monkeypatch, failure, expected_status, expected_error): + async def test_openapi_health_reports_safe_failures( + self, respx_mock, monkeypatch, failure, expected_status, expected_error + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -5077,7 +5096,10 @@ class TestMCPServerManager: @pytest.mark.asyncio @pytest.mark.parametrize("oauth2_flow", [None, "authorization_code", "client_credentials"]) async def test_health_check_server_oauth2_reports_reachability( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, oauth2_flow: Literal["authorization_code", "client_credentials"] | None + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + oauth2_flow: Literal["authorization_code", "client_credentials"] | None, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() @@ -5106,14 +5128,28 @@ class TestMCPServerManager: assert not {"authorization", "x-api-key", "cookie"}.intersection(route.calls[0].request.headers) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type", [ - MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token, - MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, MCPAuth.true_passthrough, MCPAuth.oauth_delegate, - ]) + @pytest.mark.parametrize( + "auth_type", + [ + MCPAuth.bearer_token, + MCPAuth.api_key, + MCPAuth.basic, + MCPAuth.authorization, + MCPAuth.token, + MCPAuth.oauth2_token_exchange, + MCPAuth.oauth2_id_jag, + MCPAuth.true_passthrough, + MCPAuth.oauth_delegate, + ], + ) @pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse]) @pytest.mark.parametrize("response_code", [200, 204, 302, 401, 403, 405, 503]) async def test_health_check_without_credentials_accepts_any_http_response( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, auth_type: MCPAuthType, transport: Literal[MCPTransport.http, MCPTransport.sse], + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + auth_type: MCPAuthType, + transport: Literal[MCPTransport.http, MCPTransport.sse], response_code: int, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") @@ -5144,6 +5180,7 @@ class TestMCPServerManager: self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, response_code: int ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + class UnreadBody(httpx.AsyncByteStream): def __init__(self) -> None: self.read = False @@ -5158,17 +5195,28 @@ class TestMCPServerManager: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="streaming-health", name="streaming-health", transport=MCPTransport.sse, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/events", + server_id="streaming-health", + name="streaming-health", + transport=MCPTransport.sse, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/events", ) manager.registry[server.server_id] = server bodies: Final = (UnreadBody(), UnreadBody()) - route: Final = respx_mock.get(server.url).mock(side_effect=[ - httpx.Response(response_code, stream=body, headers={ - "Content-Type": "text/event-stream", "Set-Cookie": "health=secret; Path=/", - "Location": "http://127.0.0.1/private", - }) for body in bodies - ]) + route: Final = respx_mock.get(server.url).mock( + side_effect=[ + httpx.Response( + response_code, + stream=body, + headers={ + "Content-Type": "text/event-stream", + "Set-Cookie": "health=secret; Path=/", + "Location": "http://127.0.0.1/private", + }, + ) + for body in bodies + ] + ) first: Final = await manager.health_check_server(server.server_id) second: Final = await manager.health_check_server(server.server_id) @@ -5179,19 +5227,28 @@ class TestMCPServerManager: assert all("cookie" not in call.request.headers for call in route.calls) @pytest.mark.asyncio - @pytest.mark.parametrize(("transport", "url"), [ - (MCPTransport.stdio, "https://mcp.example.test"), - (MCPTransport.http, None), (MCPTransport.http, ""), (MCPTransport.http, "not-a-url"), - (MCPTransport.http, "ftp://mcp.example.test"), - (MCPTransport.http, "https://user:secret@mcp.example.test"), - (MCPTransport.http, "https://mcp.example.test:bad/mcp"), - ]) + @pytest.mark.parametrize( + ("transport", "url"), + [ + (MCPTransport.stdio, "https://mcp.example.test"), + (MCPTransport.http, None), + (MCPTransport.http, ""), + (MCPTransport.http, "not-a-url"), + (MCPTransport.http, "ftp://mcp.example.test"), + (MCPTransport.http, "https://user:secret@mcp.example.test"), + (MCPTransport.http, "https://mcp.example.test:bad/mcp"), + ], + ) async def test_health_reachability_rejects_unprobeable_urls_without_requests( self, respx_mock: MockRouter, transport: Literal[MCPTransport.http, MCPTransport.stdio], url: str | None ) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="unprobeable", name="unprobeable", transport=transport, auth_type=MCPAuth.oauth2, url=url, + server_id="unprobeable", + name="unprobeable", + transport=transport, + auth_type=MCPAuth.oauth2, + url=url, ) manager.registry[server.server_id] = server @@ -5202,19 +5259,26 @@ class TestMCPServerManager: assert not respx_mock.calls @pytest.mark.asyncio - @pytest.mark.parametrize("failure", [ - httpx.ConnectError("TLS/connection failure with secret details"), - httpx.ReadTimeout("secret timeout details"), - httpx.RemoteProtocolError("secret malformed response"), - ]) + @pytest.mark.parametrize( + "failure", + [ + httpx.ConnectError("TLS/connection failure with secret details"), + httpx.ReadTimeout("secret timeout details"), + httpx.RemoteProtocolError("secret malformed response"), + ], + ) async def test_health_reachability_reports_no_response_without_secret_details( self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, failure: httpx.RequestError ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="failed-health", name="failed-health", transport=MCPTransport.http, - auth_type=MCPAuth.bearer_token, is_byok=True, url="https://mcp.example.test/secret?token=secret", + server_id="failed-health", + name="failed-health", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + is_byok=True, + url="https://mcp.example.test/secret?token=secret", ) manager.registry[server.server_id] = server route: Final = respx_mock.get(server.url).mock(side_effect=failure) @@ -5230,8 +5294,11 @@ class TestMCPServerManager: monkeypatch.setenv("SSL_SECURITY_LEVEL", "invalid-secret-cipher") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="bad-tls", name="bad-tls", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test", + server_id="bad-tls", + name="bad-tls", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test", ) manager.registry[server.server_id] = server @@ -5249,8 +5316,11 @@ class TestMCPServerManager: monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_HEALTH_CHECK_TIMEOUT", 0.1) manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="slow-health", name="slow-health", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/slow", + server_id="slow-health", + name="slow-health", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/slow", ) manager.registry[server.server_id] = server started: Final = asyncio.Event() @@ -5303,8 +5373,11 @@ class TestMCPServerManager: server_ids: Final = [f"health-{index}" for index in range(server_count)] manager.registry = { server_id: MCPServer( - server_id=server_id, name=server_id, transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url=f"https://health.example.test/{server_id}", + server_id=server_id, + name=server_id, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url=f"https://health.example.test/{server_id}", ) for server_id in server_ids } @@ -5631,8 +5704,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers captured["server_label"] = server_label @@ -5717,8 +5797,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers @@ -7067,8 +7154,7 @@ class TestMCPServerManager: # Mock _create_mcp_client to return our mock client manager._create_mcp_client = AsyncMock(return_value=mock_client) - # Mock user auth with no restrictions - user_api_key_auth: Final = UserAPIKeyAuth() + user_api_key_auth: Final = UserAPIKeyAuth(api_key="sk-test") # Mock proxy logging proxy_logging_obj = MagicMock() @@ -7094,6 +7180,1051 @@ class TestMCPServerManager: # Verify the MCP client call was awaited exactly once assert mock_client.call_tool.await_count == 1 + @staticmethod + def _manager_ready_for_call_tool( + listed_tools: list[MCPTool], caller: ListedToolsCaller | None = None + ) -> tuple[MCPServerManager, MagicMock]: + from mcp.types import CallToolResult + + manager = MCPServerManager() + server = MCPServer( + server_id="test-server", + name="test-server", + transport=MCPTransport.http, + url="http://test-server.com", + ) + manager.registry = {"test-server": server} + manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" + manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" + manager._create_prefixed_tools(listed_tools, server) + manager._record_listed_tools(server, listed_tools, caller) + + mock_client = AsyncMock() + mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) + manager._create_mcp_client = AsyncMock(return_value=mock_client) + + proxy_logging_obj = MagicMock() + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + return manager, proxy_logging_obj + + @staticmethod + def _unrestricted_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test") + + @pytest.mark.asyncio + async def test_call_tool_hands_listed_tool_description_and_schema_to_pre_call_hooks(self): + schema = {"type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"]} + listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] + auth = self._unrestricted_auth() + manager, proxy_logging_obj = self._manager_ready_for_call_tool( + listed, caller=ListedToolsCaller(user_api_key_auth=auth) + ) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=auth, + proxy_logging_obj=proxy_logging_obj, + ) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ("Runs the test tool", schema) + + @pytest.mark.asyncio + async def test_call_tool_hands_during_call_hooks_name_and_arguments_only_even_for_a_listed_tool(self): + """A during_mcp_call guardrail evaluates the call in flight, so it keeps seeing only the name and + arguments it always did; the listed description and schema go to the pre-call hooks alone.""" + schema = {"type": "object", "properties": {"param": {"type": "string"}}} + listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] + auth = UserAPIKeyAuth(api_key="sk-test") + manager, _ = self._manager_ready_for_call_tool(listed, caller=ListedToolsCaller(user_api_key_auth=auth)) + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=auth, + proxy_logging_obj=proxy_logging_obj, + ) + + during_data = proxy_logging_obj.during_call_hook.call_args.kwargs["data"] + assert during_data["mcp_arguments"] == {"param": "value"} + assert (during_data.get("mcp_tool_description"), during_data.get("mcp_input_schema")) == (None, None) + assert "Description:" not in during_data["messages"][0]["content"] + + @pytest.mark.asyncio + async def test_call_tool_passes_no_tool_metadata_when_tool_was_never_listed(self): + auth = self._unrestricted_auth() + manager, proxy_logging_obj = self._manager_ready_for_call_tool( + [MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})], + caller=ListedToolsCaller(user_api_key_auth=auth), + ) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=auth, + proxy_logging_obj=proxy_logging_obj, + ) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) + + def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) + manager._record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) + + latest = manager.get_listed_tool(server, "echo") + assert latest is not None and latest.description == "v2" + assert manager.get_listed_tool(server, "missing") is None + + def test_get_listed_tool_never_strips_the_bare_name_it_is_given(self): + """The lookup is exact: a never-listed tool whose bare name starts with the server prefix is not the + listed sibling that stripping the prefix again would name.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv-id", name="srv", alias="srv", transport=MCPTransport.http, url="http://srv") + caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice")) + manager._record_listed_tools( + server, + [ + MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}), + MCPTool(name="bar", description="Fetches bar records", inputSchema={"type": "object"}), + ], + caller, + ) + + assert manager.get_listed_tool(server, "srv-foo", caller) is None + listed = manager.get_listed_tool(server, "foo", caller) + assert listed is not None and listed.description == "Fetches foo records" + + @pytest.mark.asyncio + async def test_get_listed_tool_uses_admin_description_override_clients_saw(self): + schema = {"type": "object", "properties": {"text": {"type": "string"}}} + manager = _catalog_manager( + MCPTool(name="echo", description="Upstream wording", inputSchema=schema), + MCPTool(name="ping", description="Untouched", inputSchema={}), + ) + server = MCPServer( + server_id="srv", + name="srv", + transport=MCPTransport.http, + url="http://srv", + tool_name_to_description={"echo": "Admin wording"}, + ) + await manager._get_tools_from_server(server, add_prefix=True, record_listing=True) + + overridden = manager.get_listed_tool(server, "echo") + assert overridden is not None + assert (overridden.name, overridden.description, overridden.input_schema) == ("echo", "Admin wording", schema) + untouched = manager.get_listed_tool(server, "ping") + assert untouched is not None and untouched.description == "Untouched" + + @pytest.mark.asyncio + async def test_get_listed_tool_keeps_the_masked_description_over_the_admin_override(self, catalog_guardrail): + """A discovery guardrail masked the admin override in tools/list, so the tool-call hooks must see + the masked wording, not the original override the caller never saw.""" + _, proxy_logging_obj = catalog_guardrail + manager = _catalog_manager(MCPTool(name="read_note", description="Read a note", inputSchema={"type": "object"})) + server = MCPServer( + server_id="notes", + name="notes", + transport=MCPTransport.http, + tool_name_to_description={"read_note": "Read a SECRET note"}, + ) + served = await manager._get_tools_from_server( + server, add_prefix=True, proxy_logging_obj=proxy_logging_obj, record_listing=True + ) + assert [tool.description for tool in served] == ["Read a [MASKED] note"] + + listed = manager.get_listed_tool(server, "read_note") + assert listed is not None and listed.description == "Read a [MASKED] note", ( + "tools/call must be evaluated against the description tools/list served" + ) + + def test_server_definition_change_drops_listed_tools(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other") + manager._record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) + manager._record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) + + manager._invalidate_server_definition_caches(server.server_id) + + assert manager.get_listed_tool(server, "echo") is None + kept = manager.get_listed_tool(other, "ping") + assert kept is not None and kept.description == "kept" + + @pytest.mark.asyncio + async def test_server_save_during_an_in_flight_listing_is_not_undone_by_the_stale_record(self): + """A PUT /v1/mcp/server that lands while a listing awaits its upstream fetch drops the server's + catalog; the fetch completing afterwards must not write the pre-save catalog back, or hooks see + the old description next to the new definition until the next listing.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="saver") + fetch_started = asyncio.Event() + release_fetch = asyncio.Event() + + async def fetch(client, name): + fetch_started.set() + await release_fetch.wait() + return [MCPTool(name="turn", description="before save", inputSchema={})] + + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = fetch + caller = ListedToolsCaller(user_api_key_auth=user) + + async def list_tools() -> None: + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + + listing = asyncio.create_task(list_tools()) + await fetch_started.wait() + manager._invalidate_server_definition_caches(server.server_id) + release_fetch.set() + await listing + + assert manager.get_listed_tool(server, "turn", caller) is None + + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="after save", inputSchema={})] + ) + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + listed = manager.get_listed_tool(server, "turn", caller) + assert listed is not None and listed.description == "after save" + + @pytest.mark.asyncio + async def test_update_server_refreshing_openapi_tools_drops_a_listing_recorded_during_the_spec_fetch(self): + """An OpenAPI server's registry entries are rebuilt after the save is published, so a listing that + records while the spec is fetched holds the pre-save entries; the catalog is dropped again once the + registry is current.""" + manager = MCPServerManager() + old = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json" + ) + manager.registry[old.server_id] = old + new = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://new", spec_path="/new.json" + ) + caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) + + async def register_while_a_listing_records(server: MCPServer, *, initialize_mapping: bool = True) -> None: + manager._record_listed_tools( + server, + [MCPTool(name="search", description="pre-save", inputSchema={})], + caller, + manager._listed_tools_generations.get(server.server_id, 0), + ) + + manager.build_mcp_server_from_table = AsyncMock(return_value=new) + manager._maybe_register_openapi_tools = register_while_a_listing_records + manager.prime_oauth_metadata_discovery = MagicMock() + record = LiteLLM_MCPServerTable( + server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http + ) + + await manager.update_server(record) + + assert manager.registry["srv"] is new + assert manager.get_listed_tool(new, "search", caller) is None + + @pytest.mark.asyncio + @pytest.mark.parametrize("already_registered", [False, True], ids=["add_server", "update_server"]) + async def test_openapi_spec_re_read_keeps_discovery_and_oauth_metadata_filled_during_the_fetch( + self, already_registered: bool + ): + """The listed-tool catalog recorded during the spec fetch holds pre-save entries, but a prompts + discovery or OAuth protected-resource fetch answered in that window already saw the published + definition; dropping those too sends the next request upstream again.""" + manager = MCPServerManager() + if already_registered: + manager.registry["srv"] = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json" + ) + new = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://new", spec_path="/new.json" + ) + caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) + metadata_key: Final = (new.server_id, new.url) + prompt_fetches = 0 + + async def fetch_prompts() -> list[Prompt]: + nonlocal prompt_fetches + prompt_fetches += 1 + return [Prompt(name="greet")] + + async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None: + manager._record_listed_tools( + server, + [MCPTool(name="search", description="pre-save", inputSchema={})], + caller, + manager._listed_tools_generations.get(server.server_id, 0), + ) + await manager._prompt_discovery_cache.get((server.server_id, None), fetch_prompts) + discoverable_endpoints._OAUTH_METADATA_CACHE[metadata_key] = (time.time() + 300, {"resource": new.url}) + + manager.build_mcp_server_from_table = AsyncMock(return_value=new) + manager._maybe_register_openapi_tools = register_while_discovery_fills + manager.prime_oauth_metadata_discovery = MagicMock() + record = LiteLLM_MCPServerTable( + server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http + ) + save = manager.update_server if already_registered else manager.add_server + + try: + await save(record) + + 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_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} + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(metadata_key, None) + + @pytest.mark.asyncio + async def test_user_oauth_refresh_keeps_listed_tools(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) + + await manager.invalidate_user_oauth_token_cache("alice", server.server_id) + + listed = manager.get_listed_tool(server, "echo") + assert listed is not None and listed.description == "shared" + + def test_per_caller_server_keeps_listed_tools_per_identity(self): + manager = MCPServerManager() + server = MCPServer( + server_id="srv", + name="srv", + transport=MCPTransport.http, + url="http://srv", + auth_type=MCPAuth.oauth2_token_exchange, + ) + alice = UserAPIKeyAuth(user_id="alice", token="hashed-alice") + bob = UserAPIKeyAuth(user_id="bob", token="hashed-bob") + alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}} + bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}} + manager._record_listed_tools( + server, + [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], + ListedToolsCaller(user_api_key_auth=alice), + ) + manager._record_listed_tools( + server, + [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], + ListedToolsCaller(user_api_key_auth=bob), + ) + + alice_tool = manager.get_listed_tool(server, "read", ListedToolsCaller(user_api_key_auth=alice)) + bob_tool = manager.get_listed_tool(server, "read", ListedToolsCaller(user_api_key_auth=bob)) + assert alice_tool is not None and (alice_tool.description, alice_tool.input_schema) == ( + "alice view", + alice_schema, + ) + assert bob_tool is not None and (bob_tool.description, bob_tool.input_schema) == ("bob view", bob_schema) + carol = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="carol", token="k")) + assert manager.get_listed_tool(server, "read", carol) is None + + shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") + manager._record_listed_tools( + shared, + [MCPTool(name="echo", description="everyone", inputSchema={})], + ListedToolsCaller(user_api_key_auth=alice), + ) + for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob)) + assert for_bob is None, "keyed callers get their own slot even on servers without upstream per-user auth" + anonymous = manager.get_listed_tool(shared, "echo") + assert anonymous is None + + @pytest.mark.parametrize( + ("server_kwargs", "caller_a", "caller_b"), + [ + pytest.param( + {"extra_headers": ["X-Workspace"]}, + ListedToolsCaller(raw_headers={"x-workspace": "A"}), + ListedToolsCaller(raw_headers={"X-Workspace": "B"}), + id="forwarded-header", + ), + pytest.param( + {"auth_type": MCPAuth.true_passthrough}, + ListedToolsCaller(raw_headers={"authorization": "Bearer upstream-a"}), + ListedToolsCaller(raw_headers={"authorization": "Bearer upstream-b"}), + id="anonymous-passthrough-bearer", + ), + pytest.param( + {"auth_type": MCPAuth.bearer_token}, + ListedToolsCaller(mcp_auth_header="byok-a"), + ListedToolsCaller(mcp_auth_header="byok-b"), + id="per-server-auth-header", + ), + pytest.param( + {"auth_type": MCPAuth.oauth2_token_exchange}, + ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", token="hashed-shared"), + raw_headers={"x-litellm-api-key": "sk-shared", "authorization": "Bearer entra-alice"}, + ), + ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", token="hashed-shared"), + raw_headers={"x-litellm-api-key": "sk-shared", "authorization": "Bearer entra-bob"}, + ), + id="shared-key-different-obo-subjects", + ), + pytest.param( + {"transport": MCPTransport.stdio, "command": "srv", "env": {"WS": "${X-WS}"}}, + ListedToolsCaller(raw_headers={"X-WS": "A"}), + ListedToolsCaller(raw_headers={"X-WS": "B"}), + id="header-driven-stdio-env", + ), + ], + ) + def test_upstream_identity_inputs_keep_listed_tools_apart(self, server_kwargs, caller_a, caller_b): + manager = MCPServerManager() + server = MCPServer( + **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} + ) + manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) + manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) + + for_a = manager.get_listed_tool(server, "turn", caller_a) + for_b = manager.get_listed_tool(server, "turn", caller_b) + assert for_a is not None and for_a.description == "Catalog A" + assert for_b is not None and for_b.description == "Catalog B" + assert manager.get_listed_tool(server, "turn", ListedToolsCaller()) is None + + def test_shared_server_ignores_headers_it_never_forwards(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._record_listed_tools( + server, + [MCPTool(name="turn", description="everyone", inputSchema={})], + ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), + ) + + other = ListedToolsCaller(raw_headers={"authorization": "Bearer sk-other", "x-workspace": "B"}) + listed = manager.get_listed_tool(server, "turn", other) + assert listed is not None and listed.description == "everyone" + + @pytest.mark.asyncio + async def test_byok_listing_never_reads_the_credential_store(self): + """tools/list keys the caller's catalog slot by what the client supplied plus the caller's key. + Resolving the stored BYOK credential for that would fail every REST listing while the DB is + down and would seed a per-worker cache the next tools/call trusts over the store.""" + manager = MCPServerManager() + server = MCPServer( + server_id="byok-cold", + name="byok_cold", + transport=MCPTransport.http, + url="http://byok-cold", + is_byok=True, + auth_type=MCPAuth.api_key, + ) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-cold-user") + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="listed while db down", inputSchema={})] + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch( + "litellm.proxy._experimental.mcp_server.db.get_user_credential", + AsyncMock(side_effect=RuntimeError("DB DOWN")), + ), + ): + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + + listed = manager.get_listed_tool(server, "turn", listed_tools_caller_for(server, user, None, None, None, None)) + assert listed is not None and listed.description == "listed while db down" + + @pytest.mark.parametrize( + ("list_header", "call_kwargs"), + [ + pytest.param( + None, {"mcp_auth_header": "stored-secret", "catalog_auth_header": None}, id="execute-mcp-tool" + ), + pytest.param(None, {"mcp_auth_header": None}, id="responses-api"), + pytest.param("Bearer hdr", {"mcp_auth_header": "Bearer hdr"}, id="client-supplied-header"), + ], + ) + @pytest.mark.asyncio + async def test_byok_tools_call_reads_the_slot_the_clients_own_header_listed( + self, list_header: str | None, call_kwargs: dict[str, str | None] + ): + """A REST listing records under the header the client sent (none here). tools/call then swaps the + stored credential in, either before reaching ``call_tool`` (``execute_mcp_tool``) or inside it (the + Responses API), and must still read that slot rather than one keyed by the credential.""" + from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache_key, + cache_byok_credential, + ) + from litellm.proxy._experimental.mcp_server.operations import byok_credential_cache + + manager = MCPServerManager() + server = MCPServer( + server_id="byok-catalog", + name="byok_catalog", + transport=MCPTransport.http, + url="http://byok-catalog", + is_byok=True, + ) + manager.registry = {"byok-catalog": server} + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-user") + mock_client = AsyncMock() + mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) + manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="stored cred catalog", inputSchema={})] + ) + proxy_logging_obj = MagicMock() + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + cache_byok_credential("byok-user", "byok-catalog", "stored-secret") + try: + await manager._get_tools_from_server( + server=server, mcp_auth_header=list_header, user_api_key_auth=user, record_listing=True + ) + listed = manager.get_listed_tool( + server, "turn", listed_tools_caller_for(server, user, list_header, None, None, None) + ) + assert listed is not None and listed.description == "stored cred catalog" + + await manager.call_tool( + server_name="byok_catalog", + name="turn", + arguments={}, + user_api_key_auth=user, + proxy_logging_obj=proxy_logging_obj, + **call_kwargs, + ) + finally: + byok_credential_cache.delete_cache(byok_credential_cache_key("byok-user", "byok-catalog")) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert hook_kwargs["tool_description"] == "stored cred catalog" + + @pytest.mark.asyncio + async def test_byok_supplied_header_lists_without_credential_validation(self): + manager = MCPServerManager() + server = MCPServer( + server_id="byok-catalog", + name="byok_catalog", + transport=MCPTransport.http, + url="http://byok-catalog", + is_byok=True, + ) + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="t", inputSchema={})] + ) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + await manager._get_tools_from_server( + server=server, + mcp_auth_header="Bearer hdr", + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), + record_listing=True, + ) + + caller: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), mcp_auth_header="Bearer hdr" + ) + listed = manager.get_listed_tool(server, "turn", caller) + assert listed is not None and listed.description == "t" + assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" + + @pytest.mark.parametrize( + "server_auth", + [ + pytest.param( + { + "auth_type": MCPAuth.oauth2, + "client_id": "cid", + "client_secret": "csec", + "token_url": "http://cc1/token", + }, + id="oauth2", + ), + pytest.param({"auth_type": MCPAuth.api_key, "authentication_token": "STATIC-ADMIN-TOKEN"}, id="api_key"), + pytest.param( + {"auth_type": MCPAuth.bearer_token, "authentication_token": "STATIC-ADMIN-TOKEN"}, id="bearer_token" + ), + pytest.param({"auth_type": MCPAuth.none}, id="none"), + ], + ) + @pytest.mark.asyncio + async def test_byok_listing_keys_the_catalog_by_the_caller_and_never_touches_the_stored_secret( + self, server_auth: dict[str, object] + ): + """The caller's key plus what the caller supplied (nothing here) keys the catalog slot tools/call + reads, even with the stored BYOK secret at hand in the cache, and tools/list sends upstream exactly + what the caller supplied, so the static token, the M2M mint and MCPJWTSigner all behave as they + did before the catalog existed, whatever the auth_type.""" + from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache_key, + cache_byok_credential, + ) + from litellm.proxy._experimental.mcp_server.operations import byok_credential_cache + + manager = MCPServerManager() + server = MCPServer( + server_id="cc1", + name="cc1", + transport=MCPTransport.http, + url="http://cc1", + is_byok=True, + **server_auth, + ) + alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice") + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="echo", description="listed catalog", inputSchema={})] + ) + signer_headers = AsyncMock(return_value={"Authorization": "Bearer signed-jwt"}) + cache_byok_credential("alice", "cc1", "BYOK-ALICE-SECRET") + try: + with ( + patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=MagicMock(), + ), + patch( # test-quality-ok: same singleton's header injection, asserted on by call + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.inject_mcp_jwt_headers_for_upstream", + signer_headers, + ), + ): + await manager._get_tools_from_server(server=server, user_api_key_auth=alice, record_listing=True) + finally: + byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1")) + + client_kwargs = manager._create_mcp_client.await_args.kwargs + assert client_kwargs["mcp_auth_header"] is None, client_kwargs + assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"} + signer_headers.assert_awaited_once() + listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=alice)) + assert listed is not None and listed.description == "listed catalog" + + @pytest.mark.parametrize( + ("signer", "static_headers"), + [ + pytest.param(MagicMock(), None, id="signer"), + pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, id="static-authorization"), + pytest.param(None, None, id="no-signer"), + ], + ) + def test_keyed_callers_always_list_into_their_own_slot(self, signer, static_headers): + """The catalog is guardrail-shaped per key, so a keyed caller never reads another caller's + listing regardless of the signer or static authorization configuration.""" + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers + ) + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="alice", token="hashed-alice")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="bob", token="hashed-bob")) + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=signer, + ): + manager._record_listed_tools( + server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice + ) + assert manager.get_listed_tool(server, "turn", bob) is None + + for_alice = manager.get_listed_tool(server, "turn", alice) + assert for_alice is not None and for_alice.description == "alice view" + + def test_signed_server_slot_splits_on_the_callers_key_not_only_the_user(self): + """Two keys sharing a user_id get different signed JWTs, so they split; the same key + presented again lands on its own slot.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-beta")) + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=MagicMock(), + ): + manager._record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) + assert manager.get_listed_tool(server, "turn", bob) is None + + same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + listed = manager.get_listed_tool(server, "turn", same_key) + + assert listed is not None and listed.description == "slot a" + + def test_listed_tools_slot_is_split_per_team_for_keyless_callers(self): + """A team-only JWT admits a caller with neither a key nor a user, so the team keys the slot.""" + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + team_one: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one") + ) + team_two: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two") + ) + manager._record_listed_tools( + server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one + ) + + assert manager.get_listed_tool(server, "foo", team_two) is None + listed: Final = manager.get_listed_tool(server, "foo", team_one) + assert listed is not None and listed.description == "Fetch rows FLAGWORD" + + def test_listed_tools_slot_is_split_per_team_for_the_same_keyless_user(self): + """One JWT user acting in two teams is served two team-shaped catalogs, so each team is a slot.""" + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + alice_in_one: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-one") + ) + alice_in_two: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two") + ) + manager._record_listed_tools( + server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one + ) + + assert manager.get_listed_tool(server, "foo", alice_in_two) is None + listed: Final = manager.get_listed_tool(server, "foo", alice_in_one) + assert listed is not None and listed.description == "Fetch rows FLAGWORD" + + def test_listed_tools_slot_is_split_by_the_admission_bearer_of_keyless_callers_without_a_user(self): + """Two team-only JWT callers of one team differ only in the JWT they were admitted with, so that + credential keys the slot, on a server that never forwards it.""" + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + alice: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"), + raw_headers={"authorization": "Bearer jwt-alice"}, + ) + bob: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"), + raw_headers={"authorization": "Bearer jwt-bob"}, + ) + manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) + + assert manager.get_listed_tool(server, "foo", bob) is None + listed: Final = manager.get_listed_tool(server, "foo", alice) + assert listed is not None and listed.description == "alice view" + + @pytest.mark.parametrize( + ("server_kwargs", "forwards_bearer"), + [ + pytest.param( + {"auth_type": MCPAuth.oauth2, "delegate_auth_to_upstream": True, "oauth2_flow": "authorization_code"}, + True, + id="oauth2-delegated-to-upstream", + ), + pytest.param({"auth_type": MCPAuth.oauth_delegate}, True, id="oauth-delegate"), + pytest.param({"auth_type": MCPAuth.true_passthrough}, True, id="true-passthrough"), + pytest.param({"auth_type": MCPAuth.oauth2_token_exchange}, True, id="token-exchange"), + pytest.param( + {"auth_type": MCPAuth.none, "extra_headers": ["Authorization"], "oauth_passthrough": True}, + True, + id="oauth-passthrough", + ), + pytest.param({}, False, id="plain"), + pytest.param( + {"auth_type": MCPAuth.oauth2, "oauth2_flow": "authorization_code"}, + False, + id="oauth2-gateway-managed", + ), + pytest.param( + { + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "client_credentials", + "delegate_auth_to_upstream": True, + "client_id": "gateway", + "client_secret": "secret", + "token_url": "http://idp/token", + }, + False, + id="oauth2-client-credentials", + ), + ], + ) + def test_listed_tools_slot_is_split_by_the_forwarded_bearer_on_servers_that_forward_it( + self, server_kwargs: dict[str, object], forwards_bearer: bool + ): + """Two callers sharing one key but carrying different upstream bearers are served two upstream + catalogs exactly on the servers whose egress forwards or exchanges that bearer.""" + manager: Final = MCPServerManager() + server: Final = MCPServer( + **{"server_id": "dg", "name": "dg", "transport": MCPTransport.http, "url": "http://dg", **server_kwargs} + ) + caller_a: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"), + raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-A"}, + ) + caller_b: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"), + raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"}, + ) + manager._record_listed_tools( + server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a + ) + + for_b: Final = manager.get_listed_tool(server, "lookup", caller_b) + assert (for_b is None) is forwards_bearer + for_a: Final = manager.get_listed_tool(server, "lookup", caller_a) + assert for_a is not None and for_a.description == "Workspace A lookup FLAGWORD" + + @pytest.mark.asyncio + async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self): + manager = MCPServerManager() + server = MCPServer( + server_id="catalog", + name="catalog", + transport=MCPTransport.http, + url="http://catalog", + extra_headers=["X-Workspace"], + ) + manager.registry = {"catalog": server} + catalogs = { + "A": [ + MCPTool( + name="turn", description="Catalog A", inputSchema={"properties": {"turn": {"description": "A"}}} + ) + ], + "B": [ + MCPTool( + name="turn", description="Catalog B", inputSchema={"properties": {"turn": {"description": "B"}}} + ) + ], + } + mock_client = AsyncMock() + mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) + manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager._fetch_tools_with_timeout = AsyncMock(side_effect=lambda client, name: catalogs[client.workspace]) + for workspace in ("A", "B"): + manager._create_mcp_client.return_value.workspace = workspace + await manager._get_tools_from_server( + server=server, + extra_headers={"X-Workspace": workspace}, + raw_headers={"x-workspace": workspace, "authorization": "Bearer sk-litellm"}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"), + record_listing=True, + ) + + proxy_logging_obj = MagicMock() + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + await manager.call_tool( + server_name="catalog", + name="turn", + arguments={"turn": "A-1"}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"), + proxy_logging_obj=proxy_logging_obj, + raw_headers={"x-workspace": "A", "authorization": "Bearer sk-litellm"}, + ) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ( + "Catalog A", + {"properties": {"turn": {"description": "A"}}}, + ) + + def test_per_caller_listed_tools_evict_oldest_caller_and_keep_shared(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _LISTED_TOOLS_CALLERS_PER_SERVER + + manager = MCPServerManager() + server = MCPServer( + server_id="srv", + name="srv", + transport=MCPTransport.http, + url="http://srv", + auth_type=MCPAuth.oauth2_token_exchange, + ) + manager._record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) + callers = [ + ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}")) + for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1) + ] + for caller in callers: + manager._record_listed_tools( + server, + [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], + caller, + ) + manager._record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) + + assert manager.get_listed_tool(server, "read", callers[0]) is None + second = manager.get_listed_tool(server, "read", callers[1]) + assert second is not None and second.description == "u1 again" + newest = manager.get_listed_tool(server, "read", callers[-1]) + assert newest is not None and newest.description == callers[-1].user_api_key_auth.user_id + assert len(manager._listed_tools_by_server_id[server.server_id]) == _LISTED_TOOLS_CALLERS_PER_SERVER + 1 + shared = manager.get_listed_tool(server, "read") + assert shared is not None and shared.description == "shared" + + @pytest.mark.asyncio + @pytest.mark.parametrize("add_prefix", [True, False]) + async def test_openapi_listing_records_listed_tools(self, add_prefix): + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + server = MCPServer( + server_id="petstore-id", + name="petstore", + alias="petstore", + transport=MCPTransport.http, + url=None, + spec_path="/spec.yaml", + ) + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + + async def _handler(**kwargs): + return None + + global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + global_mcp_tool_registry.register_tool( + name="petstore-list_pets", + description="List pets", + input_schema={"type": "object", "properties": {"limit": {"type": "integer"}}}, + handler=_handler, + ) + try: + listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix, record_listing=True) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + + assert [t.name for t in listed] == ["petstore-list_pets" if add_prefix else "list_pets"] + tool = manager.get_listed_tool(server, "list_pets") + assert tool is not None and tool.description == "List pets" + assert tool.input_schema["properties"] == {"limit": {"type": "integer"}} + + @pytest.mark.asyncio + async def test_openapi_listing_ignores_overlapping_server_prefix(self): + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + server = MCPServer( + server_id="pet-id", + name="pet", + alias="pet", + transport=MCPTransport.http, + url=None, + spec_path="/spec.yaml", + ) + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + + async def _handler(**kwargs): + return None + + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + global_mcp_tool_registry.register_tool( + name="pet-petstore-list", + description="Local pet tool", + input_schema={"type": "object", "properties": {"limit": {"type": "integer"}}}, + handler=_handler, + ) + global_mcp_tool_registry.register_tool( + name="petstore-list", + description="Foreign petstore tool", + input_schema={"type": "object", "properties": {"status": {"type": "string"}}}, + handler=_handler, + ) + try: + listed = await manager._get_tools_from_server(server=server, add_prefix=True, record_listing=True) + finally: + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + + assert [t.name for t in listed] == ["pet-petstore-list"] + tool = manager.get_listed_tool(server, "petstore-list") + assert tool is not None and tool.description == "Local pet tool" + assert tool.input_schema["properties"] == {"limit": {"type": "integer"}} + + @pytest.mark.asyncio + @pytest.mark.parametrize("openapi", [False, True], ids=["remote", "openapi"]) + async def test_get_tools_from_server_records_the_catalog_only_when_asked_to(self, openapi): + """The startup fill, the implicit pre-call listing and the pin snapshot reuse this fetch without + serving its result, so only a listing that asks to be recorded sets what tools/call hooks see.""" + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + if openapi: + server = MCPServer( + server_id="srv", name="srv", alias="srv", transport=MCPTransport.http, url=None, spec_path="/spec.yaml" + ) + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + global_mcp_tool_registry.unregister_tools_with_prefix("srv-") + global_mcp_tool_registry.register_tool( + name="srv-echo", description="Echoes", input_schema={"type": "object"}, handler=lambda **kwargs: None + ) + else: + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister") + + try: + listed = await manager._get_tools_from_server(server=server, user_api_key_auth=user) + assert [t.name for t in listed] == ["srv-echo"] + assert server.server_id not in manager._listed_tools_by_server_id + + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("srv-") + + recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) + assert recorded is not None and recorded.description == "Echoes" + + @pytest.mark.asyncio + async def test_list_tools_records_the_served_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + manager.get_allowed_mcp_servers = AsyncMock(return_value=["srv"]) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister") + + listed = await manager.list_tools(user_api_key_auth=user) + + assert [t.name for t in listed] == ["srv-echo"] + recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) + assert recorded is not None and recorded.description == "Echoes" + + @pytest.mark.asyncio + async def test_startup_tool_name_mapping_records_no_listed_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + + await manager._initialize_tool_name_to_mcp_server_name_mapping() + + assert manager.server_exposes_tool(server, "echo") is True + assert server.server_id not in manager._listed_tools_by_server_id + assert manager.get_listed_tool(server, "echo") is None + + @pytest.mark.asyncio + async def test_get_tools_for_server_records_no_listed_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + + listed = await manager.get_tools_for_server("srv") + + assert [t.name for t in listed] == ["srv-echo"] + assert server.server_id not in manager._listed_tools_by_server_id + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_with_user_api_key_auth(self): """ @@ -9739,15 +10870,11 @@ class TestGetPublicMCPServers: if registered_in == "both" else server ) - manager.config_mcp_servers = ( - {server.server_id: config_server} if registered_in in ("config", "both") else {} - ) + manager.config_mcp_servers = {server.server_id: config_server} if registered_in in ("config", "both") else {} manager.registry = {server.server_id: server} if registered_in in ("database", "both") else {} original_server: Final = server.model_dump() original_config_server: Final = config_server.model_dump() - expected_public: Final = registered_in != "neither" and ( - public_ids == [server.server_id] or implicitly_public - ) + expected_public: Final = registered_in != "neither" and (public_ids == [server.server_id] or implicitly_public) with ( patch("litellm.public_mcp_servers", public_ids), @@ -9755,9 +10882,7 @@ class TestGetPublicMCPServers: ): public_servers: Final = manager.get_public_mcp_servers() assert manager.is_mcp_server_public(server.server_id) is expected_public - assert [item.server_id for item in public_servers] == ( - [server.server_id] if expected_public else [] - ) + assert [item.server_id for item in public_servers] == ([server.server_id] if expected_public else []) assert manager.is_mcp_server_public("server-alias") is False assert manager.is_mcp_server_public("missing-server") is False assert manager.is_mcp_server_public(server.server_id, public_ids=frozenset()) is ( @@ -10657,7 +11782,9 @@ class TestOBOConcurrencyLimit: inflight = {"current": 0, "peak": 0} class _ConcurrencyRecordingClient: - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False): + async def call_tool( + self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False + ): inflight["current"] += 1 inflight["peak"] = max(inflight["peak"], inflight["current"]) try: @@ -12992,7 +14119,9 @@ class TestConfigServerIdPinning: @pytest.mark.asyncio @pytest.mark.parametrize("aliasing_entry_first", [True, False]) - async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected(self, config_only_mcp_manager_factory, aliasing_entry_first: bool): + async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected( + self, config_only_mcp_manager_factory, aliasing_entry_first: bool + ): """A grant naming 'docs_server' reaches both servers unpinned; the pin would narrow it to one.""" manager = config_only_mcp_manager_factory() wiki = ( @@ -13008,7 +14137,9 @@ class TestConfigServerIdPinning: await manager.load_servers_from_config(dict((wiki, docs) if aliasing_entry_first else (docs, wiki))) @pytest.mark.asyncio - async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected(self, config_only_mcp_manager_factory): + async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected( + self, config_only_mcp_manager_factory + ): manager = config_only_mcp_manager_factory() with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"): @@ -13096,7 +14227,9 @@ class TestConfigServerIdPinning: assert second_round == first_round @pytest.mark.asyncio - async def test_shadow_warning_fires_again_when_the_shadowed_set_changes(self, config_only_mcp_manager_factory, caplog): + async def test_shadow_warning_fires_again_when_the_shadowed_set_changes( + self, config_only_mcp_manager_factory, caplog + ): manager = config_only_mcp_manager_factory() await manager.load_servers_from_config(self._config(server_id="docs-prod-1")) @@ -13267,7 +14400,9 @@ class TestConfigServerIdPinning: assert manager.config_mcp_servers["wiki"].url == "https://example.com/mcp" @pytest.mark.asyncio - async def test_a_row_that_shadows_one_id_still_reports_capturing_another(self, config_only_mcp_manager_factory, caplog): + async def test_a_row_that_shadows_one_id_still_reports_capturing_another( + self, config_only_mcp_manager_factory, caplog + ): """Skipping is per identifier, not per row, so the second collision is not lost.""" manager = config_only_mcp_manager_factory() await manager.load_servers_from_config( @@ -13639,7 +14774,8 @@ async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, ("none", {"Authorization": "Bearer injected"}, "extra-headers", "Bearer injected"), ], ) -async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_request_ctx, +async def test_debug_resolution_matches_final_header_conflict_winner( + _mcp_request_ctx, config: Literal["stored", "static", "none"], extra_headers: dict[str, str] | None, expected_source: str, @@ -13758,12 +14894,16 @@ async def test_debug_reports_legacy_signing_and_non_http_transport( async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + server_id="temporary-oauth-discovery", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, ) manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", ) with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: @@ -13783,13 +14923,18 @@ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publi async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", + server_id="repeated-stale", + name="stale", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + oauth2_flow="authorization_code", ) manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) with ( patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, @@ -13809,13 +14954,20 @@ async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + server_id="resolved-replacement", + name="replacement", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy( + update={ + "url": "https://new.example.com/mcp", + "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + } ) - replacement: Final = original.model_copy(update={ - "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", - "token_url": "https://new.example.com/token", - }) manager.registry[original.server_id] = replacement assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement @@ -13823,8 +14975,11 @@ async def test_stale_discovery_falls_back_to_resolved_registered_server() -> Non def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="stale-publication", name="publication", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + server_id="stale-publication", + name="publication", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) original_slot: Final = manager._oauth_discovery_slot(original.server_id) @@ -13840,9 +14995,13 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + server_id="expiring-session", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) @@ -13943,7 +15102,9 @@ async def test_openapi_health_reports_size_limit_as_unknown_and_caches_failure(r result = await manager.health_check_server(server.server_id) cached = await manager.health_check_server(server.server_id) assert result.status == "unknown" - assert result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + assert ( + result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + ) assert cached.health_check_error == result.health_check_error assert cached.last_health_check == result.last_health_check assert route.call_count == 1 @@ -13955,8 +15116,11 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( - server_id="cancelled-cache", name="cancelled-cache", transport=MCPTransport.http, - spec_path="https://93.184.216.34/cancelled-cache.json", auth_type=MCPAuth.none, + server_id="cancelled-cache", + name="cancelled-cache", + transport=MCPTransport.http, + spec_path="https://93.184.216.34/cancelled-cache.json", + auth_type=MCPAuth.none, ) manager.registry = {server.server_id: server} started = asyncio.Event() @@ -14085,7 +15249,9 @@ class _DiscoveryUpstream: def _discovery_server() -> MCPServer: - return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) + return MCPServer( + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http + ) @pytest.mark.asyncio @@ -14253,7 +15419,9 @@ async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) 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))) +@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 @@ -14570,26 +15738,45 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them( class TestProtectedCredentialPreparation: @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,credential", [ - (MCPAuth.bearer_token, None), - (MCPAuth.bearer_token, "Bearer"), - (MCPAuth.api_key, None), - (MCPAuth.basic, "Basic"), - ]) + @pytest.mark.parametrize( + "auth_type,credential", + [ + (MCPAuth.bearer_token, None), + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.api_key, None), + (MCPAuth.basic, "Basic"), + ], + ) @pytest.mark.parametrize("dispatch", ["managed", "local"]) async def test_openapi_dispatch_rejects_unusable_effective_credentials( - self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - auth_type: MCPAuthType, credential: str | None, dispatch: str, + self, + tmp_path: Path, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + auth_type: MCPAuthType, + credential: str | None, + dispatch: str, ) -> None: from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix spec_path: Final = tmp_path / "openapi.json" - spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"}, - "paths": {"/echo": {"get": {"operationId": "echo"}}}})) + spec_path.write_text( + json.dumps( + { + "openapi": "3.0.0", + "info": {"title": "Auth", "version": "1"}, + "paths": {"/echo": {"get": {"operationId": "echo"}}}, + } + ) + ) server: Final = MCPServer( - server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential, + server_id="dispatch-auth", + name="dispatch-auth", + url="https://upstream.example", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=credential, ) manager: Final = MCPServerManager() await manager._register_openapi_tools(str(spec_path), server, server.url) @@ -14612,14 +15799,21 @@ class TestProtectedCredentialPreparation: self, transport: MCPTransport, client_secret: str | None, subject: str | None ) -> None: server = MCPServer( - server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp", - transport=transport, auth_type=MCPAuth.oauth2_token_exchange, - client_id="gateway", client_secret=client_secret, - token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback", + server_id="incomplete-obo", + name="incomplete-obo", + url="https://upstream.example/mcp", + transport=transport, + auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway", + client_secret=client_secret, + token_exchange_endpoint="https://idp.example/token", + authentication_token="static-fallback", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header="Bearer override", subject_token=subject, + server, + mcp_auth_header="Bearer override", + subject_token=subject, ) assert exc.value.status_code == (401 if subject is None else 500) assert "static-fallback" not in str(exc.value.detail) @@ -14632,8 +15826,11 @@ class TestProtectedCredentialPreparation: self, auth_type: MCPAuthType, credential: str | dict[str, str] | None ) -> None: server = MCPServer( - server_id="empty-static", name="empty-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-static", + name="empty-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential) @@ -14641,16 +15838,22 @@ class TestProtectedCredentialPreparation: assert "credential" in str(exc.value.detail).lower() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,headers", [ - (MCPAuth.api_key, {"X-API-Key": "key"}), - (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), - ]) + @pytest.mark.parametrize( + "auth_type,headers", + [ + (MCPAuth.api_key, {"X-API-Key": "key"}), + (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), + ], + ) async def test_static_auth_accepts_actual_forwarded_credential( self, auth_type: MCPAuthType, headers: dict[str, str] ) -> None: server = MCPServer( - server_id="header-static", name="header-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="header-static", + name="header-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers) assert client._get_auth_headers() == headers @@ -14659,29 +15862,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange]) async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None: server = MCPServer( - server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="openapi-empty", + name="openapi-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, token_exchange_endpoint="https://idp.example/token", ) with pytest.raises(HTTPException) as exc: await MCPServerManager().resolve_openapi_upstream_auth( - mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None, - user_api_key_auth=None, forwarded_headers=None, + mcp_server=server, + oauth2_headers=None, + raw_headers=None, + mcp_auth_header=None, + user_api_key_auth=None, + forwarded_headers=None, ) assert exc.value.status_code in (401, 500) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,slot,value", [ - (MCPAuth.api_key, "X-API-Key", "token"), - (MCPAuth.authorization, "Authorization", "opaque-secret-value"), - (MCPAuth.authorization, "Authorization", "Bearer abc"), - (MCPAuth.authorization, "Authorization", "Custom abc"), - ]) + @pytest.mark.parametrize( + "auth_type,slot,value", + [ + (MCPAuth.api_key, "X-API-Key", "token"), + (MCPAuth.authorization, "Authorization", "opaque-secret-value"), + (MCPAuth.authorization, "Authorization", "Bearer abc"), + (MCPAuth.authorization, "Authorization", "Custom abc"), + ], + ) async def test_raw_static_credentials_are_forwarded_unchanged( - self, auth_type: MCPAuthType, slot: str, value: str, + self, + auth_type: MCPAuthType, + slot: str, + value: str, ) -> None: - server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value) + server = MCPServer( + server_id="raw-key", + name="raw-key", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, + ) client = await MCPServerManager()._create_mcp_client(server) assert client._resolved_auth is not None request = httpx.Request("GET", server.url) @@ -14695,17 +15917,24 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"]) @pytest.mark.parametrize("source", ["configured", "caller", "forwarded"]) async def test_raw_authorization_rejects_bare_schemes_before_dispatch( - self, respx_mock: MockRouter, value: str, source: str, + self, + respx_mock: MockRouter, + value: str, + source: str, ) -> None: server: Final = MCPServer( - server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.authorization, + server_id="raw-empty", + name="raw-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.authorization, authentication_token=value if source == "configured" else None, ) destination: Final = respx_mock.route().respond(200) with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, + server, + mcp_auth_header=value if source == "caller" else None, extra_headers={"Authorization": value} if source == "forwarded" else None, ) assert exc.value.status_code == 500 @@ -14713,9 +15942,15 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None: - server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True, - token_exchange_endpoint="https://idp.example/token") + server = MCPServer( + server_id="obo-byok", + name="obo-byok", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + is_byok=True, + token_exchange_endpoint="https://idp.example/token", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override") assert exc.value.status_code == 401 @@ -14723,41 +15958,66 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")]) async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None: - server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured) + server = MCPServer( + server_id="override", + name="override", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=configured, + ) client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override) assert client._get_auth_headers()["Authorization"] == override @pytest.mark.asyncio @pytest.mark.parametrize("token", [None, "shared"]) async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None: - server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token) + server = MCPServer( + server_id="empty-header", + name="empty-header", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=token, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "}) assert exc.value.status_code == 500 @pytest.mark.asyncio async def test_custom_slot_uses_its_actual_credential(self) -> None: - server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, - upstream_token_header="X-Custom", authentication_token="key") + server = MCPServer( + server_id="custom", + name="custom", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", + authentication_token="key", + ) client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"}) assert client._credential_slot == "X-Custom" assert await client.discovery_auth_fingerprint() @pytest.mark.asyncio - @pytest.mark.parametrize("static_headers,accepted", [ - ({"apikey": "static-key"}, True), - ({"apikey": ""}, False), - ({"X-Tenant": "tenant"}, True), - ]) + @pytest.mark.parametrize( + "static_headers,accepted", + [ + ({"apikey": "static-key"}, True), + ({"apikey": ""}, False), + ({"X-Tenant": "tenant"}, True), + ], + ) async def test_api_key_carried_by_static_header_passes_fail_closed_check( self, static_headers: dict[str, str], accepted: bool ) -> None: server: Final = MCPServer( - server_id="static-slot", name="static-slot", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, static_headers=static_headers, + server_id="static-slot", + name="static-slot", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + static_headers=static_headers, ) if not accepted: with pytest.raises(HTTPException) as exc: @@ -14769,21 +16029,36 @@ class TestProtectedCredentialPreparation: assert all(request.headers[name] == value for name, value in static_headers.items()) @pytest.mark.asyncio - @pytest.mark.parametrize("static,forwarded,caller", [ - ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), - ({}, {"X-API-Key": "forwarded"}, None), - ({}, None, "ApiKey caller"), - ({"X-API-Key": "static"}, {"Authorization": ""}, None), - ]) + @pytest.mark.parametrize( + "static,forwarded,caller", + [ + ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), + ({}, {"X-API-Key": "forwarded"}, None), + ({}, None, "ApiKey caller"), + ({"X-API-Key": "static"}, {"Authorization": ""}, None), + ], + ) async def test_openapi_static_credentials_remain_supported( - self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None + self, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + static: dict[str, str], + forwarded: dict[str, str] | None, + caller: str | None, ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, _request_extra_headers, create_tool_function, + _request_auth_header, + _request_extra_headers, + create_tool_function, ) + tool: Final = create_tool_function( - "/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key, + "/echo", + "get", + {}, + "https://upstream.example", + headers=static, + auth_type=MCPAuth.api_key, ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") @@ -14817,8 +16092,13 @@ class TestProtectedCredentialPreparation: self.closed = True auth = CancelledAuth() - server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key) + server = MCPServer( + server_id="cancel", + name="cancel", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + ) client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth) with pytest.raises(asyncio.CancelledError): await prepare_mcp_client(server, client) @@ -14827,8 +16107,14 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization]) async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None: - server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ") + server = MCPServer( + server_id="blank-static", + name="blank-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=" ", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server) assert exc.value.status_code == 500 @@ -14836,8 +16122,13 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="]) async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None: - server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic) + server = MCPServer( + server_id="bad-basic", + name="bad-basic", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header}) assert exc.value.status_code == 500 @@ -14846,34 +16137,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None: - server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, - authentication_token=value if source == "configured" else None) + server = MCPServer( + server_id="basic-scheme", + name="basic-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value if source == "configured" else None, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None) assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,default_slot", [ - (MCPAuth.api_key, "fixture-key", "X-API-Key"), - (MCPAuth.bearer_token, "fixture-key", "Authorization"), - (MCPAuth.basic, "user:pass", "Authorization"), - (MCPAuth.token, "fixture-key", "Authorization"), - (MCPAuth.authorization, "fixture-key", "Authorization"), - ]) + @pytest.mark.parametrize( + "auth_type,value,default_slot", + [ + (MCPAuth.api_key, "fixture-key", "X-API-Key"), + (MCPAuth.bearer_token, "fixture-key", "Authorization"), + (MCPAuth.basic, "user:pass", "Authorization"), + (MCPAuth.token, "fixture-key", "Authorization"), + (MCPAuth.authorization, "fixture-key", "Authorization"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_usable_credential_survives_an_empty_alternate_header( self, auth_type: MCPAuthType, value: str, default_slot: str, source: str ) -> None: server: Final = MCPServer( - server_id="alternate", name="alternate", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom", + server_id="alternate", + name="alternate", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + upstream_token_header="X-Custom", authentication_token=value if source == "configured" else None, ) empty_slot: Final = default_slot if source == "configured" else "X-Custom" selected_slot: Final = "X-Custom" if source == "configured" else default_slot client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""}, + server, + mcp_auth_header=value if source == "caller" else None, + extra_headers={empty_slot: ""}, ) request: Final = await client.prepare_request_auth() assert request.headers[selected_slot] @@ -14882,8 +16187,12 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None: server: Final = MCPServer( - server_id="both-empty", name="both-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom", + server_id="both-empty", + name="both-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""}) @@ -14896,12 +16205,17 @@ class TestProtectedCredentialPreparation: self, custom_slot: str | None, source: str ) -> None: server: Final = MCPServer( - server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot, + server_id="caller-auth", + name="caller-auth", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header=custom_slot, ) headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""} client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=headers if source == "caller" else None, + server, + mcp_auth_header=headers if source == "caller" else None, extra_headers=headers if source == "forwarded" else None, ) request: Final = await client.prepare_request_auth() @@ -14910,14 +16224,29 @@ class TestProtectedCredentialPreparation: assert custom_slot is None or custom_slot not in request.headers @pytest.mark.asyncio - @pytest.mark.parametrize("value", [ - "", " ", "Bearer", "Basic", "token", "ApiKey", - "Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY", - ]) + @pytest.mark.parametrize( + "value", + [ + "", + " ", + "Bearer", + "Basic", + "token", + "ApiKey", + "Bearer Bearer", + "ApiKey ApiKey", + "token token", + "bEaReR BEARER", + "aPiKeY\tAPIKEY", + ], + ) async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None: server: Final = MCPServer( - server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, + server_id="caller-empty", + name="caller-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value}) @@ -14928,8 +16257,11 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None: server: Final = MCPServer( - server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, + server_id="basic-pair", + name="basic-pair", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14942,8 +16274,12 @@ class TestProtectedCredentialPreparation: import base64 server: Final = MCPServer( - server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value, + server_id="basic-valid", + name="basic-valid", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14952,17 +16288,27 @@ class TestProtectedCredentialPreparation: assert base64.b64decode(encoded) == value.encode() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value", [ - (MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"), - (MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"), - ]) + @pytest.mark.parametrize( + "auth_type,value", + [ + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.bearer_token, "Bearer "), + (MCPAuth.bearer_token, "bearer"), + (MCPAuth.token, "token"), + (MCPAuth.token, "token "), + (MCPAuth.token, "TOKEN"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix( self, auth_type: MCPAuthType, value: str, source: str ) -> None: server: Final = MCPServer( - server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-scheme", + name="empty-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14970,17 +16316,24 @@ class TestProtectedCredentialPreparation: assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,expected", [ - (MCPAuth.bearer_token, "token", "Bearer token"), - (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), - (MCPAuth.token, "tokenish", "token tokenish"), - ]) + @pytest.mark.parametrize( + "auth_type,value,expected", + [ + (MCPAuth.bearer_token, "token", "Bearer token"), + (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), + (MCPAuth.token, "tokenish", "token tokenish"), + ], + ) async def test_static_credentials_that_resemble_schemes_remain_usable( self, auth_type: MCPAuthType, value: str, expected: str ) -> None: server: Final = MCPServer( - server_id="real-token", name="real-token", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value, + server_id="real-token", + name="real-token", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -15019,16 +16372,31 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon registry.register_tool("observer-execute", "Execute", {"type": "object"}, upstream) monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) manager = MCPServerManager() - manager.registry = {"observer": MCPServer( - server_id="observer", name="observer", server_name="observer", transport="http", - url="https://observer.example/mcp", spec_path="observer.json", auth_type="none", - )} + manager.registry = { + "observer": MCPServer( + server_id="observer", + name="observer", + server_name="observer", + transport="http", + url="https://observer.example/mcp", + spec_path="observer.json", + auth_type="none", + ) + } manager.tool_name_to_mcp_server_name_mapping = {"observer-execute": "observer"} - result = await asyncio.wait_for(manager.call_tool( - server_name="observer", name="execute", arguments={"text": "hello"}, - user_api_key_auth=UserAPIKeyAuth(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), - guardrail_context=MCPRequestContext.resolve_guardrail_context({"metadata": {"guardrails": ["observe"] if selected else []}}), - ), timeout=5) + result = await asyncio.wait_for( + manager.call_tool( + server_name="observer", + name="execute", + arguments={"text": "hello"}, + user_api_key_auth=UserAPIKeyAuth(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + guardrail_context=MCPRequestContext.resolve_guardrail_context( + {"metadata": {"guardrails": ["observe"] if selected else []}} + ), + ), + timeout=5, + ) assert tool_started.is_set() assert guardrail_started.is_set() is selected assert result.is_error is False @@ -15057,11 +16425,21 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback - upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) + upstream = MCPServer( + server_id="explicit-empty", + name="explicit_empty", + url="https://example.invalid/mcp", + transport=MCPTransport.http, + allow_sampling=True, + ) token = auth_context_var.set(None) sampling = AsyncMock() try: - legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99") + legacy_server.set_auth_context( + UserAPIKeyAuth(user_id="unrelated"), + raw_headers={"authorization": "unrelated-credential"}, + client_ip="192.0.2.99", + ) with ( patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as factory, patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), @@ -15069,7 +16447,9 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie if legacy_factory: callback = _create_sampling_callback(user_api_key_auth=UserAPIKeyAuth(user_id="explicit")) else: - await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) + await MCPServerManager()._create_mcp_client( + upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None + ) callback = factory.call_args.kwargs["sampling_callback"] await callback(None, None) captured = sampling.await_args.kwargs @@ -15092,16 +16472,28 @@ class TestSharedIdentifierPrefixWarning: manager = MCPServerManager() rows = [ LiteLLM_MCPServerTable( - server_id="srv-a", server_name="alpha", alias="shared", url="https://a.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-a", + server_name="alpha", + alias="shared", + url="https://a.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-b", server_name="beta", alias="Shared", url="https://b.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-b", + server_name="beta", + alias="Shared", + url="https://b.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-c", server_name="gamma", alias="lonely", url="https://c.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-c", + server_name="gamma", + alias="lonely", + url="https://c.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), ] raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] @@ -15193,7 +16585,9 @@ async def test_reload_warns_once_about_a_blocked_stdio_row_that_is_rebuilt_every @pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]) async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision): manager = config_only_mcp_manager_factory() - await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}}) + await manager.load_servers_from_config( + {"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}} + ) server = next(iter(manager.config_mcp_servers.values())) client = await manager._create_mcp_client(server) assert server.protocol_version == revision @@ -15205,11 +16599,15 @@ async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_m def test_runtime_protocol_metadata_preserves_explicit_precedence( revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None ) -> None: - server: Final = MCPServer.model_validate({ - "server_id": "preview", "name": "preview", "transport": "http", - "mcp_info": {"protocol_version": revision}, - **({"protocol_version": explicit} if explicit is not None else {}), - }) + server: Final = MCPServer.model_validate( + { + "server_id": "preview", + "name": "preview", + "transport": "http", + "mcp_info": {"protocol_version": revision}, + **({"protocol_version": explicit} if explicit is not None else {}), + } + ) assert server.protocol_version == (explicit if explicit is not None else revision) @@ -15387,9 +16785,7 @@ class TestToolCatalogGuard: proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=asyncio.CancelledError) with pytest.raises(asyncio.CancelledError): - await manager._get_tools_from_server( - _notes_server(), add_prefix=False, proxy_logging_obj=proxy_logging_obj - ) + await manager._get_tools_from_server(_notes_server(), add_prefix=False, proxy_logging_obj=proxy_logging_obj) proxy_logging_obj.pre_call_hook.assert_awaited_once() proxy_logging_obj.slack_alerting_instance.send_alert.assert_not_awaited() @@ -15446,7 +16842,10 @@ class TestToolCatalogGuard: guardrail, proxy_logging_obj = catalog_guardrail monkeypatch.setattr(signer_module, "_mcp_jwt_signer_instance", None) signer = signer_module.MCPJWTSigner( - guardrail_name="jwt-signer", event_hook="pre_mcp_call", default_on=True, issuer="https://litellm.example.com" + guardrail_name="jwt-signer", + event_hook="pre_mcp_call", + default_on=True, + issuer="https://litellm.example.com", ) monkeypatch.setattr(litellm, "callbacks", [signer, guardrail]) manager = _catalog_manager(LIST_NOTES, POISONED_DELETE) @@ -15666,7 +17065,10 @@ class TestToolCatalogGuard: widened = MCPTool( name="read_note", description="Read a note", - inputSchema={"type": "object", "properties": {"id": {"type": "string"}, "callback_url": {"type": "string"}}}, + inputSchema={ + "type": "object", + "properties": {"id": {"type": "string"}, "callback_url": {"type": "string"}}, + }, ) manager = _catalog_manager(widened) @@ -15728,8 +17130,12 @@ class TestToolCatalogGuard: return "ok" with patch.dict(global_mcp_tool_registry.tools, {}, clear=True): - global_mcp_tool_registry.register_tool("petstore-list_pets", "List pets, newest first", {"type": "object"}, handler) - global_mcp_tool_registry.register_tool("petstore-delete_pets", POISONED_DELETE.description, {"type": "object"}, handler) + global_mcp_tool_registry.register_tool( + "petstore-list_pets", "List pets, newest first", {"type": "object"}, handler + ) + global_mcp_tool_registry.register_tool( + "petstore-delete_pets", POISONED_DELETE.description, {"type": "object"}, handler + ) global_mcp_tool_registry.register_tool("petstore-find_pet", "Find a pet", {"type": "object"}, handler) served = await manager._get_tools_from_server( server, add_prefix=add_prefix, proxy_logging_obj=proxy_logging_obj diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 4cc7794d4ad..545b2757ffd 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -9,6 +9,7 @@ from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -23,18 +24,20 @@ from mcp.types import ( TextContent, TextResourceContents, ) +from mcp.types import Tool as MCPTool from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS from pydantic import TypeAdapter from starlette.types import Message, Receive, Scope, Send from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller from litellm.proxy._types import ( LiteLLM_MCPServerTable, MCPTransport, UserAPIKeyAuth, ) from litellm.types.mcp import MCPAuth -from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer +from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool def test_mcp_available_on_sdk2(): @@ -85,9 +88,6 @@ def cleanup_mcp_global_state(): yield - - - def _call_tool_params(name, arguments=None): from mcp.types import CallToolRequestParams @@ -99,6 +99,7 @@ def _paged_params(): return PaginatedRequestParams() + @pytest.mark.asyncio async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx): """Test that proxy_server_request body contains name and arguments""" @@ -295,7 +296,9 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r ): with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()): with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger): - result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})) + result = await mcp_server_tool_call( + _mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}) + ) assert result.is_error is True # The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this @@ -1167,20 +1170,32 @@ async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind, else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata) ) with ( - patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))), + patch.object( + server, + "get_or_extract_auth_context", + AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None)), + ), patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])), - patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))), + patch.object( + operations.global_mcp_server_manager, + "read_resource_from_server", + AsyncMock(return_value=ReadResourceResult(contents=[content])), + ), ): result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri)) assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == { - "cacheScope": "private", "resultType": "complete", "ttlMs": 0, - "contents": [{ - "uri": uri, - "mimeType": "text/plain" if kind == "text" else "image/png", - "text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=", - **({"_meta": metadata} if metadata is not None else {}), - }], + "cacheScope": "private", + "resultType": "complete", + "ttlMs": 0, + "contents": [ + { + "uri": uri, + "mimeType": "text/plain" if kind == "text" else "image/png", + "text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=", + **({"_meta": metadata} if metadata is not None else {}), + } + ], } @@ -1674,7 +1689,9 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error( with ( patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam "litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context", - new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None), + new=AsyncMock( + return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None + ), ), patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam "litellm.proxy._experimental.mcp_server.operations._list_mcp_tools", @@ -1913,8 +1930,8 @@ async def test_streamable_http_session_manager_is_stateless(): ("DELETE", b"", False), ), ) -async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_request_ctx, - debug: bool, method: str, request_body: bytes, stateful: bool +async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless( + _mcp_request_ctx, debug: bool, method: str, request_body: bytes, stateful: bool ) -> None: from starlette.requests import Request from starlette.types import Message, Receive, Scope, Send @@ -4056,7 +4073,8 @@ async def test_truncated_jsonrpc_response_with_nested_method_skips_lock( # parsed, with a nested "method" key in the first bytes to trip a flat # substring heuristic. response_prefix: Final = ( - '{"jsonrpc":"2.0","id":99,"' + response_field + '{"jsonrpc":"2.0","id":99,"' + + response_field + '":{"code":-32000,"message":"test","data":{"method":"GET","payload":"' ).encode() response_body: Final = ( @@ -6564,8 +6582,12 @@ class TestGatewayCreateInitializationOptions: yield (None, None) async def record_request( - serving_server: object, read_stream: object, write_stream: object, - *, lifespan_state: object, init_options: InitializationOptions, + serving_server: object, + read_stream: object, + write_stream: object, + *, + lifespan_state: object, + init_options: InitializationOptions, ) -> None: captured["server_name"] = init_options.server_name @@ -6877,7 +6899,6 @@ async def test_probe_upstream_auth_surfaces_httpx_status_error(): returning the response. The probe must catch that specifically (before the fail-open `except Exception`) so the auth check is not silently defeated. """ - import httpx from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth @@ -7412,7 +7433,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool return_value=oauth_server, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7658,7 +7680,8 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator(): return_value=alias_less_server, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7933,7 +7956,8 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req return_value=None, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7989,9 +8013,12 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): fake_server.server_name = "openapi-petstore" fake_server.alias = None fake_server.short_prefix = None + fake_server.tool_name_to_description = None fake_tool = MagicMock() fake_tool.name = "list_pets" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} start_time = datetime.now(timezone.utc) litellm_logging_obj, _ = function_setup( @@ -8042,6 +8069,348 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): assert litellm_logging_obj.model == "MCP: list_pets" +@pytest.mark.asyncio +async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing_before_a_listing(): + """A local-registry tools/call with no prior tools/list hands the pre-call hooks name and arguments + only, as before this metadata existed, so a pre_mcp_call policy never scans a description the caller was + not served. Once the caller has listed, the same call hands the entry that listing served.""" + from litellm.caching.caching import DualCache + from litellm.proxy._experimental.mcp_server import operations as mcp_module + from litellm.proxy.utils import ProxyLogging + + petstore = MCPServer( + server_id="petstore-id", + name="petstore", + server_name="petstore", + transport=MCPTransport.http, + url=None, + spec_path="https://example.com/petstore.yaml", + tool_name_to_description={"list_pets": "ADMIN DESC"}, + ) + schema = {"type": "object", "properties": {"limit": {"type": "integer"}}} + mcp_module.global_mcp_tool_registry.register_tool( + name="petstore-list_pets", description="List the pets", input_schema=schema, handler=lambda limit: "ok" + ) + manager = mcp_module.global_mcp_server_manager + manager._listed_tools_by_server_id.pop(petstore.server_id, None) + alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging.pre_call_hook = AsyncMock(return_value={}) + pre_call_tool_check = AsyncMock(wraps=manager.pre_call_tool_check) + + async def call() -> tuple[MCPTool | None, dict]: + await mcp_module.execute_mcp_tool( + name="petstore-list_pets", + arguments={"limit": 10}, + allowed_mcp_servers=[petstore], + start_time=datetime.now(), + user_api_key_auth=alice, + ) + return pre_call_tool_check.call_args.kwargs["tool"], proxy_logging.pre_call_hook.call_args.kwargs["data"] + + try: + with ( + patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + ): + never_listed_tool, never_listed_data = await call() + manager._record_listed_tools( + petstore, + [MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)], + ListedToolsCaller(user_api_key_auth=alice), + ) + listed_tool, listed_data = await call() + finally: + mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + manager._listed_tools_by_server_id.pop(petstore.server_id, None) + + assert never_listed_tool is None + assert (never_listed_data.get("mcp_tool_description"), never_listed_data.get("mcp_input_schema")) == (None, None) + assert listed_tool is not None and (listed_tool.description, listed_tool.input_schema) == ("ADMIN DESC", schema) + assert (listed_data["mcp_tool_description"], listed_data["mcp_input_schema"]) == ("ADMIN DESC", schema) + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_clients_saw(): + """When tools/list pinned the schema and masked the description of an OpenAPI tool, the local-registry + call path must hand the pre-call hooks that served entry, not the raw registry one.""" + from litellm.proxy._experimental.mcp_server import operations as mcp_module + + petstore = MCPServer( + server_id="petstore-id", + name="petstore", + server_name="petstore", + transport=MCPTransport.http, + url=None, + spec_path="https://example.com/petstore.yaml", + tool_name_to_description={"getpetbyid": "Find a SECRET pet"}, + ) + registry_schema = {"type": "object", "properties": {"petId": {"type": "integer"}, "dump_all": {"type": "boolean"}}} + pinned_schema = {"type": "object", "properties": {"petId": {"type": "integer"}}} + mcp_module.global_mcp_tool_registry.register_tool( + name="petstore-getpetbyid", + description="Find pet by ID", + input_schema=registry_schema, + handler=lambda petId: "ok", + ) + manager = mcp_module.global_mcp_server_manager + alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") + manager._record_listed_tools( + petstore, + [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)], + ListedToolsCaller(user_api_key_auth=alice), + ) + pre_call_tool_check = AsyncMock(return_value={}) + + try: + with ( + patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), + ): + await mcp_module.execute_mcp_tool( + name="petstore-getpetbyid", + arguments={"petId": 1}, + allowed_mcp_servers=[petstore], + start_time=datetime.now(), + user_api_key_auth=alice, + ) + finally: + mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + manager._listed_tools_by_server_id.pop(petstore.server_id, None) + + handed_tool = pre_call_tool_check.call_args.kwargs["tool"] + assert (handed_tool.description, handed_tool.input_schema) == ("Find a [MASKED] pet", pinned_schema), ( + "the pre-call policy must evaluate the entry tools/list served, not the raw registry entry" + ) + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entry(): + """Two keys can be shown differently guarded OpenAPI catalogs. The call path must evaluate each key + against the entry its own tools/list served, not the entry the most recent listing left behind.""" + from litellm.proxy._experimental.mcp_server import operations as mcp_module + + petstore = MCPServer( + server_id="petstore-id", + name="petstore", + server_name="petstore", + transport=MCPTransport.http, + url=None, + spec_path="https://example.com/petstore.yaml", + ) + schema = {"type": "object", "properties": {"petId": {"type": "integer"}}} + mcp_module.global_mcp_tool_registry.register_tool( + name="petstore-getpetbyid", description="Find a SECRET pet", input_schema=schema, handler=lambda petId: "ok" + ) + manager = mcp_module.global_mcp_server_manager + guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice") + opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob") + manager._record_listed_tools( + petstore, + [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)], + ListedToolsCaller(user_api_key_auth=guarded), + ) + manager._record_listed_tools( + petstore, + [MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)], + ListedToolsCaller(user_api_key_auth=opted_out), + ) + pre_call_tool_check = AsyncMock(return_value={}) + + try: + with ( + patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), + ): + for caller in (guarded, opted_out): + await mcp_module.execute_mcp_tool( + name="petstore-getpetbyid", + arguments={"petId": 1}, + allowed_mcp_servers=[petstore], + start_time=datetime.now(), + user_api_key_auth=caller, + ) + finally: + mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + manager._listed_tools_by_server_id.pop(petstore.server_id, None) + + handed = [call.kwargs["tool"].description for call in pre_call_tool_check.call_args_list] + assert handed == ["Find a [MASKED] pet", "Find a SECRET pet"], ( + "each key's tools/call must be evaluated against the OpenAPI entry its own listing served" + ) + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_hooks_no_registry_metadata(): + """An OpenAPI operation whose name starts with its own server prefix runs instead of the shorter one, and + with no prior listing the pre-call hooks get name and arguments only, never either registry entry.""" + from litellm.proxy._experimental.mcp_server import operations as mcp_module + + petstore = MCPServer( + server_id="petstore-id", + name="petstore", + server_name="petstore", + transport=MCPTransport.http, + url=None, + spec_path="https://example.com/petstore.yaml", + ) + registry = mcp_module.global_mcp_tool_registry + registry.register_tool(name="petstore-get_pet", description="short", input_schema={}, handler=lambda: "short") + registry.register_tool( + name="petstore-petstore-get_pet", + description="long", + input_schema={"type": "object", "properties": {"petId": {"type": "integer"}}}, + handler=lambda: "long", + ) + manager = mcp_module.global_mcp_server_manager + pre_call_tool_check = AsyncMock(return_value={}) + + try: + with ( + patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), + ): + result = await mcp_module.execute_mcp_tool( + name="petstore-petstore-get_pet", + arguments={}, + allowed_mcp_servers=[petstore], + start_time=datetime.now(), + user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"), + ) + finally: + registry.unregister_tools_with_prefix("petstore-") + + assert pre_call_tool_check.call_args.kwargs["tool"] is None + assert result.content[0].text == "long" + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation_named_after_a_listed_one(): + """After the caller listed ``get_pet``, a call to the never-listed ``petstore-get_pet`` operation hands the + pre-call hooks name and arguments only, not the listed sibling's description and schema.""" + from litellm.proxy._experimental.mcp_server import operations as mcp_module + + petstore = MCPServer( + server_id="petstore-id", + name="petstore", + server_name="petstore", + transport=MCPTransport.http, + url=None, + spec_path="https://example.com/petstore.yaml", + ) + registry = mcp_module.global_mcp_tool_registry + registry.register_tool( + name="petstore-petstore-get_pet", description="long", input_schema={}, handler=lambda: "long" + ) + manager = mcp_module.global_mcp_server_manager + alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") + manager._record_listed_tools( + petstore, + [MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})], + ListedToolsCaller(user_api_key_auth=alice), + ) + pre_call_tool_check = AsyncMock(return_value={}) + + try: + with ( + patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), + ): + result = await mcp_module.execute_mcp_tool( + name="petstore-petstore-get_pet", + arguments={}, + allowed_mcp_servers=[petstore], + start_time=datetime.now(), + user_api_key_auth=alice, + ) + finally: + registry.unregister_tools_with_prefix("petstore-") + manager._listed_tools_by_server_id.pop(petstore.server_id, None) + + assert pre_call_tool_check.call_args.kwargs["tool"] is None + assert result.content[0].text == "long" + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hooks_no_description(): + """The listing tools/call runs on its own when this worker does not yet expose the tool is never served + to the caller, so it leaves the caller's listed slot empty and the pre-call hooks still get name and + arguments only, as on main.""" + manager = mcp_operations.global_mcp_server_manager + server = _never_listed_passthrough_server() + manager.registry[server.server_id] = server + manager._listed_tools_by_server_id.pop(server.server_id, None) + upstream = AsyncMock() + upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) + proxy_logging = _mock_mcp_proxy_logging() + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.pre_call_hook = AsyncMock(return_value={}) + proxy_logging.during_call_hook = AsyncMock(return_value=None) + fetch_tools = AsyncMock( + return_value=[MCPTool(name="add", description="Adds. FLAGWORD", inputSchema={"type": "object"})] + ) + + with ( + patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)), + patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + ): + result = await mcp_operations.execute_mcp_tool( + name="lazy_map-add", + arguments={"a": 1, "b": 2}, + allowed_mcp_servers=[server], + start_time=datetime.now(), + mcp_auth_header="Bearer caller-token", + raw_headers={"authorization": "Bearer caller-token"}, + ) + + assert fetch_tools.await_count == 1 + assert upstream.call_tool.await_count == 1 + assert result.content[0].text == "ok" + hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) + assert server.server_id not in manager._listed_tools_by_server_id + + +@pytest.mark.asyncio +async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_admin(): + """The pin snapshot lists the raw upstream catalog, without the catalog guard or the admin's description + overrides, so it must not become what the admin's own later tools/call is evaluated against.""" + from litellm.caching.caching import DualCache + from litellm.proxy._experimental.mcp_server.rest_endpoints import fetch_pinnable_tool_catalog + from litellm.proxy.utils import ProxyLogging + + manager = mcp_operations.global_mcp_server_manager + server = MCPServer( + server_id="pin-srv", + name="pin_srv", + transport=MCPTransport.http, + url="https://up.example.com/mcp", + tool_name_to_description={"add": "Admin wording"}, + ) + manager._listed_tools_by_server_id.pop(server.server_id, None) + admin = UserAPIKeyAuth(api_key="sk-admin", user_id="admin") + request = MagicMock() + request.client.host = "10.1.2.3" + request.headers = {"x-litellm-api-key": "sk-admin"} + fetch_tools = AsyncMock( + return_value=[MCPTool(name="add", description="Upstream wording", inputSchema={"type": "object"})] + ) + + with ( + patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())), + patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), + patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())), + ): + snapshot = await fetch_pinnable_tool_catalog(server, request, admin) + + assert snapshot == {"add": PinnedMCPTool(description="Upstream wording", input_schema={"type": "object"})} + assert server.server_id not in manager._listed_tools_by_server_id + assert manager.get_listed_tool(server, "add", ListedToolsCaller(user_api_key_auth=admin)) is None + + @pytest.mark.asyncio async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requested_server(): """A prefixed REST name that resolves to no tool must still dispatch to the server_id. @@ -8098,7 +8467,8 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste return_value=None, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -8597,7 +8967,9 @@ class TestMCPMetaTraceCarrier: assert _mcp_meta_trace_carrier(None) is None assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None - only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta + only_progress = CallToolRequestParams.model_validate( + {"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False + ).meta assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None @@ -10394,7 +10766,9 @@ async def test_mcp_origin_admission_precedes_authentication( patch("litellm.proxy.proxy_server.origins", allowed_origins), patch.object(server, "extract_mcp_auth_context", authenticate), ): - async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client: + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=server.app), base_url="http://gateway" + ) as client: response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers)) assert response.status_code == expected_status @@ -10477,12 +10851,15 @@ async def test_streamable_http_rejects_modern_protocol_version( @pytest.mark.asyncio -@pytest.mark.parametrize("handler_name,field", [ - ("handle_list_tools", "tools"), - ("list_prompts", "prompts"), - ("list_resources", "resources"), - ("list_resource_templates", "resource_templates"), -]) +@pytest.mark.parametrize( + "handler_name,field", + [ + ("handle_list_tools", "tools"), + ("list_prompts", "prompts"), + ("list_resources", "resources"), + ("list_resource_templates", "resource_templates"), + ], +) async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field): from litellm.proxy._experimental.mcp_server import server @@ -10500,7 +10877,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai auth = UserAPIKeyAuth(user_id="denied-caller") denial = HTTPException(status_code=403, detail="scope denied") logger = MagicMock() - logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None) + logger.post_call_failure_hook = AsyncMock( + side_effect=RuntimeError("log unavailable") if failure_hook_raises else None + ) upstream = AsyncMock() with ( patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)), @@ -10509,7 +10888,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), ): with pytest.raises(HTTPException) as rejected: - await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True) + await operations._get_tools_from_mcp_servers( + user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True + ) assert rejected.value is denial upstream.assert_not_awaited() logger.post_call_failure_hook.assert_awaited_once() @@ -10521,7 +10902,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai @pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/"))) @pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS)) async def test_legacy_sse_mount_emits_message_endpoint( - prefix: str, suffix: str, opening_protocol: str | None, + prefix: str, + suffix: str, + opening_protocol: str | None, ) -> None: from starlette.applications import Starlette from starlette.routing import Mount @@ -10586,16 +10969,20 @@ async def test_legacy_sse_mount_emits_message_endpoint( return (await messages.get())["status"] if opening_protocol is not None: - discover: Final = json.dumps({ - "jsonrpc": "2.0", - "id": 0, - "method": "server/discover", - "params": {"_meta": { - "io.modelcontextprotocol/protocolVersion": opening_protocol, - "io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"}, - "io.modelcontextprotocol/clientCapabilities": {}, - }}, - }).encode() + discover: Final = json.dumps( + { + "jsonrpc": "2.0", + "id": 0, + "method": "server/discover", + "params": { + "_meta": { + "io.modelcontextprotocol/protocolVersion": opening_protocol, + "io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"}, + "io.modelcontextprotocol/clientCapabilities": {}, + } + }, + } + ).encode() assert await post(discover) == 202 discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode() discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0]) @@ -10628,7 +11015,16 @@ async def test_legacy_sse_mount_emits_message_endpoint( patch.object( mcp_server, "extract_mcp_auth_context", - AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})), + AsyncMock( + return_value=( + post_auth, + None, + [marker], + {marker: {"Authorization": marker}}, + {"Authorization": marker}, + {"x-request-marker": marker}, + ) + ), ), patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing), ): @@ -10677,7 +11073,11 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct dispatched = AsyncMock(return_value=expected) auth = UserAPIKeyAuth(user_id="discover-caller") with ( - patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))), + patch.object( + server, + "get_or_extract_auth_context", + AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None)), + ), patch.object(server.operations.GatewayOperations, "execute", dispatched), ): result = await server.discover(_mcp_request_ctx(), RequestParams()) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 1cfe6198b6f..1c987778da6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -23,6 +23,7 @@ from mcp.types import Tool import litellm from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller from litellm.proxy._experimental.mcp_server.tool_search import ( AGENT_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME, @@ -32,12 +33,14 @@ from litellm.proxy._experimental.mcp_server.tool_search import ( ToolSearchResult, coerce_top_k, get_virtual_tool_definitions, + handle_mcp_tool_search, search_mcp_tools, search_tools, ) from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector -from litellm.types.mcp import MCPToolSearchSettings +from litellm.types.mcp import MCPToolSearchSettings, MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer def _make_tools(specs: list[tuple[str, str]]) -> tuple[Tool, ...]: @@ -1353,3 +1356,33 @@ async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> N assert exc_info.value.status_code == 403 assert "MCP server 'github'" in exc_info.value.detail["error"] assert "agent 'agent-123'" in exc_info.value.detail["error"] + + +@pytest.mark.asyncio +async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: pytest.MonkeyPatch) -> None: + """The search lists the whole catalog but serves only its hits, so the listing must not fill the + caller's listed-tools slot: a later call to a tool the search never returned is not a listed tool.""" + monkeypatch.setattr(litellm, "mcp_tool_search", None) + manager = mcp_operations.global_mcp_server_manager + server = MCPServer(server_id="search-slot", name="search-slot", transport=MCPTransport.http, url="http://slot") + user = UserAPIKeyAuth(api_key="sk-search-slot", user_id="searcher") + upstream = [ + Tool(name="echo", description="Echo text back", inputSchema={"type": "object"}), + Tool(name="delete_note", description="Delete a note", inputSchema={"type": "object"}), + ] + with ( + patch.dict(manager.tool_name_to_mcp_server_name_mapping), + patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), + patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), + ): + try: + result = await handle_mcp_tool_search(query="echo", top_k=1, user_api_key_dict=user) + caller = ListedToolsCaller(user_api_key_auth=user) + listed = [manager.get_listed_tool(server, tool.name, caller) for tool in upstream] + finally: + manager._drop_listed_tools(server.server_id) + + assert result.is_error is False + assert [hit["name"] for hit in json.loads(result.content[0].text)] == ["search-slot-echo"] + assert listed == [None, None] diff --git a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index a8ec7be55f0..60157193a16 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -42,9 +42,12 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): fake_server.server_name = "openapi-petstore" fake_server.alias = None fake_server.short_prefix = None + fake_server.tool_name_to_description = None fake_tool = MagicMock() fake_tool.name = "list_pets" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} pre_call = AsyncMock(return_value={}) handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False)) @@ -125,9 +128,12 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): fake_server.server_name = "openapi-petstore" fake_server.alias = None fake_server.short_prefix = None + fake_server.tool_name_to_description = None fake_tool = MagicMock() fake_tool.name = "delete_pet" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} pre_call = AsyncMock( side_effect=HTTPException(status_code=403, detail="not allowed") @@ -190,6 +196,8 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable(): fake_tool = MagicMock() fake_tool.name = "list_pets" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} pre_call = AsyncMock(return_value={}) handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False)) @@ -274,6 +282,8 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): fake_tool = MagicMock() fake_tool.name = "get_values" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} captured: dict = {} async def handle_local(_name, _arguments, _wire_compat): @@ -620,6 +630,8 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc if dispatch_arm == "local_registry": fake_tool = MagicMock() fake_tool.name = "list_reports" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} with ( patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server), patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool), @@ -691,6 +703,8 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st fake_tool = MagicMock() fake_tool.name = "list_reports" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} fake_tool.handler = raising_handler server = MCPServer( server_id="srv-openapi", diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index b16b27ac919..f710bc7f1d7 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -1,15 +1,101 @@ import asyncio +from typing import Final from unittest.mock import AsyncMock, patch import pytest from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult +from mcp.types import Tool as MCPTool +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._experimental.mcp_server import operations +from litellm.proxy._experimental.mcp_server import rest_endpoints +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context +from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.utils import ProxyLogging from litellm.types.mcp import MCPAuth, MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer +class _CatalogHookCapture(CustomLogger): + data: dict[str, object] | None = None + + async def async_pre_call_hook( + self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, data: dict[str, object], call_type: str + ) -> None: + if call_type == "call_mcp_tool": + self.data = data.copy() + + +async def _served_catalog_tool() -> str: + return "ok" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("surface", ["mcp", "rest"]) +@pytest.mark.parametrize("restriction", ["key", "server"]) +async def test_listing_records_only_tools_the_caller_received( + monkeypatch: pytest.MonkeyPatch, surface: str, restriction: str +) -> None: + manager: Final = operations.global_mcp_server_manager + server: Final = MCPServer( + server_id="served-catalog", name="served-catalog", transport=MCPTransport.http, + spec_path="/catalog.yaml", allow_all_keys=True, + allowed_tools=["echo"] if restriction == "server" else None, + ) + auth: Final = UserAPIKeyAuth( + api_key="sk-served-catalog", user_id="lister", + object_permission={ + "object_permission_id": "served-permission", + "mcp_servers": [server.server_id], + "mcp_tool_permissions": {server.server_id: ["echo"]} if restriction == "key" else None, + }, + ) + monkeypatch.setitem(manager.registry, server.server_id, server) + monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "status", server.server_id) + monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "served-catalog-status", server.server_id) + capture: Final = _CatalogHookCapture() + monkeypatch.setattr(litellm, "callbacks", [capture]) + for name in ("echo", "status"): + global_mcp_tool_registry.register_tool( + name=f"served-catalog-{name}", description=f"{name} description", + input_schema={"type": "object"}, handler=_served_catalog_tool, + ) + try: + if surface == "mcp": + listing: Final = await operations._list_mcp_tools( + user_api_key_auth=auth, mcp_servers=[server.server_id], record_listing=True, + ) + assert [tool.name for tool in listing.tools] == ["served-catalog-echo"] + else: + rest_listing: Final = await rest_endpoints._get_tools_for_single_server( + server, None, user_api_key_auth=auth, + ) + assert [tool.name for tool in rest_listing] == ["echo"] + granted: Final = auth.model_copy(update={"object_permission": None}) + caller: Final = ListedToolsCaller(user_api_key_auth=granted) + assert manager.get_listed_tool(server, "status", caller) is None + served: Final = manager.get_listed_tool(server, "echo", caller) + assert served is not None + assert (served.description, served.input_schema) == ("echo description", {"type": "object"}) + server.allowed_tools = None + result: Final = await manager.call_tool( + server_name=server.server_id, name="status", arguments={}, user_api_key_auth=granted, + proxy_logging_obj=ProxyLogging(user_api_key_cache=UserApiKeyCache()), + ) + assert result.is_error is False + assert capture.data is not None + assert capture.data["messages"] == [{"role": "user", "content": "Tool: status\nArguments: {}"}] + assert (capture.data.get("mcp_tool_description"), capture.data.get("mcp_input_schema")) == (None, None) + finally: + manager._drop_listed_tools(server.server_id) + global_mcp_tool_registry.unregister_tools_with_prefix("served-catalog-") + + @pytest.mark.asyncio async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(caplog): from litellm.proxy._experimental.mcp_server.operations import _prefetch_oauth_creds_for_user @@ -665,3 +751,30 @@ async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled): ) assert result.tools == [] assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled + assert listing.await_args.kwargs["record_listing"] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("listing_kwargs", "recorded"), [({}, False), ({"record_listing": True}, True)]) +async def test_list_mcp_tools_records_the_catalog_only_when_asked( + listing_kwargs: dict[str, bool], recorded: bool +) -> None: + """The aggregate listing fills the caller's listed-tools slot only when asked: a listing an internal + caller never serves must not hand a later tools/call a description the caller never saw.""" + manager = operations.global_mcp_server_manager + server = MCPServer(server_id="listing-slot", name="listing-slot", transport=MCPTransport.http, url="http://slot") + user = UserAPIKeyAuth(api_key="sk-listing-slot", user_id="lister") + upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})] + with ( + patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), + patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), + patch.dict(manager.tool_name_to_mcp_server_name_mapping), + ): + try: + listing = await operations._list_mcp_tools(user_api_key_auth=user, **listing_kwargs) + listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) + finally: + manager._drop_listed_tools(server.server_id) + assert [tool.name for tool in listing.tools] == ["listing-slot-echo"] + assert (listed is not None) is recorded diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py index 23131321938..e72b716665c 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -346,6 +346,29 @@ class TestAllowFlow: assert evaluate_call.json["conversationId"] == "sess-123" assert evaluate_call.json["agentId"] == "my-agent-key" + @pytest.mark.asyncio + async def test_evaluate_payload_includes_listed_tool_metadata(self): + handler: Final = FakeHandler([_token_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler) + schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]} + await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema)) + assert handler.calls[1].json["tool"] == { + "name": "send_email", + "description": "Send an email", + "inputSchema": schema, + } + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("description", "schema"), + [(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")], + ) + async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema): + handler: Final = FakeHandler([_token_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler) + await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema)) + assert handler.calls[1].json["tool"] == {"name": "send_email"} + @pytest.mark.asyncio async def test_non_mcp_call_type_skipped(self): handler: Final = FakeHandler([]) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index 8b9439e25d2..d3a723fe1d3 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -13,6 +13,7 @@ from typing import Any, Dict, Final, List, Optional import pytest from fastapi import HTTPException +from pydantic import TypeAdapter import litellm from litellm import Router @@ -21,6 +22,7 @@ from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PARALLEL_REQUEST_SLOT_TTL_SECONDS, ParallelSlotAcquisition, @@ -39,6 +41,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.mcp import MCPPreCallRequestObject from litellm.types.utils import ( EmbeddingResponse, ModelResponse, @@ -108,6 +111,159 @@ def test_api_key_descriptor_applies_budget_throttle( assert api_key_descriptor["rate_limit"]["tokens_per_unit"] == expected_tpm +@pytest.mark.asyncio +@pytest.mark.parametrize( + "description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"] +) +@pytest.mark.parametrize("arguments_rewritten", [False, True]) +async def test_mcp_description_does_not_change_admission_or_reserved_tokens( + description: str | None, arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + cache: Final = DualCache() + handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()) + schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}} + request: Final = MCPPreCallRequestObject( + tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema + ) + data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {})) + messages: Final = data["messages"] + caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-description-reservation"), tpm_limit=64) + + if arguments_rewritten: + data["mcp_arguments"] = {"q": "Transformed arguments " * 100} + monkeypatch.setattr(litellm, "callbacks", [handler]) + await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool") + + stash: Final = get_request_stash() + assert stash is not None + assert stash.reserved_tokens == 25 + assert ( + await cache.async_get_cache( + key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True + ) + == 25 + ) + assert data["messages"] is messages + assert data.get("mcp_tool_description") == description + assert data["mcp_input_schema"] == schema + assert messages == [ + { + "role": "user", + "content": "Tool: echo\nArguments: {'q': 'hello'}", + } + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"] +) +@pytest.mark.parametrize("itpm_limit,otpm_limit", [(64, 4096), (4096, 64), (4096, 4096)]) +@pytest.mark.parametrize("arguments_rewritten", [False, True]) +async def test_mcp_description_preserves_project_input_and_output_reservations( + description: str | None, itpm_limit: int, otpm_limit: int, + arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + cache: Final = DualCache() + handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()) + schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}} + request: Final = MCPPreCallRequestObject( + tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema + ) + data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {})) + messages: Final = data["messages"] + base_data: Final[dict[str, object]] = { + "messages": [{"role": "user", "content": "Tool: echo\nArguments: {'q': 'hello'}"}] + } + expected_input: Final = handler._estimate_precise_input_tokens(base_data, "mcp-tool-call", "call_mcp_tool") + expected_output: Final = handler.no_max_tokens_output_floor(otpm_limit) + expected_combined: Final = handler._estimate_tokens_for_request( + base_data, min_configured_tpm_limit=4096, call_type="call_mcp_tool" + ) + caller: Final = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-project-reservation"), + tpm_limit=4096, + project_id="mcp-project-reservation", + project_metadata={ + "model_itpm_limit": {"mcp-tool-call": itpm_limit}, + "model_otpm_limit": {"mcp-tool-call": otpm_limit}, + }, + ) + + if arguments_rewritten: + data["mcp_arguments"] = {"q": "Transformed arguments " * 100} + monkeypatch.setattr(litellm, "callbacks", [handler]) + await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool") + + stash: Final = get_request_stash() + assert stash is not None + assert (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) == ( + expected_combined, + expected_input, + expected_output, + ) + assert ( + await cache.async_get_cache( + key=handler.create_rate_limit_keys( + "model_per_project_itpm", f"{caller.project_id}:mcp-tool-call", "tokens" + ), + local_only=True, + ) + == expected_input + ) + assert ( + await cache.async_get_cache( + key=handler.create_rate_limit_keys( + "model_per_project_otpm", f"{caller.project_id}:mcp-tool-call", "tokens" + ), + local_only=True, + ) + == expected_output + ) + assert data["messages"] is messages + assert data.get("mcp_tool_description") == description + assert data["mcp_input_schema"] == schema + assert messages == [ + { + "role": "user", + "content": "Tool: echo\nArguments: {'q': 'hello'}", + } + ] + + +def test_llm_tpm_estimation_still_counts_messages_with_mcp_metadata() -> None: + handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + data: Final[dict[str, object]] = { + "messages": [{"role": "user", "content": "x" * 400}], + "max_tokens": 1, + "mcp_tool_name": "echo", + "mcp_arguments": {}, + } + assert handler._estimate_tokens_for_request(data, call_type="acompletion") == 101 + + +@pytest.mark.asyncio +async def test_unconverted_mcp_request_keeps_its_reservation() -> None: + cache: Final = DualCache() + handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-raw-mcp-request"), tpm_limit=64) + data: Final[dict[str, object]] = {"name": "echo", "arguments": {"q": "hello"}, "server_id": "fixture"} + + await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool") + + stash: Final = get_request_stash() + assert stash is not None + assert stash.reserved_tokens == 16 + assert ( + await cache.async_get_cache( + key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True + ) + == 16 + ) + + @pytest.mark.flaky(reruns=3) @pytest.mark.asyncio async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller): diff --git a/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py index c5bd89645b4..25be3b5de6b 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py +++ b/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py @@ -403,6 +403,26 @@ def test_create_mcp_request_object_from_kwargs_full(proxy_logging, make_user_api assert snapshot == {"tool_name": "calc", "arguments": {"x": 1}, "server_name": "math", "auth_user_id": "u-1"} +def test_mcp_tool_metadata_flows_from_kwargs_to_synthetic_data(proxy_logging): + schema = {"type": "object", "properties": {"x": {"type": "integer"}}} + obj = proxy_logging._create_mcp_request_object_from_kwargs( + kwargs={ + "name": "calc", + "arguments": {"x": 1}, + "tool_description": "Adds numbers", + "tool_input_schema": schema, + } + ) + out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={}) + assert (out["mcp_tool_description"], out["mcp_input_schema"]) == ("Adds numbers", schema) + + +def test_mcp_tool_metadata_absent_when_tool_was_never_listed(proxy_logging): + obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={"name": "calc", "arguments": {}}) + out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={}) + assert "mcp_tool_description" not in out and "mcp_input_schema" not in out + + def test_create_mcp_request_object_from_kwargs_empty(proxy_logging): obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={}) snapshot = { diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 2699f9445c9..73bee304fc2 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1,10 +1,11 @@ +import asyncio import importlib import subprocess import sys import textwrap import types -from typing import Any, Final, cast -from unittest.mock import AsyncMock, MagicMock +from typing import Any, Final, Literal, cast +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException @@ -12,15 +13,26 @@ from mcp.types import CallToolResult, TextContent from mcp.types import Tool as MCPTool from openai.types.responses.tool_param import Mcp +import litellm +from litellm.caching.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.proxy._experimental.mcp_server import operations as mcp_operations from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import ProxyLogging from litellm.responses import main as responses_main from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.mcp import MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.responses.main import OutputFunctionToolCall -from litellm.types.utils import ModelResponse +from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse class _DummyMCPResult: @@ -1310,6 +1322,167 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel" +@pytest.mark.asyncio +@pytest.mark.parametrize("real_listing", [False, True]) +@pytest.mark.parametrize( + ("allowed_tools", "expected_names"), + [ + ([], ["responses_slot-echo", "responses_slot-status"]), + (["echo"], ["responses_slot-echo"]), + (["responses_slot-echo"], ["responses_slot-echo"]), + (["absent"], []), + ], +) +async def test_bridge_listing_leaves_the_callers_catalog_unchanged( + monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str], real_listing: bool +) -> None: + manager: Final = mcp_operations.global_mcp_server_manager + server: Final = MCPServer( + server_id="responses-slot", name="responses_slot", alias="responses_slot", transport=MCPTransport.http + ) + user: Final = UserAPIKeyAuth(api_key="sk-responses-slot", user_id="responder") + upstream: Final = [ + MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"}), + MCPTool(name="status", description="Report status", inputSchema={"type": "object"}), + MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}), + ] + fake_manager: Final = types.SimpleNamespace( + get_registry=MagicMock(return_value={}), + get_allowed_mcp_servers=AsyncMock(return_value=[]), + get_mcp_servers_from_ids=MagicMock(return_value=[]), + get_mcp_server_by_name=MagicMock(return_value=None), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + with ( + patch.dict(manager.tool_name_to_mcp_server_name_mapping), + patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), + patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), + ): + try: + if real_listing: + await manager._get_tools_from_server(server, user_api_key_auth=user, record_listing=True) + caller: Final = ListedToolsCaller(user_api_key_auth=user) + before: Final = { + tool.name: (listed.description, listed.input_schema) + for tool in upstream + if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None + } + assert bool(before) is real_listing + tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + user_api_key_auth=user, + mcp_tools_with_litellm_proxy=[ + { + "type": "mcp", + "server_url": "litellm_proxy/mcp/responses-slot", + "allowed_tools": allowed_tools, + } + ], + ) + recorded: Final = { + tool.name: (listed.description, listed.input_schema) + for tool in upstream + if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None + } + assert recorded == before + assert ( + manager.get_listed_tool( + server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller")) + ) + is None + ) + finally: + manager._drop_listed_tools(server.server_id) + + assert [tool.name for tool in tools] == expected_names + + +class _BridgeMetadataGuardrail(CustomGuardrail): + def __init__(self) -> None: + super().__init__(guardrail_name="bridge-metadata", event_hook=GuardrailEventHooks.pre_mcp_call, default_on=True) + self.calls: tuple[tuple[object, object], ...] = () + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + logging_obj: Logging | None = None, + ) -> GenericGuardrailAPIInputs: + if request_data.get("mcp_arguments") == {"probe": "bridge"}: + self.calls += ((request_data.get("mcp_tool_description"), request_data.get("mcp_input_schema")),) + return inputs + + +@pytest.mark.asyncio +async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="bridge", name="bridge", transport=MCPTransport.http, url="http://upstream") + manager.registry = {server.server_id: server} + user: Final = UserAPIKeyAuth(api_key="sk-bridge", user_id="bridge-user") + upstream: Final = [ + MCPTool( + name="echo", + description="Echo text", + inputSchema={"type": "object", "properties": {"text": {"type": "string"}}}, + ), + MCPTool(name="status", description="Read status", inputSchema={"type": "object"}), + ] + client: Final = AsyncMock() + client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")]) + manager._create_mcp_client = AsyncMock(return_value=client) + manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream) + guardrail: Final = _BridgeMetadataGuardrail() + logger: Final = ProxyLogging(user_api_key_cache=DualCache()) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", logger) + monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager) + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])) + monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager) + first_listed: Final = asyncio.Event() + second_listed: Final = asyncio.Event() + + async def bridge(name: str, first: bool) -> None: + if not first: + await first_listed.wait() + tools, server_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + user_api_key_auth=user, + mcp_tools_with_litellm_proxy=[ + {"type": "mcp", "server_url": "litellm_proxy/mcp/bridge", "allowed_tools": [name]} + ], + ) + (first_listed if first else second_listed).set() + await second_listed.wait() + result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map=server_map, + tool_calls=[ + {"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name} + ], + user_api_key_auth=user, + served_tools=tools, + ) + assert [entry["result"] for entry in result] == ["ok"] + + try: + await asyncio.gather(bridge("echo", True), bridge("status", False)) + assert sorted(guardrail.calls, key=str) == sorted( + ((tool.description, tool.input_schema) for tool in upstream), key=str + ) + await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger) + assert guardrail.calls[-1] == (None, None) + await manager._get_tools_from_server( + server, user_api_key_auth=user, proxy_logging_obj=logger, record_listing=True + ) + await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger) + assert guardrail.calls[-1] == (upstream[0].description, upstream[0].input_schema) + finally: + manager._drop_listed_tools(server.server_id) + ProxyLogging._callback_capabilities_cache.clear() + + def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace: return types.SimpleNamespace( get_registry=MagicMock(return_value={}),