From 3c13e3b457ffa555d504056094452a60b17bcc33 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 15 Sep 2026 00:42:16 +0000 Subject: [PATCH 01/47] feat(mcp): hand listed-tool metadata to pre-call hooks with per-caller catalog identity Track the tools each MCP server listed per caller identity so pre_mcp_call and during_mcp_call hooks receive the tool description and input schema the client saw. Servers with no caller-dependent inputs share one slot; user identity, forwarded headers, stdio env, relayed bearers, and server-specific auth get their own. Local registry and OpenAPI paths pass the registered metadata and admin description overrides. The Agent 365 guardrail reads the new fields into its evaluate payload. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 231 ++++++--- .../_experimental/mcp_server/operations.py | 12 +- .../guardrail_hooks/agent_365/agent_365.py | 26 +- litellm/proxy/utils.py | 4 + litellm/types/mcp.py | 4 + .../test_mcp_guardrail_usage_monitor.py | 1 + .../mcp_server/test_mcp_server.py | 140 +++++- .../mcp_server/test_mcp_server_manager.py | 460 ++++++++++++++++++ .../mcp_server/test_openapi_tool_auth.py | 14 + .../guardrail_hooks/test_agent_365.py | 23 + .../utils/proxy_logging/test_mcp_bridging.py | 20 + 11 files changed, 876 insertions(+), 59 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 6baa695433c..81f1d677d87 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -250,6 +250,21 @@ _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]] +_NO_LISTED_TOOLS: Final[_ListedToolsByCaller] = MappingProxyType({}) +_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 @@ -1128,6 +1143,25 @@ 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 _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str: """Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection. @@ -1931,6 +1965,7 @@ class MCPServerManager: "gmail_send_email": "zapier_mcp_server", } """ + self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list 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 @@ -2629,7 +2664,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) _warn_legacy_delegate_auth_if_applicable(new_server, source="config") _warn_config_id_jag_server_outruns_sso(new_server) - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2833,7 +2868,7 @@ class MCPServerManager: global_mcp_tool_registry, ) - self._invalidate_discovery_lists(server.server_id) + self._invalidate_server_definition_caches(server.server_id) prefix_root: Final = normalize_server_name(get_server_prefix(server)) if server.spec_path and prefix_root: openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR @@ -3240,7 +3275,7 @@ class MCPServerManager: # env_vars_are_encrypted=False. new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3277,7 +3312,7 @@ class MCPServerManager: previous_server=self.registry[mcp_server.server_id], ) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3732,19 +3767,7 @@ 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( @@ -3830,7 +3853,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.""" @@ -4379,6 +4402,12 @@ 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=mcp_auth_header, + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ) try: # Tool *listing* must not be blocked by missing per-user env vars — @@ -4457,29 +4486,25 @@ class MCPServerManager: if server.spec_path: # OpenAPI tools were stored in the registry under the prefix # active at registration time — fetch by that same prefix. - _tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server)) + registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR + _tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix) tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(_tools) # OpenAPI tools are stored in the registry with their prefix already # 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". - if not add_prefix: - prefix: Final = get_server_prefix(server) - sep: Final = MCP_TOOL_PREFIX_SEPARATOR - tools = [ - ( - t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]}) - if t.name.startswith(f"{prefix}{sep}") - else t - ) - for t in tools - ] - return tools + unprefixed_tools: Final = [ # mutable-ok: returned through the list[MCPTool] listing contract + t.model_copy(update=MappingProxyType({"name": t.name[len(registry_prefix) :]})) for t in tools + ] + self._record_listed_tools(server, unprefixed_tools, listed_caller) + return tools if add_prefix else unprefixed_tools else: tools = await self._fetch_tools_with_timeout(client, server.name) self._remember_upstream_initialize_instructions(server, client) - prefixed_or_original_tools: Final = self._create_prefixed_tools(tools, server, add_prefix=add_prefix) + prefixed_or_original_tools: Final = self._create_prefixed_tools( + tools, server, add_prefix=add_prefix, caller=listed_caller + ) return prefixed_or_original_tools @@ -4523,6 +4548,86 @@ class MCPServerManager: self._resource_discovery_cache.invalidate(server_id) self._template_discovery_cache.invalidate(server_id) + def _invalidate_server_definition_caches(self, server_id: str) -> None: + self._invalidate_discovery_lists(server_id) + self._listed_tools_by_server_id.pop(server_id, None) + + def _discovers_per_caller(self, server: MCPServer) -> bool: + return ( + 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) + or self._signs_caller_identity_upstream(server) + ) + + @staticmethod + def _signs_caller_identity_upstream(server: MCPServer) -> bool: + """Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may + tailor its catalog to the caller even though the server itself is configured as shared.""" + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server + get_mcp_jwt_signer, + ) + + if get_mcp_jwt_signer() is None: + return False + return not any(k.lower() == "authorization" for k in (server.static_headers or {})) + + 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 upstream catalog. + + Forwarded headers, header-driven stdio env, a relayed caller bearer, and the + server-specific auth header all reach upstream, so two callers differing in any of + them may be shown different tools. Shared servers with none of those stay on the + shared (``None``) slot. OpenAPI servers list from the process-wide registry. + """ + if server.spec_path or caller is None: + return None + auth: Final = caller.user_api_key_auth + identity: Final = ( + (auth.user_id, auth.api_key) if auth is not None and self._discovers_per_caller(server) else None + ) + forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) + 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 + relayed_bearer: Final = ( + self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) + if server.is_client_forwarded_token + else None + ) + inputs: Final = (identity, caller.mcp_auth_header, forwarded, stdio_env, relayed_bearer) + if not any(inputs): + return None + material: Final = json.dumps(inputs, 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 _record_listed_tools( + self, server: MCPServer, tools: Sequence[MCPTool], caller: ListedToolsCaller | None + ) -> None: + 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, _NO_LISTED_TOOLS) + 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, @@ -4533,12 +4638,7 @@ class MCPServerManager: subject_token: str | None, credential_fingerprint: str | None = None, ) -> _DiscoveryKey: - per_user: Final = ( - server.requires_per_user_auth - or self._references_per_user_env_var(server) - or server.delegate_auth_to_upstream - or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) - ) + per_user: Final = self._discovers_per_caller(server) if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): return server.server_id, None identity: Final = ( @@ -5368,7 +5468,13 @@ class MCPServerManager: "attempts; the 3-character prefix space is too crowded." ) - def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]: + def _create_prefixed_tools( + self, + tools: list[MCPTool], + server: MCPServer, + add_prefix: bool = True, + caller: ListedToolsCaller | None = None, + ) -> list[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -5393,9 +5499,21 @@ class MCPServerManager: for spelling in iter_known_tool_name_spellings(original_name, server): self.tool_name_to_mcp_server_name_mapping[spelling] = prefix + self._record_listed_tools(server, tools, caller) 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, _NO_LISTED_TOOLS).get(identity) + if not listed: + return None + tool: Final = listed.get(name) or listed.get(strip_known_server_prefix(name, server)) + if tool is None: + return None + description: Final = (server.tool_name_to_description or {}).get(tool.name) + return tool if description is None else tool.model_copy(update={"description": description}) + def _create_prefixed_prompts( self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True ) -> list[Prompt]: @@ -5629,6 +5747,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. @@ -5642,6 +5761,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 @@ -5696,6 +5818,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 @@ -5751,6 +5875,7 @@ class MCPServerManager: start_time: datetime.datetime, litellm_logging_obj: "LiteLLMLoggingObj | None" = None, guardrail_context: Mapping[str, object] | None = None, + tool: MCPTool | None = None, ): """Create and return a during hook task for MCP tool calls. @@ -5765,6 +5890,8 @@ class MCPServerManager: tool_name=name, arguments=arguments, server_name=server_name_from_prefix, + tool_description=tool.description if tool is not None else None, + tool_input_schema=tool.input_schema if tool is not None else None, start_time=start_time.timestamp() if start_time else None, hidden_params=HiddenParams(), ) @@ -5911,21 +6038,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 @@ -6376,6 +6489,12 @@ class MCPServerManager: user_api_key_auth, mcp_auth_header, ) + listed_caller: Final = ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=_server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header), + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ) ######################################################### # Pre MCP Tool Call Hook @@ -6392,6 +6511,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 "arguments" in hook_result: arguments = hook_result["arguments"] @@ -6408,6 +6528,7 @@ class MCPServerManager: start_time=start_time, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + tool=self.get_listed_tool(mcp_server, name, listed_caller), ) tasks.append(during_hook_task) @@ -6669,7 +6790,7 @@ class MCPServerManager: for server_id in previous_registry.keys() | registered_registry.keys(): if previous_registry.get(server_id) != registered_registry.get(server_id): - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.registry = registered_registry _warn_on_shared_identifier_prefixes(registered_registry.values()) # A discovery task may have published into ``previous_registry`` while diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index a19246b6e90..1395b808c2f 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -139,6 +139,7 @@ from litellm.types.mcp import ( without_header, ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer +from litellm.types.mcp_server.tool_registry import MCPTool as RegisteredTool from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall from litellm.utils import Rules, client, function_setup @@ -1610,6 +1611,12 @@ async def _list_mcp_resource_templates( return managed_resource_templates +def _registered_tool_metadata(name: str, registered: RegisteredTool, server: MCPServer) -> MCPTool: + overrides: Final = server.tool_name_to_description + description: Final = overrides.get(name, registered.description) if overrides else registered.description + return MCPTool(name=name, description=description, input_schema=registered.input_schema) + + def _resolve_display_name_to_original( name: str, allowed_mcp_servers: list[MCPServer], @@ -2079,6 +2086,7 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + tool=_registered_tool_metadata(original_tool_name, local_tool, mcp_server), ) # `pre_call_tool_check` may return guardrail-modified # arguments; honor them on the local path too. @@ -2147,7 +2155,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 @@ -2189,6 +2198,7 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + tool=_registered_tool_metadata(original_tool_name, registered_local_tool, prefix_server), ) if "arguments" in hook_result: arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args 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 975d321104d..151d0bea231 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 @@ -60,6 +60,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 @@ -81,6 +82,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``.""" @@ -99,6 +107,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] @@ -397,8 +413,14 @@ class Agent365Guardrail(CustomGuardrail): arguments: Final = data.get("mcp_arguments") server_name: Final = str(data.get("mcp_server_name") or "litellm") agent_id: Final = self.agent_id or 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_tool_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/utils.py b/litellm/proxy/utils.py index b8cc30ad8a7..dc4ea58d8d6 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1498,6 +1498,8 @@ class ProxyLogging: "user_api_key_request_route": kwargs.get("user_api_key_request_route"), "mcp_tool_name": request_obj.tool_name, # Keep original for reference "mcp_arguments": request_obj.arguments, # Keep original for reference + "mcp_tool_description": request_obj.tool_description, + "mcp_tool_input_schema": request_obj.tool_input_schema, # Surface the per-MCP-server rate-limit identity so the # ParallelRequestLimiterV3 hook can apply mcp_rpm_limit on the # synthetic call_mcp_tool payload (otherwise a key with @@ -1728,6 +1730,8 @@ class ProxyLogging: tool_name=kwargs.get("name", ""), arguments=kwargs.get("arguments", {}), server_name=kwargs.get("server_name"), + tool_description=kwargs.get("tool_description"), + tool_input_schema=kwargs.get("tool_input_schema"), user_api_key_auth=user_api_key_auth_dict, hidden_params=HiddenParams(), ) diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index fec5e84c8df..f654bff08fb 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -422,6 +422,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() @@ -445,6 +447,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/test_litellm/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py index 24e6d2de10d..956f87f0da1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py +++ b/tests/test_litellm/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/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ab00ec4da1e..48a1da8605b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.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 @@ -6881,7 +6882,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 @@ -7993,9 +7993,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( @@ -8046,6 +8049,141 @@ 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_registered_tool_metadata_to_pre_call_hooks(): + """OpenAPI-generated tools dispatch through the local registry, so the pre-call hooks must get the + registered description and input schema on that path too, even when no tools/list ran first.""" + 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": {"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) + 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-list_pets", + arguments={"limit": 10}, + allowed_mcp_servers=[petstore], + start_time=datetime.now(), + user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"), + ) + finally: + mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + + handed_tool = pre_call_tool_check.call_args.kwargs["tool"] + assert (handed_tool.name, handed_tool.description, handed_tool.input_schema) == ( + "list_pets", + "List the pets", + schema, + ) + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_hands_openapi_hooks_the_admin_description_clients_saw(): + """tools/list shows the admin's tool_name_to_description wording, so the local-registry call path + must hand the pre-call hooks that same wording rather than the generated 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": "ADMIN DESC"}, + ) + 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=schema, handler=lambda petId: "ok" + ) + manager = mcp_module.global_mcp_server_manager + manager._listed_tools_by_server_id.pop(petstore.server_id, None) + 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=UserAPIKeyAuth(api_key="sk-user", user_id="alice"), + ) + finally: + mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + + handed_tool = pre_call_tool_check.call_args.kwargs["tool"] + assert (handed_tool.description, handed_tool.input_schema) == ("ADMIN DESC", schema) + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_hands_hooks_the_metadata_of_the_operation_it_runs_when_names_collide(): + """An OpenAPI operation whose name starts with its own server prefix must not be reported to the + pre-call hooks with the metadata of the shorter operation, since that is not the one that runs.""" + 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-") + + handed_tool = pre_call_tool_check.call_args.kwargs["tool"] + assert (handed_tool.description, handed_tool.input_schema) == ( + "long", + {"type": "object", "properties": {"petId": {"type": "integer"}}}, + ) + assert result.content[0].text == "long" + + @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. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 16bffa1a356..246db7ae7f1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -41,6 +41,7 @@ from pydantic import AnyUrl, TypeAdapter from litellm.constants import MCP_METADATA_TIMEOUT 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, @@ -6772,6 +6773,465 @@ 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]) -> 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) + + 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() -> MagicMock: + user_api_key_auth = MagicMock() + user_api_key_auth.object_permission = None + user_api_key_auth.object_permission_id = None + return user_api_key_auth + + @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)] + manager, proxy_logging_obj = self._manager_ready_for_call_tool(listed) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=self._unrestricted_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_listed_tool_metadata_to_during_call_hooks_through_real_conversion(self): + schema = {"type": "object", "properties": {"param": {"type": "string"}}} + listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] + manager, _ = self._manager_ready_for_call_tool(listed) + 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=UserAPIKeyAuth(api_key="sk-test"), + proxy_logging_obj=proxy_logging_obj, + ) + + during_data = proxy_logging_obj.during_call_hook.call_args.kwargs["data"] + assert (during_data["mcp_tool_description"], during_data["mcp_tool_input_schema"]) == ( + "Runs the test tool", + schema, + ) + + @pytest.mark.asyncio + async def test_call_tool_passes_no_tool_metadata_when_tool_was_never_listed(self): + manager, proxy_logging_obj = self._manager_ready_for_call_tool( + [MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})] + ) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=self._unrestricted_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_prefixed_name_and_latest_listing(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._create_prefixed_tools([MCPTool(name="echo", description="v1", inputSchema={})], server) + manager._create_prefixed_tools([MCPTool(name="echo", description="v2", inputSchema={})], server) + + by_prefixed_name = manager.get_listed_tool(server, "srv-echo") + assert by_prefixed_name is not None and by_prefixed_name.description == "v2" + assert manager.get_listed_tool(server, "missing") is None + + def test_get_listed_tool_uses_admin_description_override_clients_saw(self): + manager = MCPServerManager() + server = MCPServer( + server_id="srv", + name="srv", + transport=MCPTransport.http, + url="http://srv", + tool_name_to_description={"echo": "Admin wording"}, + ) + schema = {"type": "object", "properties": {"text": {"type": "string"}}} + manager._create_prefixed_tools( + [ + MCPTool(name="echo", description="Upstream wording", inputSchema=schema), + MCPTool(name="ping", description="Untouched", inputSchema={}), + ], + server, + ) + + overridden = manager.get_listed_tool(server, "srv-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" + + 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._create_prefixed_tools([MCPTool(name="echo", description="old", inputSchema={})], server) + manager._create_prefixed_tools([MCPTool(name="ping", description="kept", inputSchema={})], other) + + 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_user_oauth_refresh_keeps_listed_tools(self): + """Tool definitions are server-wide, so one user's re-auth must not blank the metadata other + callers' tool calls hand to pre-call guardrails.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._create_prefixed_tools([MCPTool(name="echo", description="shared", inputSchema={})], server) + + 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", api_key="hashed-alice") + bob = UserAPIKeyAuth(user_id="bob", api_key="hashed-bob") + alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}} + bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}} + manager._create_prefixed_tools( + [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], + server, + caller=ListedToolsCaller(user_api_key_auth=alice), + ) + manager._create_prefixed_tools( + [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], + server, + caller=ListedToolsCaller(user_api_key_auth=bob), + ) + + alice_tool = manager.get_listed_tool(server, "srv-read", ListedToolsCaller(user_api_key_auth=alice)) + bob_tool = manager.get_listed_tool(server, "srv-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", api_key="k")) + assert manager.get_listed_tool(server, "srv-read", carol) is None + + shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") + manager._create_prefixed_tools( + [MCPTool(name="echo", description="everyone", inputSchema={})], + shared, + caller=ListedToolsCaller(user_api_key_auth=alice), + ) + for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob)) + assert for_bob is not None and for_bob.description == "everyone" + + @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( + {"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): + """Whatever reaches upstream and can change its catalog must also split the listed-tool cache.""" + manager = MCPServerManager() + server = MCPServer( + **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} + ) + manager._create_prefixed_tools( + [MCPTool(name="turn", description="Catalog A", inputSchema={})], server, caller=caller_a + ) + manager._create_prefixed_tools( + [MCPTool(name="turn", description="Catalog B", inputSchema={})], server, caller=caller_b + ) + + for_a = manager.get_listed_tool(server, "srv-turn", caller_a) + for_b = manager.get_listed_tool(server, "srv-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, "srv-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._create_prefixed_tools( + [MCPTool(name="turn", description="everyone", inputSchema={})], + server, + caller=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.parametrize( + ("signer", "static_headers", "shared"), + [ + pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), + pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), + pytest.param(None, None, True, id="no-signer-stays-shared"), + ], + ) + def test_jwt_signer_makes_a_shared_server_list_per_caller(self, signer, static_headers, shared): + """MCPJWTSigner hands upstream a JWT naming the caller on an otherwise shared ``auth_type: none`` + server, so the upstream may tailor the catalog and the cache must not hand one caller another's.""" + 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", api_key="hashed-alice")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="bob", api_key="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._create_prefixed_tools( + [MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice + ) + for_bob = manager.get_listed_tool(server, "srv-turn", bob) + + if shared: + assert for_bob is not None and for_bob.description == "alice view" + else: + assert for_bob is None + + @pytest.mark.asyncio + async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self): + """Interleaved callers on a forwarded-header server: the hook must see the caller's own catalog.""" + 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"), + ) + + 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="catalog-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._create_prefixed_tools([MCPTool(name="read", description="shared", inputSchema={})], server) + 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._create_prefixed_tools( + [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], + server, + caller=caller, + ) + manager._create_prefixed_tools( + [MCPTool(name="read", description="u1 again", inputSchema={})], server, caller=callers[1] + ) + + assert manager.get_listed_tool(server, "srv-read", callers[0]) is None + second = manager.get_listed_tool(server, "srv-read", callers[1]) + assert second is not None and second.description == "u1 again" + newest = manager.get_listed_tool(server, "srv-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, "srv-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) + 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"] + for name in ("list_pets", "petstore-list_pets"): + tool = manager.get_listed_tool(server, name) + 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) + 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 async def test_get_allowed_mcp_servers_with_user_api_key_auth(self): """ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index 15d3b67e641..972496218cf 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/test_litellm/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/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index f9b7561b9d3..de2d3a95853 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -300,6 +300,29 @@ class TestAllowFlow: assert evaluate_call.json["conversationId"] == "sess-123" assert evaluate_call.json["agentId"] == "agent-007" + @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_tool_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_tool_input_schema=schema)) + assert handler.calls[1].json["tool"] == {"name": "send_email"} + @pytest.mark.asyncio async def test_agent_id_falls_back_to_key_alias(self): handler: Final = FakeHandler([_token_response(), _allow_response()]) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py index 4e02124e1b3..c078b6878e5 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py +++ b/tests/test_litellm/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_tool_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 (out["mcp_tool_description"], out["mcp_tool_input_schema"]) == (None, None) + + def test_create_mcp_request_object_from_kwargs_empty(proxy_logging): obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={}) snapshot = { From 504870e1f9237b85c092604ea1c07a2b8d270a41 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 15 Sep 2026 01:40:32 +0000 Subject: [PATCH 02/47] refactor(mcp): drop the listed-tools empty sentinel and routine test docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/mcp_server_manager.py | 5 ++--- .../_experimental/mcp_server/test_mcp_server_manager.py | 6 ------ 2 files changed, 2 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 81f1d677d87..1a90435203a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -251,7 +251,6 @@ _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]] -_NO_LISTED_TOOLS: Final[_ListedToolsByCaller] = MappingProxyType({}) _LISTED_TOOLS_CALLERS_PER_SERVER: Final = 256 @@ -4617,7 +4616,7 @@ class MCPServerManager: ) -> None: 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, _NO_LISTED_TOOLS) + existing: Final[_ListedToolsByCaller] = self._listed_tools_by_server_id.get(server.server_id, {}) 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) @@ -5505,7 +5504,7 @@ class MCPServerManager: 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, _NO_LISTED_TOOLS).get(identity) + listed: Final = self._listed_tools_by_server_id.get(server.server_id, {}).get(identity) if not listed: return None tool: Final = listed.get(name) or listed.get(strip_known_server_prefix(name, server)) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 246db7ae7f1..e40cc943b45 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6913,8 +6913,6 @@ class TestMCPServerManager: @pytest.mark.asyncio async def test_user_oauth_refresh_keeps_listed_tools(self): - """Tool definitions are server-wide, so one user's re-auth must not blank the metadata other - callers' tool calls hand to pre-call guardrails.""" manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") manager._create_prefixed_tools([MCPTool(name="echo", description="shared", inputSchema={})], server) @@ -6997,7 +6995,6 @@ class TestMCPServerManager: ], ) def test_upstream_identity_inputs_keep_listed_tools_apart(self, server_kwargs, caller_a, caller_b): - """Whatever reaches upstream and can change its catalog must also split the listed-tool cache.""" manager = MCPServerManager() server = MCPServer( **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} @@ -7037,8 +7034,6 @@ class TestMCPServerManager: ], ) def test_jwt_signer_makes_a_shared_server_list_per_caller(self, signer, static_headers, shared): - """MCPJWTSigner hands upstream a JWT naming the caller on an otherwise shared ``auth_type: none`` - server, so the upstream may tailor the catalog and the cache must not hand one caller another's.""" manager = MCPServerManager() server = MCPServer( server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers @@ -7062,7 +7057,6 @@ class TestMCPServerManager: @pytest.mark.asyncio async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self): - """Interleaved callers on a forwarded-header server: the hook must see the caller's own catalog.""" manager = MCPServerManager() server = MCPServer( server_id="catalog", From 8cfbf5dc2eb55bc344ceb5adbea2adfba120b93e Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 15 Sep 2026 01:52:26 +0000 Subject: [PATCH 03/47] fix(mcp): mark the listed-tools cache digest as a non-security hash Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/mcp_server_manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1a90435203a..9de3179d053 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4598,7 +4598,7 @@ class MCPServerManager: if not any(inputs): return None material: Final = json.dumps(inputs, sort_keys=True, separators=(",", ":")) - return hashlib.sha256(material.encode()).hexdigest() + return hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest() @staticmethod def _forwarded_header_values( From 0ab0e49e6a9d6f5e53ce4828e6a91ecca78b67c8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 15 Sep 2026 02:09:05 +0000 Subject: [PATCH 04/47] fix(mcp): key the listed-tools cache by the OBO subject token token_exchange servers list upstream with the caller's own Entra bearer, so two callers on one LiteLLM key with different subjects were sharing a catalog slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 15 ++++++++------- .../mcp_server/test_mcp_server_manager.py | 12 ++++++++++++ 2 files changed, 20 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 9de3179d053..6de4be72e23 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4575,10 +4575,11 @@ class MCPServerManager: 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 upstream catalog. - Forwarded headers, header-driven stdio env, a relayed caller bearer, and the - server-specific auth header all reach upstream, so two callers differing in any of - them may be shown different tools. Shared servers with none of those stay on the - shared (``None``) slot. OpenAPI servers list from the process-wide registry. + Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or + exchanged as the OBO subject), and the server-specific auth header all reach + upstream, so two callers differing in any of them may be shown different tools. Shared + servers with none of those stay on the shared (``None``) slot. OpenAPI servers list from + the process-wide registry. """ if server.spec_path or caller is None: return None @@ -4589,12 +4590,12 @@ class MCPServerManager: forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) 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 - relayed_bearer: Final = ( + caller_bearer: Final = ( self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) - if server.is_client_forwarded_token + if server.is_client_forwarded_token or server.auth_type == MCPAuth.oauth2_token_exchange else None ) - inputs: Final = (identity, caller.mcp_auth_header, forwarded, stdio_env, relayed_bearer) + inputs: Final = (identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer) if not any(inputs): return None material: Final = json.dumps(inputs, sort_keys=True, separators=(",", ":")) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index e40cc943b45..c016a63acd6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6986,6 +6986,18 @@ class TestMCPServerManager: 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", api_key="hashed-shared"), + raw_headers={"x-litellm-api-key": "sk-shared", "authorization": "Bearer entra-alice"}, + ), + ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", api_key="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"}), From bef557f3ed971cf338941cf75cac51ea546bd389 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 08:34:38 +0000 Subject: [PATCH 05/47] fix(mcp): resolve the BYOK credential before keying the listed-tools slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 13 ++++++-- .../mcp_server/test_mcp_server_manager.py | 31 +++++++++++++++++++ 2 files changed, 41 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 6de4be72e23..b1f52bd5bcc 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4401,9 +4401,16 @@ class MCPServerManager: verbose_logger.info("_get_tools_from_server for %s...", server.name) client = None + # tools/call resolves the BYOK credential before keying its listed-tools slot; resolve the + # same value here or a stored-credential server would list into a slot the call never reads. + resolved_mcp_auth_header: Final = ( + mcp_auth_header + if not server.is_byok or isinstance(mcp_auth_header, dict) + else await _resolve_byok_mcp_auth_header(server, user_api_key_auth, mcp_auth_header) + ) listed_caller: Final = ListedToolsCaller( user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, + mcp_auth_header=resolved_mcp_auth_header, raw_headers=raw_headers, oauth2_headers=oauth2_headers, ) @@ -4449,7 +4456,7 @@ class MCPServerManager: if ( get_mcp_jwt_signer() is not None and not has_static_authorization - and not mcp_auth_header + and not resolved_mcp_auth_header and not has_extra_authorization ): extra_headers = await inject_mcp_jwt_headers_for_upstream( @@ -4472,7 +4479,7 @@ class MCPServerManager: client = await self._create_mcp_client( server=server, - mcp_auth_header=mcp_auth_header, + mcp_auth_header=resolved_mcp_auth_header, extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index c016a63acd6..368f0d09e3c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7037,6 +7037,37 @@ class TestMCPServerManager: listed = manager.get_listed_tool(server, "turn", other) assert listed is not None and listed.description == "everyone" + @pytest.mark.asyncio + async def test_byok_stored_credential_lists_into_the_slot_tools_call_reads(self): + 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, + ) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-user") + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="stored cred catalog", inputSchema={})] + ) + cache_byok_credential("byok-user", "byok-catalog", "stored-secret") + try: + await manager._get_tools_from_server(server=server, user_api_key_auth=user) + finally: + byok_credential_cache.delete_cache(byok_credential_cache_key("byok-user", "byok-catalog")) + + call_side = ListedToolsCaller(user_api_key_auth=user, mcp_auth_header="stored-secret") + listed = manager.get_listed_tool(server, "turn", call_side) + assert listed is not None and listed.description == "stored cred catalog" + @pytest.mark.parametrize( ("signer", "static_headers", "shared"), [ From 3650243b907bb7a077e6021523b6fdddfbff50ea Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 08:44:33 +0000 Subject: [PATCH 06/47] fix(mcp): drop the OAuth discovery cache when a server definition changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 11 ++++++ .../mcp_server/mcp_server_manager.py | 5 +++ .../mcp_server/test_discoverable_endpoints.py | 36 +++++++++++++++++++ 3 files changed, 52 insertions(+) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..91da7204435 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -141,6 +141,17 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) +def invalidate_oauth_metadata_cache(server_id: str) -> None: + """Drop cached upstream IdP metadata for a server whose definition changed.""" + for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: + del _OAUTH_METADATA_CACHE[cache_key] + for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: + lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) + if lock is None or lock.locked(): + continue + _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + + def encode_state_with_base_url( base_url: str, original_state: str, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b1f52bd5bcc..5647a44f814 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4555,8 +4555,13 @@ class MCPServerManager: self._template_discovery_cache.invalidate(server_id) def _invalidate_server_definition_caches(self, server_id: str) -> None: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton + invalidate_oauth_metadata_cache, + ) + self._invalidate_discovery_lists(server_id) self._listed_tools_by_server_id.pop(server_id, None) + invalidate_oauth_metadata_cache(server_id) def _discovers_per_caller(self, server: MCPServer) -> bool: return ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f9a0075e530..e6aba5f1727 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12611,3 +12611,39 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session( proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called() proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_server_drops_cached_upstream_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import LiteLLM_MCPServerTable + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + server = MCPServer( + server_id="oauth-cache-server", + name="oauth_cache_server", + url="http://old-upstream/mcp", + transport=MCPTransport.http, + ) + manager.registry[server.server_id] = server + stale_key: Final = (server.server_id, server.url) + other_key: Final = ("other-server", "http://other/mcp") + discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"}) + discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"}) + try: + await manager.update_server( + LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.name, + url="http://new-upstream/mcp", + transport=MCPTransport.http, + ) + ) + assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) From 7ca6686e82271e12f3c2162235410996d34a7b9b Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 08:45:18 +0000 Subject: [PATCH 07/47] refactor(mcp): drop a diff-narrating comment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/mcp_server_manager.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5647a44f814..4c257451345 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4401,8 +4401,6 @@ class MCPServerManager: verbose_logger.info("_get_tools_from_server for %s...", server.name) client = None - # tools/call resolves the BYOK credential before keying its listed-tools slot; resolve the - # same value here or a stored-credential server would list into a slot the call never reads. resolved_mcp_auth_header: Final = ( mcp_auth_header if not server.is_byok or isinstance(mcp_auth_header, dict) From c93d88ed797b83b4433801564866f3b56e331a60 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 22:26:08 +0000 Subject: [PATCH 08/47] fix(mcp): never validate a supplied header on the tools/list BYOK path The pre-listing resolver ran the tool-call byok_auth_required check even when the caller already supplied x-mcp-auth, and it ran outside the per-server error boundary, so a single deprecated-header caller dropped the server from the aggregate list. Listing now returns a supplied header unchanged and falls back to the stored credential without raising Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 21 +++++++++++++----- .../mcp_server/test_mcp_server_manager.py | 22 +++++++++++++++++++ 2 files changed, 38 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4c257451345..074a30eb0db 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1286,6 +1286,21 @@ async def _resolve_byok_mcp_auth_header( return mcp_auth_header +async def _byok_listing_auth_header( + mcp_server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | dict[str, str] | None, +) -> str | dict[str, str] | None: + """The credential a tools/list may use: a supplied header forwards unchanged, and a missing one + falls back to the stored credential without the tool-call path's byok_auth_required raise.""" + if not mcp_server.is_byok or mcp_auth_header is not None: + return mcp_auth_header + + from litellm.proxy._experimental.mcp_server.operations import _get_byok_credential + + return await _get_byok_credential(mcp_server, user_api_key_auth) + + def _client_forwarded_authorization_headers( mcp_server: MCPServer, oauth2_headers: dict[str, str] | None, @@ -4401,11 +4416,7 @@ class MCPServerManager: verbose_logger.info("_get_tools_from_server for %s...", server.name) client = None - resolved_mcp_auth_header: Final = ( - mcp_auth_header - if not server.is_byok or isinstance(mcp_auth_header, dict) - else await _resolve_byok_mcp_auth_header(server, user_api_key_auth, mcp_auth_header) - ) + resolved_mcp_auth_header: Final = await _byok_listing_auth_header(server, user_api_key_auth, mcp_auth_header) listed_caller: Final = ListedToolsCaller( user_api_key_auth=user_api_key_auth, mcp_auth_header=resolved_mcp_auth_header, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 368f0d09e3c..3683d00d2a3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7068,6 +7068,28 @@ class TestMCPServerManager: listed = manager.get_listed_tool(server, "turn", call_side) assert listed is not None and listed.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"), + ) + + assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" + @pytest.mark.parametrize( ("signer", "static_headers", "shared"), [ From 46cd428b65f92a445ae18b9ce51996219af8e9d0 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 23:43:36 +0000 Subject: [PATCH 09/47] test(mcp): assert the BYOK listing lands in the caller's listed-tool slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/test_mcp_server_manager.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 3683d00d2a3..e44c66f6e35 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7088,6 +7088,9 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), ) + caller: Final = ListedToolsCaller(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( From efc4f85fe68c37fdc2770cff6651f4aeb2e4ec8e Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 00:57:56 +0000 Subject: [PATCH 10/47] test(mcp): cover the deprecated string x-mcp-auth header on a BYOK tools/list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/mcp/test_mcp_credentials.py | 23 +++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/tests/integration/mcp/test_mcp_credentials.py b/tests/integration/mcp/test_mcp_credentials.py index 95a5e46646b..b3f5725040a 100644 --- a/tests/integration/mcp/test_mcp_credentials.py +++ b/tests/integration/mcp/test_mcp_credentials.py @@ -192,3 +192,26 @@ 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 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 = tuple( + item + for item in peer.drain() + if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list" + ) + assert len(listings) == 1, listings + assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr" From 25797cbf2a99583f92c7c8bdf0ef12797beec317 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 02:46:18 +0000 Subject: [PATCH 11/47] fix(mcp): key the per-caller listed-tool slot by the hashed token Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 45 ++++++++++++++++--- 2 files changed, 39 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 074a30eb0db..4c14b3ac1b8 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4606,7 +4606,7 @@ class MCPServerManager: return None auth: Final = caller.user_api_key_auth identity: Final = ( - (auth.user_id, auth.api_key) if auth is not None and self._discovers_per_caller(server) else None + (auth.user_id, auth.token) if auth is not None and self._discovers_per_caller(server) else None ) forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) header_env: Final = self._build_stdio_env(server, caller.raw_headers) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index e44c66f6e35..019b93925ea 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6931,8 +6931,8 @@ class TestMCPServerManager: url="http://srv", auth_type=MCPAuth.oauth2_token_exchange, ) - alice = UserAPIKeyAuth(user_id="alice", api_key="hashed-alice") - bob = UserAPIKeyAuth(user_id="bob", api_key="hashed-bob") + 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._create_prefixed_tools( @@ -6953,7 +6953,7 @@ class TestMCPServerManager: 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", api_key="k")) + carol = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="carol", token="k")) assert manager.get_listed_tool(server, "srv-read", carol) is None shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") @@ -6989,11 +6989,11 @@ class TestMCPServerManager: pytest.param( {"auth_type": MCPAuth.oauth2_token_exchange}, ListedToolsCaller( - user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", api_key="hashed-shared"), + 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", api_key="hashed-shared"), + 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", @@ -7106,8 +7106,8 @@ class TestMCPServerManager: 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", api_key="hashed-alice")) - bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="bob", api_key="hashed-bob")) + 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", @@ -7123,6 +7123,37 @@ class TestMCPServerManager: else: assert for_bob is None + def test_per_caller_slot_identity_is_the_token_not_the_api_key(self): + """Two callers sharing user_id split on the hashed token the admission validator stamps; + the raw api_key never enters the identity, so a caller carrying only that token lands on + the same slot.""" + from litellm.proxy._types import hash_token + + 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._create_prefixed_tools( + [MCPTool(name="turn", description="slot a", inputSchema={})], server, caller=alice + ) + assert manager.get_listed_tool(server, "srv-turn", bob) is None + + same_token = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="same-user", token=hash_token("sk-alpha")) + ) + listed = manager.get_listed_tool(server, "srv-turn", same_token) + + assert listed is not None and listed.description == "slot a" + @pytest.mark.asyncio async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self): manager = MCPServerManager() From b6f8480812228ec38c8e2e7e2a353d39c1742286 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 04:13:19 +0000 Subject: [PATCH 12/47] fix(mcp): key discovery cache by the hashed token Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_experimental/mcp_server/mcp_server_manager.py | 6 ++---- .../mcp_server/test_mcp_server_manager.py | 14 ++++++++++++++ 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4c14b3ac1b8..52022e02b0c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4663,16 +4663,14 @@ class MCPServerManager: if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): return server.server_id, None identity: Final = ( - (user_api_key_auth.user_id, user_api_key_auth.api_key) - if per_user and user_api_key_auth is not None - else None + (user_api_key_auth.user_id, user_api_key_auth.token) if per_user and user_api_key_auth is not None else None ) material: Final = json.dumps( (identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint), sort_keys=True, separators=(",", ":"), ) - return server.server_id, hashlib.sha256(material.encode()).hexdigest() + return server.server_id, hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest() async def get_prompts_from_server( self, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 019b93925ea..8a81a8e7bcb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14376,6 +14376,20 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "first" not in str(first) assert "second" not in str(second) + from litellm.proxy._types import hash_token + + same_user_other_token: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-second")), None, None, None, None + ) + same_token_no_key: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-first")), None, None, None, None + ) + with_key: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", api_key="sk-first"), None, None, None, None + ) + assert same_user_other_token != with_key + assert same_token_no_key == with_key + @pytest.mark.asyncio async def test_discovery_cache_retries_cancelled_fetches() -> None: From 444894be2ad0493600de50e673593ba8bf04597c Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 20:01:28 +0000 Subject: [PATCH 13/47] fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates Discovery-list cache identity now uses the hashed token instead of the raw api_key and treats MCPJWTSigner-signed servers as per caller. Server definition changes also drop the cached upstream OAuth metadata. OpenAPI listings look tools up under the normalized registry prefix with the separator, so an overlapping sibling prefix no longer leaks into the list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 11 +++ .../mcp_server/mcp_server_manager.py | 70 ++++++++------ .../mcp_server/test_discoverable_endpoints.py | 35 +++++++ .../mcp_server/test_mcp_server_manager.py | 92 +++++++++++++++++++ 4 files changed, 180 insertions(+), 28 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..91da7204435 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -141,6 +141,17 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) +def invalidate_oauth_metadata_cache(server_id: str) -> None: + """Drop cached upstream IdP metadata for a server whose definition changed.""" + for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: + del _OAUTH_METADATA_CACHE[cache_key] + for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: + lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) + if lock is None or lock.locked(): + continue + _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + + def encode_state_with_base_url( base_url: str, original_state: str, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 31896d9ddc5..d2a4b6ebad9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2664,7 +2664,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) _warn_legacy_delegate_auth_if_applicable(new_server, source="config") _warn_config_id_jag_server_outruns_sso(new_server) - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2868,7 +2868,7 @@ class MCPServerManager: global_mcp_tool_registry, ) - self._invalidate_discovery_lists(server.server_id) + self._invalidate_server_definition_caches(server.server_id) prefix_root: Final = normalize_server_name(get_server_prefix(server)) if server.spec_path and prefix_root: openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR @@ -3275,7 +3275,7 @@ class MCPServerManager: # env_vars_are_encrypted=False. new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3312,7 +3312,7 @@ class MCPServerManager: previous_server=self.registry[mcp_server.server_id], ) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -4492,24 +4492,16 @@ class MCPServerManager: if server.spec_path: # OpenAPI tools were stored in the registry under the prefix # active at registration time — fetch by that same prefix. - _tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server)) + registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR + _tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix) tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(_tools) # OpenAPI tools are stored in the registry with their prefix already # 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". - if not add_prefix: - prefix: Final = get_server_prefix(server) - sep: Final = MCP_TOOL_PREFIX_SEPARATOR - tools = [ - ( - t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]}) - if t.name.startswith(f"{prefix}{sep}") - else t - ) - for t in tools - ] - return tools + if add_prefix: + return tools + return [t.model_copy(update=MappingProxyType({"name": t.name[len(registry_prefix) :]})) for t in tools] else: tools = await self._fetch_tools_with_timeout(client, server.name) self._remember_upstream_initialize_instructions(server, client) @@ -4558,6 +4550,35 @@ class MCPServerManager: self._resource_discovery_cache.invalidate(server_id) self._template_discovery_cache.invalidate(server_id) + def _invalidate_server_definition_caches(self, server_id: str) -> None: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton + invalidate_oauth_metadata_cache, + ) + + self._invalidate_discovery_lists(server_id) + invalidate_oauth_metadata_cache(server_id) + + def _discovers_per_caller(self, server: MCPServer) -> bool: + return ( + 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) + or self._signs_caller_identity_upstream(server) + ) + + @staticmethod + def _signs_caller_identity_upstream(server: MCPServer) -> bool: + """Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may + tailor its catalog to the caller even though the server itself is configured as shared.""" + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server + get_mcp_jwt_signer, + ) + + if get_mcp_jwt_signer() is None: + return False + return not any(k.lower() == "authorization" for k in (server.static_headers or {})) + def _discovery_key( self, server: MCPServer, @@ -4568,25 +4589,18 @@ class MCPServerManager: subject_token: str | None, credential_fingerprint: str | None = None, ) -> _DiscoveryKey: - per_user: Final = ( - server.requires_per_user_auth - or self._references_per_user_env_var(server) - or server.delegate_auth_to_upstream - or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) - ) + per_user: Final = self._discovers_per_caller(server) if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): return server.server_id, None identity: Final = ( - (user_api_key_auth.user_id, user_api_key_auth.api_key) - if per_user and user_api_key_auth is not None - else None + (user_api_key_auth.user_id, user_api_key_auth.token) if per_user and user_api_key_auth is not None else None ) material: Final = json.dumps( (identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint), sort_keys=True, separators=(",", ":"), ) - return server.server_id, hashlib.sha256(material.encode()).hexdigest() + return server.server_id, hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest() async def get_prompts_from_server( self, @@ -6704,7 +6718,7 @@ class MCPServerManager: for server_id in previous_registry.keys() | registered_registry.keys(): if previous_registry.get(server_id) != registered_registry.get(server_id): - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.registry = registered_registry _warn_on_shared_identifier_prefixes(registered_registry.values()) # A discovery task may have published into ``previous_registry`` while diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f9a0075e530..0a92b548d62 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12611,3 +12611,38 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session( proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called() proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_server_drops_cached_upstream_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + server = MCPServer( + server_id="oauth-cache-server", + name="oauth_cache_server", + url="http://old-upstream/mcp", + transport=MCPTransport.http, + ) + manager.registry[server.server_id] = server + stale_key: Final = (server.server_id, server.url) + other_key: Final = ("other-server", "http://other/mcp") + discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"}) + discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"}) + try: + await manager.update_server( + LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.name, + url="http://new-upstream/mcp", + transport=MCPTransport.http, + ) + ) + assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 70ef4312f4c..85b4aaa11eb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14058,6 +14058,98 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "first" not in str(first) assert "second" not in str(second) + from litellm.proxy._types import hash_token + + same_user_other_token: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-second")), None, None, None, None + ) + same_token_no_key: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-first")), None, None, None, None + ) + with_key: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", api_key="sk-first"), None, None, None, None + ) + assert same_user_other_token != with_key + assert same_token_no_key == with_key + + +@pytest.mark.parametrize( + ("signer", "static_headers", "shared"), + [ + pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), + pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), + pytest.param(None, None, True, id="no-signer-stays-shared"), + ], +) +def test_jwt_signer_makes_a_shared_server_discover_per_caller(signer, static_headers, shared) -> None: + manager: Final = MCPServerManager() + server: Final = _discovery_server().model_copy(update={"static_headers": static_headers}) + alice: Final = UserAPIKeyAuth(user_id="alice", token="hashed-alice") + bob: Final = 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, + ): + for_alice: Final = manager._discovery_key(server, alice, None, None, None, None) + for_bob: Final = manager._discovery_key(server, bob, None, None, None, None) + + assert (for_alice == for_bob) is shared + + +def _register_local_tool(name: str, description: str) -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + async def _handler(**kwargs): + return None + + global_mcp_tool_registry.register_tool( + name=name, description=description, input_schema={"type": "object"}, handler=_handler + ) + + +def _openapi_server(name: str) -> MCPServer: + return MCPServer( + server_id=f"{name}-id", name=name, alias=name, transport=MCPTransport.http, url=None, spec_path="/spec.yaml" + ) + + +@pytest.mark.asyncio +async def test_openapi_listing_ignores_overlapping_server_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + _register_local_tool("pet-list", "Local pet tool") + _register_local_tool("petstore-list", "Foreign petstore tool") + try: + prefixed: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=True) + bare: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=False) + finally: + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + + assert [t.name for t in prefixed] == ["pet-list"] + assert [t.name for t in bare] == ["list"] + + +@pytest.mark.asyncio +async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + _register_local_tool("pet_store-list", "Pet store tool") + try: + listed: Final = await manager._get_tools_from_server(server=_openapi_server("pet store"), add_prefix=False) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + + assert [t.name for t in listed] == ["list"] + @pytest.mark.asyncio async def test_discovery_cache_retries_cancelled_fetches() -> None: From cee44351e306c75be29b739544e71d3a5209f9b9 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 20:17:43 +0000 Subject: [PATCH 14/47] fix(mcp): keep the discovery cache digest call unchanged so CodeQL matches the existing alert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/mcp_server_manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d2a4b6ebad9..4504f4a85d2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4600,7 +4600,7 @@ class MCPServerManager: sort_keys=True, separators=(",", ":"), ) - return server.server_id, hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest() + return server.server_id, hashlib.sha256(material.encode()).hexdigest() async def get_prompts_from_server( self, From db2082858c5cbe9e80d9abe36d3b79c8b648c091 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 20:26:03 +0000 Subject: [PATCH 15/47] refactor(mcp): derive the listed-tool caller identity from the discovery cache key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_experimental/mcp_server/mcp_server_manager.py | 12 +++--------- 1 file changed, 3 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4b27c015a41..11159a1c72b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4640,10 +4640,7 @@ class MCPServerManager: if server.spec_path or caller is None: return None auth: Final = caller.user_api_key_auth - identity: Final = ( - (auth.user_id, auth.token) if auth is not None and self._discovers_per_caller(server) else None - ) - forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) + forwarded: Final = dict(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 = ( @@ -4651,11 +4648,8 @@ class MCPServerManager: if server.is_client_forwarded_token or server.auth_type == MCPAuth.oauth2_token_exchange else None ) - inputs: Final = (identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer) - if not any(inputs): - return None - material: Final = json.dumps(inputs, sort_keys=True, separators=(",", ":")) - return hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest() + _, digest = self._discovery_key(server, auth, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer) + return digest @staticmethod def _forwarded_header_values( From 9ce2803f429d2370fd378c5a8d4cbd30f4ad7791 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 22:03:22 +0000 Subject: [PATCH 16/47] fix(mcp): guard OAuth metadata cache writes with a per-server generation and drop unproven per-caller discovery keys An upstream metadata fetch that started before a server edit could store its stale reply after invalidate_oauth_metadata_cache ran. Invalidation now bumps a per-server generation and the fetch only stores when the generation it captured before I/O is unchanged. The MCPJWTSigner-based per-caller discovery classification and the api_key to token key change had no reproduction (the signer only injects on tools/list, and UserAPIKeyAuth hashes api_key in place), so both go back to the merge-base behavior. Integration coverage under tests/integration/mcp: overlapping OpenAPI aliases, a config-declared server name with a space, OAuth metadata refetch after a save, and the in-flight stale-write race Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 26 ++-- .../mcp_server/mcp_server_manager.py | 32 ++--- tests/integration/mcp/test_mcp_management.py | 52 ++++++++ .../mcp/test_oauth_configuration.py | 115 +++++++++++++++++- .../mcp_server/test_discoverable_endpoints.py | 43 +++++++ .../mcp_server/test_mcp_server_manager.py | 38 ------ 6 files changed, 228 insertions(+), 78 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 91da7204435..d2851cf4fbf 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -107,6 +107,9 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # Per-(server_id, resource_url) async locks so concurrent discovery requests # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} +# Per-server_id generation, bumped on invalidation so a fetch that started before the server +# definition changed cannot repopulate the cache with the stale reply. +_OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( tags=["mcp"], @@ -143,6 +146,7 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: def invalidate_oauth_metadata_cache(server_id: str) -> None: """Drop cached upstream IdP metadata for a server whose definition changed.""" + _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: @@ -2377,6 +2381,14 @@ async def fetch_upstream_oauth_protected_resource( cached = _OAUTH_METADATA_CACHE.get(cache_key) if cached is not None and cached[0] > now: return cached[1] + generation: Final = _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) + + def store(payload: dict | None, ttl_seconds: int) -> None: + if _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) != generation: + return + stored_at: Final = time.time() + _OAUTH_METADATA_CACHE[cache_key] = (stored_at + ttl_seconds, payload) + _prune_oauth_metadata_cache(stored_at) host_base: Final = f"{upstream.scheme}://{upstream.netloc}" candidates: Final = [f"{host_base}/.well-known/oauth-protected-resource"] @@ -2418,12 +2430,7 @@ async def fetch_upstream_oauth_protected_resource( ) continue if isinstance(payload, dict): - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_CACHE_TTL_SECONDS, - payload, - ) - _prune_oauth_metadata_cache(now) + store(payload, _OAUTH_METADATA_CACHE_TTL_SECONDS) return payload if len(network_errors) == len(candidates): @@ -2432,12 +2439,7 @@ async def fetch_upstream_oauth_protected_resource( # Negative-result caching: when no candidate yielded a usable payload, # remember that for a shorter TTL so we don't re-fetch on every # subsequent discovery request (and so the per-key lock can be pruned). - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS, - None, - ) - _prune_oauth_metadata_cache(now) + store(None, _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS) return None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4504f4a85d2..ca05266c827 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4558,27 +4558,6 @@ class MCPServerManager: self._invalidate_discovery_lists(server_id) invalidate_oauth_metadata_cache(server_id) - def _discovers_per_caller(self, server: MCPServer) -> bool: - return ( - 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) - or self._signs_caller_identity_upstream(server) - ) - - @staticmethod - def _signs_caller_identity_upstream(server: MCPServer) -> bool: - """Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may - tailor its catalog to the caller even though the server itself is configured as shared.""" - from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server - get_mcp_jwt_signer, - ) - - if get_mcp_jwt_signer() is None: - return False - return not any(k.lower() == "authorization" for k in (server.static_headers or {})) - def _discovery_key( self, server: MCPServer, @@ -4589,11 +4568,18 @@ class MCPServerManager: subject_token: str | None, credential_fingerprint: str | None = None, ) -> _DiscoveryKey: - per_user: Final = self._discovers_per_caller(server) + per_user: Final = ( + server.requires_per_user_auth + or self._references_per_user_env_var(server) + or server.delegate_auth_to_upstream + or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) + ) if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): return server.server_id, None identity: Final = ( - (user_api_key_auth.user_id, user_api_key_auth.token) if per_user and user_api_key_auth is not None else None + (user_api_key_auth.user_id, user_api_key_auth.api_key) + if per_user and user_api_key_auth is not None + else None ) material: Final = json.dumps( (identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint), diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index 66c30a62bde..67cdbbff5a4 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -1,3 +1,4 @@ +import itertools import uuid from pathlib import Path from typing import Final @@ -7,10 +8,13 @@ import yaml from integration._support.client import Gateway, eventually from integration._support.mcp import ( McpCaller, + McpPeer, call_tool, delete_mcp, forget_mcp, + listed_tools, mcp_peer, + openapi_peer, register_mcp, tool_calls, tool_names, @@ -189,6 +193,54 @@ def test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide(gateway: Ga scenario.cleanups.callback(forget_mcp, gateway, winner) +def _openapi_server_lists_and_calls_only_its_own_tools( + gateway: Gateway, key: str, peer: McpPeer, identity: str +) -> None: + listed: Final = set(listed_tools(gateway, key, identity)) + assert listed == {"getpet", "createpet"}, (identity, listed) + peer.drain() + called: Final = call_tool(gateway, key, identity, "getpet", {"petId": "7"}) + assert called.status_code == 200, called.text + assert [(item["method"], item["path"]) for item in peer.drain()] == [("GET", "/pets/7")], identity + + +def test_openapi_listing_is_scoped_to_the_exact_alias_when_aliases_overlap(gateway: Gateway) -> None: + with openapi_peer() as short, openapi_peer() as long, gateway.scenario() as scenario: + stem: Final = "pet" + uuid.uuid4().hex[:8] + servers: Final = tuple( + (peer, alias, register_mcp(scenario, peer, alias)) + for peer, alias in ((short, stem), (long, stem + "store")) + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity for _, _, identity in servers]}) + for peer, _, identity in servers: + _openapi_server_lists_and_calls_only_its_own_tools(gateway, key, peer, identity) + aggregate: Final = McpCaller(gateway, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + assert sorted(aggregate.tools) == sorted( + f"{prefix}-{tool}" for prefix, tool in itertools.product((stem, stem + "store"), ("getpet", "createpet")) + ), aggregate.tools + assert all(peer.drain() == () for peer, _, _ in servers), "listing must not reach any OpenAPI upstream" + + +def test_config_declared_openapi_server_with_a_space_in_its_name_lists_its_tools( + gateway: Gateway, tmp_path: Path +) -> None: + with openapi_peer() as peer: + config: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text()) + name: Final = "pet store " + uuid.uuid4().hex[:8] + config["mcp_servers"] = {name: peer.registration()} + path: Final = tmp_path / "openapi-space.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + identity: Final = next(i for i, s in _servers(candidate).items() if s["server_name"] == name) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + _openapi_server_lists_and_calls_only_its_own_tools(candidate, key, peer, identity) + aggregate: Final = McpCaller(candidate, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + prefix: Final = name.replace(" ", "_") + assert sorted(aggregate.tools) == [f"{prefix}-createpet", f"{prefix}-getpet"], aggregate.tools + + def test_invalid_registrations_are_rejected(gateway: Gateway) -> None: with mcp_peer() as peer, gateway.scenario() as scenario: alias: Final = "mgmt" + uuid.uuid4().hex[:8] diff --git a/tests/integration/mcp/test_oauth_configuration.py b/tests/integration/mcp/test_oauth_configuration.py index 4c46c706054..3a023f6d38e 100644 --- a/tests/integration/mcp/test_oauth_configuration.py +++ b/tests/integration/mcp/test_oauth_configuration.py @@ -1,17 +1,23 @@ import json import queue +import threading import uuid -from urllib.parse import parse_qs, urlsplit -from typing import Final, Literal +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field from pathlib import Path +from typing import Final, Literal +from urllib.parse import parse_qs, urlsplit import pytest - -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, Scenario, eventually from integration._support.database import read_rows from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names from integration._support.process import owned_proxy -from integration._support.wire import Reply, Request, wire_server +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +_Upstream = Callable[[Request], Reply] @pytest.mark.covers("other.mcp.oauth.discovery_cannot_erase_configured_authorization_endpoint") @@ -104,6 +110,105 @@ def test_partial_discovery_and_unrelated_edit_keep_actual_authorization_destinat assert updated.status_code == 202, updated.text +@dataclass(frozen=True, slots=True) +class _Hold: + armed: threading.Event = field(default_factory=threading.Event) + released: threading.Event = field(default_factory=threading.Event) + + +def _idp_upstream(origin: Callable[[], str], moved: threading.Event, hold: _Hold | None = None) -> _Upstream: + def issuer() -> str: + return origin() + ("/idp-after" if moved.is_set() else "/idp-before") + + def respond(request: Request) -> Reply: + if "oauth-authorization-server" in request.target or "openid-configuration" in request.target: + current: Final = issuer() + return Reply( + body=json.dumps( + { + "issuer": current, + "authorization_endpoint": current + "/authorize", + "token_endpoint": current + "/token", + } + ).encode() + ) + if request.target.startswith("/.well-known/oauth-protected-resource"): + body: Final = json.dumps({"resource": origin() + "/mcp", "authorization_servers": [issuer()]}).encode() + if hold is not None and hold.armed.is_set(): + assert hold.released.wait(timeout=15), "the held upstream metadata reply was never released" + return Reply(body=body) + return Reply(status=404, body=b'{"error":"unexpected"}') + + return respond + + +def _register_pass_through(scenario: Scenario, wire: Wire, alias: str) -> str: + return register_mcp(scenario, McpPeer(wire.url + "/mcp", queue.Queue()), alias, auth_type="true_passthrough") + + +def _wire_requests(wire: Wire, seen: list[Request]) -> Callable[[], tuple[Request, ...]]: + def observed() -> tuple[Request, ...]: + seen.extend(wire.drain()) + return tuple(seen) + + return observed + + +def _registration_discovery_settled(requests: tuple[Request, ...]) -> bool: + return any( + "oauth-authorization-server" in item.target or "openid-configuration" in item.target for item in requests + ) + + +def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str, ...]: + response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp") + assert response.status_code == 200, response.text + return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"])) + + +def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + moved.set() + wire.drain() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + assert any(request.target.startswith("/.well-known/oauth-protected-resource") for request in wire.drain()), ( + "the save must send protected-resource discovery back to the upstream" + ) + + +def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None: + moved: Final = threading.Event() + hold: Final = _Hold() + with ( + wire_server(_idp_upstream(lambda: wire.url, moved, hold)) as wire, + gateway.scenario() as scenario, + ThreadPoolExecutor(max_workers=1) as pool, + ): + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + seen: Final[list[Request]] = [] + observed: Final = _wire_requests(wire, seen) + eventually(observed, _registration_discovery_settled, seconds=10) + settled: Final = len(seen) + hold.armed.set() + stale: Final = pool.submit(_advertised_authorization_servers, gateway, alias) + eventually(observed, lambda requests: len(requests) > settled, seconds=10) + assert seen[settled].target.startswith("/.well-known/oauth-protected-resource"), seen[settled:] + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + hold.released.set() + assert stale.result(timeout=30) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + + @pytest.mark.covers("other.mcp.oauth.same_url_credentials_are_isolated_by_user_and_server") @pytest.mark.parametrize("transition", ("revoke", "expire")) def test_same_url_oauth_credentials_and_revocation_are_isolated_by_user_and_server( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 0a92b548d62..7a2af8dfcca 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12646,3 +12646,46 @@ async def test_update_server_drops_cached_upstream_oauth_metadata(): finally: discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) + + +@pytest.mark.asyncio +async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cache(): + import asyncio + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + fetch_upstream_oauth_protected_resource, + invalidate_oauth_metadata_cache, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="stale-write-server", name="stale_write", url="http://upstream/mcp", transport=MCPTransport.http + ) + cache_key: Final = (server.server_id, server.url) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get(url: str, headers: dict[str, str]) -> MagicMock: + started.set() + await release.wait() + return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["old-idp"]})) + + client = MagicMock() + client.get = slow_get + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ): + in_flight: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + release.set() + assert await in_flight == {"authorization_servers": ["old-idp"]} + assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 85b4aaa11eb..130d47bbafa 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14058,44 +14058,6 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "first" not in str(first) assert "second" not in str(second) - from litellm.proxy._types import hash_token - - same_user_other_token: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-second")), None, None, None, None - ) - same_token_no_key: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-first")), None, None, None, None - ) - with_key: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", api_key="sk-first"), None, None, None, None - ) - assert same_user_other_token != with_key - assert same_token_no_key == with_key - - -@pytest.mark.parametrize( - ("signer", "static_headers", "shared"), - [ - pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), - pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), - pytest.param(None, None, True, id="no-signer-stays-shared"), - ], -) -def test_jwt_signer_makes_a_shared_server_discover_per_caller(signer, static_headers, shared) -> None: - manager: Final = MCPServerManager() - server: Final = _discovery_server().model_copy(update={"static_headers": static_headers}) - alice: Final = UserAPIKeyAuth(user_id="alice", token="hashed-alice") - bob: Final = 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, - ): - for_alice: Final = manager._discovery_key(server, alice, None, None, None, None) - for_bob: Final = manager._discovery_key(server, bob, None, None, None, None) - - assert (for_alice == for_bob) is shared - def _register_local_tool(name: str, description: str) -> None: from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry From c0099a45deb9b4c122c95a6b768e6c3583964e9a Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 22:37:53 +0000 Subject: [PATCH 17/47] fix(mcp): keep OAuth metadata generations only while a fetch is in flight Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 15 +++++++++++++-- .../mcp_server/test_discoverable_endpoints.py | 16 ++++++++++++++++ 2 files changed, 29 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d2851cf4fbf..d4e95e7ad70 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -108,7 +108,8 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} # Per-server_id generation, bumped on invalidation so a fetch that started before the server -# definition changed cannot repopulate the cache with the stale reply. +# definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch +# in flight carry an entry; the rest are pruned with the cache. _OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( @@ -143,10 +144,20 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + for server_id in [sid for sid in _OAUTH_METADATA_GENERATIONS if not _oauth_metadata_fetch_in_flight(sid)]: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) + + +def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: + return any(lock.locked() for cache_key, lock in _OAUTH_METADATA_FETCH_LOCKS.items() if cache_key[0] == server_id) + def invalidate_oauth_metadata_cache(server_id: str) -> None: """Drop cached upstream IdP metadata for a server whose definition changed.""" - _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 + if _oauth_metadata_fetch_in_flight(server_id): + _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 + else: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 7a2af8dfcca..d9135a999b6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12686,6 +12686,22 @@ async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cach release.set() assert await in_flight == {"authorization_servers": ["old-idp"]} assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + discoverable_endpoints._prune_oauth_metadata_cache() + assert server.server_id not in discoverable_endpoints._OAUTH_METADATA_GENERATIONS finally: discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) + + +def test_invalidating_an_idle_server_leaves_no_generation_behind(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import invalidate_oauth_metadata_cache + + server_ids: Final = tuple(f"churned-server-{i}" for i in range(50)) + try: + for server_id in server_ids: + invalidate_oauth_metadata_cache(server_id) + assert not set(server_ids) & set(discoverable_endpoints._OAUTH_METADATA_GENERATIONS) + finally: + for server_id in server_ids: + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server_id, None) From 9786e3509aee0ab937c85946c91e084f3513918b Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 22:58:32 +0000 Subject: [PATCH 18/47] fix(mcp): count queued OAuth metadata fetchers so invalidation survives lock handoff Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 38 ++++++++----- .../mcp_server/test_discoverable_endpoints.py | 54 +++++++++++++++++++ 2 files changed, 80 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d4e95e7ad70..dd72328732b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -3,7 +3,8 @@ import html as _html import json import secrets import time -from collections.abc import Callable, Mapping +from collections.abc import AsyncIterator, Callable, Mapping +from contextlib import asynccontextmanager from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Optional from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse @@ -107,6 +108,10 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # Per-(server_id, resource_url) async locks so concurrent discovery requests # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} +# Callers inside ``_oauth_metadata_fetch_slot`` per cache key, lock waiters included. ``Lock.locked()`` +# reads False between one holder's release and the next waiter's wake-up, so it cannot tell an +# idle lock from one being handed off. +_OAUTH_METADATA_FETCHERS: Final[dict[tuple[str, str], int]] = {} # Per-server_id generation, bumped on invalidation so a fetch that started before the server # definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch # in flight carry an entry; the rest are pruned with the cache. @@ -134,13 +139,10 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: for cache_key in cache_keys_by_expiry[:overflow]: _OAUTH_METADATA_CACHE.pop(cache_key, None) - # Drop locks whose cache entry has been evicted and that aren't currently - # held; held locks stay so in-flight callers continue to coalesce. + # Drop locks whose cache entry has been evicted and that nobody holds or + # waits on; the rest stay so in-flight callers continue to coalesce. for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS): - if cache_key in _OAUTH_METADATA_CACHE: - continue - lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) - if lock is None or lock.locked(): + if cache_key in _OAUTH_METADATA_CACHE or cache_key in _OAUTH_METADATA_FETCHERS: continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -149,7 +151,21 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: - return any(lock.locked() for cache_key, lock in _OAUTH_METADATA_FETCH_LOCKS.items() if cache_key[0] == server_id) + return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS) + + +@asynccontextmanager +async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]: + _OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1 + try: + async with _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()): + yield + finally: + remaining: Final = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) - 1 + if remaining > 0: + _OAUTH_METADATA_FETCHERS[cache_key] = remaining + else: + _OAUTH_METADATA_FETCHERS.pop(cache_key, None) def invalidate_oauth_metadata_cache(server_id: str) -> None: @@ -161,8 +177,7 @@ def invalidate_oauth_metadata_cache(server_id: str) -> None: for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: - lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) - if lock is None or lock.locked(): + if cache_key in _OAUTH_METADATA_FETCHERS: continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -2386,8 +2401,7 @@ async def fetch_upstream_oauth_protected_resource( if cached is not None and cached[0] > now: return cached[1] - lock: Final = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()) - async with lock: + async with _oauth_metadata_fetch_slot(cache_key): now = time.time() cached = _OAUTH_METADATA_CACHE.get(cache_key) if cached is not None and cached[0] > now: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index d9135a999b6..b1f0b3fa67e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12693,6 +12693,60 @@ async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cach discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) +@pytest.mark.asyncio +async def test_fetch_waiting_on_a_lock_handoff_stays_tracked_through_invalidation(): + import asyncio + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + fetch_upstream_oauth_protected_resource, + invalidate_oauth_metadata_cache, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="handoff-server", name="handoff", url="http://upstream/mcp", transport=MCPTransport.http + ) + cache_key: Final = (server.server_id, server.url) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get(url: str, headers: dict[str, str]) -> MagicMock: + started.set() + await release.wait() + return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["pre-save-idp"]})) + + client = MagicMock() + client.get = slow_get + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ): + async with discoverable_endpoints._oauth_metadata_fetch_slot(cache_key): + shared_lock: Final = discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS[cache_key] + waiting: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + for _ in range(3): + await asyncio.sleep(0) + assert not started.is_set() and not waiting.done() + invalidate_oauth_metadata_cache(server.server_id) + assert discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.get(cache_key) is shared_lock + assert discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + release.set() + assert await waiting == {"authorization_servers": ["pre-save-idp"]} + assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert not discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id) + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_FETCHERS.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) + + def test_invalidating_an_idle_server_leaves_no_generation_behind(): from litellm.proxy._experimental.mcp_server import discoverable_endpoints from litellm.proxy._experimental.mcp_server.discoverable_endpoints import invalidate_oauth_metadata_cache From a3c3ca08cd43ba03d4864c259bab0468e64099e3 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 23:04:40 +0000 Subject: [PATCH 19/47] fix(mcp): keep a held OAuth metadata lock registered even when no fetcher slot claims it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index dd72328732b..7a0f59c3c2b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -142,7 +142,7 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: # Drop locks whose cache entry has been evicted and that nobody holds or # waits on; the rest stay so in-flight callers continue to coalesce. for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS): - if cache_key in _OAUTH_METADATA_CACHE or cache_key in _OAUTH_METADATA_FETCHERS: + if cache_key in _OAUTH_METADATA_CACHE or not _oauth_metadata_lock_idle(cache_key): continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -154,6 +154,13 @@ def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS) +def _oauth_metadata_lock_idle(cache_key: tuple[str, str]) -> bool: + if cache_key in _OAUTH_METADATA_FETCHERS: + return False + lock: Final = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) + return lock is None or not lock.locked() + + @asynccontextmanager async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]: _OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1 @@ -177,7 +184,7 @@ def invalidate_oauth_metadata_cache(server_id: str) -> None: for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: - if cache_key in _OAUTH_METADATA_FETCHERS: + if not _oauth_metadata_lock_idle(cache_key): continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) From 83ae2a9bbe3eaa7064f958ba3f9633475d710ea6 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 23:43:18 +0000 Subject: [PATCH 20/47] test(mcp): prove a peer worker drops stale upstream OAuth metadata after a save elsewhere Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp/test_oauth_configuration.py | 24 +++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/tests/integration/mcp/test_oauth_configuration.py b/tests/integration/mcp/test_oauth_configuration.py index 3a023f6d38e..fe2b1069f04 100644 --- a/tests/integration/mcp/test_oauth_configuration.py +++ b/tests/integration/mcp/test_oauth_configuration.py @@ -166,6 +166,14 @@ def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"])) +def _eventually_advertises(gateway: Gateway, alias: str, issuer: str) -> None: + eventually( + lambda: gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp"), + lambda response: response.status_code == 200 and response.json()["authorization_servers"] == [issuer], + seconds=40, + ) + + def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None: moved: Final = threading.Event() with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: @@ -183,6 +191,22 @@ def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gate ) +def test_peer_worker_stops_advertising_the_old_idp_after_a_save_on_another_worker( + gateway: Gateway, peer: Gateway +) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + _eventually_advertises(peer, alias, wire.url + "/idp-before") + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + _eventually_advertises(peer, alias, wire.url + "/idp-after") + + def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None: moved: Final = threading.Event() hold: Final = _Hold() From d14fb9e6e20eef1d4cfc2650cad5efe9121dc116 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 02:16:02 +0000 Subject: [PATCH 21/47] test(mcp): return one masked text per scanned string in the selected-guardrail REST test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/test_rest_endpoints.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index e82ab28bb4c..a27dd8d04bd 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -3300,7 +3300,8 @@ async def test_request_selected_tool_specific_guardrail_applies_to_virtual_execu default_on=False, custom_code="def apply_guardrail(inputs, request_data, input_type):\n" ' if inputs.get("tools", [{}])[0].get("function", {}).get("name") == "execute":\n' - f' return {{"action": "{action}", "reason": "resolved tool blocked", "texts": ["redacted"]}}\n' + ' texts = [t.replace("confidential", "redacted") for t in inputs.get("texts", [])]\n' + f' return {{"action": "{action}", "reason": "resolved tool blocked", "texts": texts}}\n' " return allow()\n", ) manager: Final = mcp_server_manager.MCPServerManager() From 75900240d550e6a2559d86a2e1b8faf9f876dec8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 02:40:47 +0000 Subject: [PATCH 22/47] fix(mcp): fold the signed caller into the discovery digest instead of a second key hash Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 23 +++++++++++-------- 1 file changed, 13 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 6719198629a..ca729e1b75a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4650,11 +4650,6 @@ class MCPServerManager: if server.spec_path or caller is None: return None auth: Final = caller.user_api_key_auth - signed_caller: Final = ( - f"{auth.user_id}:{auth.api_key}" - if auth is not None and self._signs_caller_identity_upstream(server) - else None - ) forwarded: Final = dict(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 @@ -4663,10 +4658,16 @@ class MCPServerManager: if server.is_client_forwarded_token or server.auth_type == MCPAuth.oauth2_token_exchange else None ) - _, digest = self._discovery_key(server, auth, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer) - if signed_caller is None: - return digest - return hashlib.sha256(f"{digest}:{signed_caller}".encode()).hexdigest() + _, digest = self._discovery_key( + server, + auth, + caller.mcp_auth_header, + forwarded, + stdio_env, + caller_bearer, + per_caller=self._signs_caller_identity_upstream(server), + ) + return digest @staticmethod def _signs_caller_identity_upstream(server: MCPServer) -> bool: @@ -4714,9 +4715,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) From f8734fd81513b0a4788fa04852446aace7c75e12 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 00:39:01 +0000 Subject: [PATCH 23/47] fix(mcp): satisfy type discipline gate on listed-tool identity Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2866a50bf5e..02d7255700d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4443,7 +4443,7 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None = None, client_ip: str | None = None, proxy_logging_obj: ProxyLogging | None = None, - ) -> list[MCPTool]: + ) -> Sequence[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -4566,7 +4566,7 @@ 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 = list(guarded_openapi) + unprefixed_tools: Final = guarded_openapi self._record_listed_tools(server, unprefixed_tools, listed_caller) if not add_prefix: return unprefixed_tools @@ -4583,7 +4583,7 @@ class MCPServerManager: raw_headers=raw_headers, ) prefixed_or_original_tools: Final = self._create_prefixed_tools( - list(guarded_tools), server, add_prefix=add_prefix, caller=listed_caller + guarded_tools, server, add_prefix=add_prefix, caller=listed_caller ) return prefixed_or_original_tools @@ -4649,7 +4649,7 @@ class MCPServerManager: if server.spec_path or caller is None: return None auth: Final = caller.user_api_key_auth - forwarded: Final = dict(self._forwarded_header_values(server, caller.raw_headers)) or None + 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 = ( @@ -4676,7 +4676,7 @@ class MCPServerManager: if get_mcp_jwt_signer() is None: return False - return not any(k.lower() == "authorization" for k in (server.static_headers or {})) + return server.static_headers is None or not any(k.lower() == "authorization" for k in server.static_headers) @staticmethod def _forwarded_header_values( @@ -4694,7 +4694,7 @@ class MCPServerManager: ) -> None: identity: Final = self._listed_tools_identity(server, caller) listing: Final = MappingProxyType({tool.name: tool for tool in tools}) - existing: Final[_ListedToolsByCaller] = self._listed_tools_by_server_id.get(server.server_id, {}) + 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) @@ -5609,7 +5609,7 @@ class MCPServerManager: def _create_prefixed_tools( self, - tools: list[MCPTool], + tools: Sequence[MCPTool], server: MCPServer, add_prefix: bool = True, caller: ListedToolsCaller | None = None, @@ -5644,13 +5644,13 @@ class MCPServerManager: 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, {}).get(identity) + listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity) if not listed: return None tool: Final = listed.get(name) or listed.get(strip_known_server_prefix(name, server)) if tool is None: return None - description: Final = (server.tool_name_to_description or {}).get(tool.name) + description: Final = server.tool_name_to_description.get(tool.name) if server.tool_name_to_description else None return tool if description is None else tool.model_copy(update={"description": description}) def _create_prefixed_prompts( From 2e1c6bfc70eeac8e3f9168d6a93b2a408d0630e8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 08:32:47 +0000 Subject: [PATCH 24/47] fix(mcp): hand tools/call hooks the exact catalog entry tools/list served get_listed_tool re-applied the admin description override on top of the cached listing, so a guardrail-masked description was restored to its original wording at call time, and the OpenAPI / local-registry call path built its metadata from the registry instead of the guarded caller catalog. Both paths now return the cached entry as served, falling back to the registry only when no listing was recorded Adds tests/integration/mcp/test_mcp_listed_tool_metadata.py (red on the prior head for the two regressions, red on the merge base for the feature, green on this head) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 6 +- .../_experimental/mcp_server/operations.py | 5 + tests/integration/_support/mcp.py | 14 +- .../mcp/test_mcp_listed_tool_metadata.py | 169 +++++++++++++++ .../mcp_server/test_mcp_server.py | 197 ++++++++++++++---- .../mcp_server/test_mcp_server_manager.py | 60 ++++-- 6 files changed, 380 insertions(+), 71 deletions(-) create mode 100644 tests/integration/mcp/test_mcp_listed_tool_metadata.py diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 02d7255700d..a247ce64bf5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5647,11 +5647,7 @@ class MCPServerManager: listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity) if not listed: return None - tool: Final = listed.get(name) or listed.get(strip_known_server_prefix(name, server)) - if tool is None: - return None - description: Final = server.tool_name_to_description.get(tool.name) if server.tool_name_to_description else None - return tool if description is None else tool.model_copy(update={"description": description}) + return listed.get(name) or listed.get(strip_known_server_prefix(name, server)) def _create_prefixed_prompts( self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index cccdf4c59c0..5a63fedce73 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1608,6 +1608,11 @@ async def _list_mcp_resource_templates( def _registered_tool_metadata(name: str, registered: RegisteredTool, server: MCPServer) -> MCPTool: + """The tool as ``tools/list`` served it (pinned, overridden, guardrail-masked) when a listing was + recorded for ``server``, else the registry entry with the admin description override applied.""" + listed: Final = global_mcp_server_manager.get_listed_tool(server, name) + if listed is not None: + return listed overrides: Final = server.tool_name_to_description description: Final = overrides.get(name, registered.description) if overrides else registered.description return MCPTool(name=name, description=description, input_schema=registered.input_schema) 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/mcp/test_mcp_listed_tool_metadata.py b/tests/integration/mcp/test_mcp_listed_tool_metadata.py new file mode 100644 index 00000000000..0381e8abb00 --- /dev/null +++ b/tests/integration/mcp/test_mcp_listed_tool_metadata.py @@ -0,0 +1,169 @@ +"""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 uuid +from collections.abc import Iterator, Mapping +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, gateway_from_environment +from integration._support.mcp import ( + EntryPoint, + McpCaller, + ScriptedTool, + listed_tools, + openapi_peer, + register_mcp, + scripted_peer, + text_result, +) +from integration._support.process import owned_proxy + +_ECHO: Final = "catalog-echo:" +_PROBE: Final = "catalog-probe" +_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, + }, + } + ] + path: Final = directory / "config.yaml" + path.write_text(yaml.safe_dump(config)) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path) as 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) + + +@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" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 62a7830a85d..57268902614 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -24,6 +24,7 @@ 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 @@ -86,9 +87,6 @@ def cleanup_mcp_global_state(): yield - - - def _call_tool_params(name, arguments=None): from mcp.types import CallToolRequestParams @@ -100,6 +98,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""" @@ -296,7 +295,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 @@ -1168,20 +1169,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 {}), + } + ], } @@ -1675,7 +1688,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", @@ -1914,8 +1929,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 @@ -4057,7 +4072,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 = ( @@ -6565,8 +6581,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 @@ -7412,7 +7432,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 +7679,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 +7955,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( @@ -8132,6 +8155,57 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_admin_description_client assert (handed_tool.description, handed_tool.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 + manager._record_listed_tools( + petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)], None + ) + 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=UserAPIKeyAuth(api_key="sk-user", user_id="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_hooks_the_metadata_of_the_operation_it_runs_when_names_collide(): """An OpenAPI operation whose name starts with its own server prefix must not be reported to the @@ -8236,7 +8310,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( @@ -8735,7 +8810,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 @@ -10532,7 +10609,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 @@ -10615,12 +10694,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 @@ -10638,7 +10720,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)), @@ -10647,7 +10731,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() @@ -10659,7 +10745,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 @@ -10724,16 +10812,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]) @@ -10766,7 +10858,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), ): @@ -10815,7 +10916,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/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 4cfc9d36714..7000497769f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7137,8 +7137,13 @@ class TestMCPServerManager: assert by_prefixed_name is not None and by_prefixed_name.description == "v2" assert manager.get_listed_tool(server, "missing") is None - def test_get_listed_tool_uses_admin_description_override_clients_saw(self): - manager = MCPServerManager() + @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", @@ -7146,14 +7151,7 @@ class TestMCPServerManager: url="http://srv", tool_name_to_description={"echo": "Admin wording"}, ) - schema = {"type": "object", "properties": {"text": {"type": "string"}}} - manager._create_prefixed_tools( - [ - MCPTool(name="echo", description="Upstream wording", inputSchema=schema), - MCPTool(name="ping", description="Untouched", inputSchema={}), - ], - server, - ) + await manager._get_tools_from_server(server, add_prefix=True) overridden = manager.get_listed_tool(server, "srv-echo") assert overridden is not None @@ -7161,6 +7159,26 @@ class TestMCPServerManager: 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) + assert [tool.description for tool in served] == ["Read a [MASKED] note"] + + listed = manager.get_listed_tool(server, "notes-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") @@ -16029,9 +16047,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() @@ -16088,7 +16104,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) @@ -16308,7 +16327,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) @@ -16370,8 +16392,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 From 7dee160b9f70cf0f0e0b8d178f246b7ec66bfb1d Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 09:08:00 +0000 Subject: [PATCH 25/47] fix(mcp): key OpenAPI listed-tool entries per caller so tools/call reads its own guarded listing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 50 ++++++++------- .../_experimental/mcp_server/operations.py | 35 ++++++++-- .../mcp/test_mcp_listed_tool_metadata.py | 25 ++++++++ .../mcp_server/test_mcp_server.py | 64 ++++++++++++++++++- 4 files changed, 143 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a247ce64bf5..6341f7972f0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1203,6 +1203,23 @@ def _server_auth_header_for( 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 _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str: """Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection. @@ -4638,15 +4655,15 @@ class MCPServerManager: invalidate_oauth_metadata_cache(server_id) 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 upstream catalog. + """Key the listed-tool cache by every request input that can change the served catalog. - Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or - exchanged as the OBO subject), the server-specific auth header, and the per-caller JWT - MCPJWTSigner mints for tools/list all reach upstream, so two callers differing in any of - them may be shown different tools. Shared servers with none of those stay on the shared - (``None``) slot. OpenAPI servers list from the process-wide registry. + The catalog is guardrail-shaped for the caller's own key (default-on guardrails, key or team + selections and opt-outs), so every keyed caller gets its own slot, on OpenAPI servers too. + Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or exchanged 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 server.spec_path or caller is None: + 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 @@ -4664,20 +4681,10 @@ class MCPServerManager: forwarded, stdio_env, caller_bearer, - per_caller=self._signs_caller_identity_upstream(server), + per_caller=auth is not None, ) return digest - @staticmethod - def _signs_caller_identity_upstream(server: MCPServer) -> bool: - from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server - get_mcp_jwt_signer, - ) - - if get_mcp_jwt_signer() is None: - return False - return server.static_headers is None or not any(k.lower() == "authorization" for k in server.static_headers) - @staticmethod def _forwarded_header_values( server: MCPServer, raw_headers: Mapping[str, str] | None @@ -6633,11 +6640,8 @@ class MCPServerManager: user_api_key_auth, mcp_auth_header, ) - listed_caller: Final = ListedToolsCaller( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=_server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header), - raw_headers=raw_headers, - oauth2_headers=oauth2_headers, + listed_caller: Final = listed_tools_caller_for( + mcp_server, user_api_key_auth, mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers ) ######################################################### diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 5a63fedce73..891c37f4fe7 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, @@ -1607,10 +1609,12 @@ async def _list_mcp_resource_templates( return managed_resource_templates -def _registered_tool_metadata(name: str, registered: RegisteredTool, server: MCPServer) -> MCPTool: - """The tool as ``tools/list`` served it (pinned, overridden, guardrail-masked) when a listing was - recorded for ``server``, else the registry entry with the admin description override applied.""" - listed: Final = global_mcp_server_manager.get_listed_tool(server, name) +def _registered_tool_metadata( + name: str, registered: RegisteredTool, server: MCPServer, caller: ListedToolsCaller +) -> MCPTool: + """The tool as ``tools/list`` served it to this caller (pinned, overridden, guardrail-masked) when a + listing was recorded, else the registry entry with the admin description override applied.""" + listed: Final = global_mcp_server_manager.get_listed_tool(server, name, caller) if listed is not None: return listed overrides: Final = server.tool_name_to_description @@ -2087,7 +2091,14 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, - tool=_registered_tool_metadata(original_tool_name, local_tool, mcp_server), + tool=_registered_tool_metadata( + original_tool_name, + local_tool, + mcp_server, + listed_tools_caller_for( + mcp_server, user_api_key_auth, mcp_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. @@ -2199,7 +2210,19 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, - tool=_registered_tool_metadata(original_tool_name, registered_local_tool, prefix_server), + tool=_registered_tool_metadata( + original_tool_name, + registered_local_tool, + prefix_server, + listed_tools_caller_for( + prefix_server, + user_api_key_auth, + mcp_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 diff --git a/tests/integration/mcp/test_mcp_listed_tool_metadata.py b/tests/integration/mcp/test_mcp_listed_tool_metadata.py index 0381e8abb00..92cf56e9899 100644 --- a/tests/integration/mcp/test_mcp_listed_tool_metadata.py +++ b/tests/integration/mcp/test_mcp_listed_tool_metadata.py @@ -167,3 +167,28 @@ def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_serve 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" + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 57268902614..070d87eb012 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -30,6 +30,7 @@ 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, @@ -8179,8 +8180,11 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl 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)], None + 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={}) @@ -8194,7 +8198,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl arguments={"petId": 1}, allowed_mcp_servers=[petstore], start_time=datetime.now(), - user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"), + user_api_key_auth=alice, ) finally: mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") @@ -8206,6 +8210,62 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl ) +@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_hands_hooks_the_metadata_of_the_operation_it_runs_when_names_collide(): """An OpenAPI operation whose name starts with its own server prefix must not be reported to the From 001fcad93076ed27b651e3221882549456a9b8fc Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 09:31:42 +0000 Subject: [PATCH 26/47] test(mcp): align listed-tool slot tests with per-caller keying Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/test_mcp_server_manager.py | 65 ++++++++++--------- 1 file changed, 36 insertions(+), 29 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 7000497769f..26d4263eac1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7007,10 +7007,8 @@ 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 = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + # Real auth: the listed-tool slot identity is hashed from these fields + user_api_key_auth = UserAPIKeyAuth(api_key="sk-test") # Mock proxy logging proxy_logging_obj = MagicMock() @@ -7037,7 +7035,9 @@ class TestMCPServerManager: assert mock_client.call_tool.await_count == 1 @staticmethod - def _manager_ready_for_call_tool(listed_tools: list[MCPTool]) -> tuple[MCPServerManager, MagicMock]: + def _manager_ready_for_call_tool( + listed_tools: list[MCPTool], caller: ListedToolsCaller | None = None + ) -> tuple[MCPServerManager, MagicMock]: from mcp.types import CallToolResult manager = MCPServerManager() @@ -7050,7 +7050,7 @@ class TestMCPServerManager: 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._create_prefixed_tools(listed_tools, server, caller=caller) mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) @@ -7064,23 +7064,23 @@ class TestMCPServerManager: return manager, proxy_logging_obj @staticmethod - def _unrestricted_auth() -> MagicMock: - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None - return user_api_key_auth + 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)] - manager, proxy_logging_obj = self._manager_ready_for_call_tool(listed) + 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=self._unrestricted_auth(), + user_api_key_auth=auth, proxy_logging_obj=proxy_logging_obj, ) @@ -7091,7 +7091,8 @@ class TestMCPServerManager: async def test_call_tool_hands_listed_tool_metadata_to_during_call_hooks_through_real_conversion(self): schema = {"type": "object", "properties": {"param": {"type": "string"}}} listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] - manager, _ = self._manager_ready_for_call_tool(listed) + 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) @@ -7100,7 +7101,7 @@ class TestMCPServerManager: server_name="test-server", name="test_tool", arguments={"param": "value"}, - user_api_key_auth=UserAPIKeyAuth(api_key="sk-test"), + user_api_key_auth=auth, proxy_logging_obj=proxy_logging_obj, ) @@ -7112,15 +7113,17 @@ class TestMCPServerManager: @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"})] + [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=self._unrestricted_auth(), + user_api_key_auth=auth, proxy_logging_obj=proxy_logging_obj, ) @@ -7244,7 +7247,9 @@ class TestMCPServerManager: caller=ListedToolsCaller(user_api_key_auth=alice), ) for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob)) - assert for_bob is not None and for_bob.description == "everyone" + 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"), @@ -7371,20 +7376,24 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), ) - caller: Final = ListedToolsCaller(mcp_auth_header="Bearer hdr") + 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( - ("signer", "static_headers", "shared"), + ("signer", "static_headers"), [ - pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), - pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), - pytest.param(None, None, True, id="no-signer-stays-shared"), + 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_jwt_signer_makes_a_shared_server_list_per_caller(self, signer, static_headers, shared): + 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 @@ -7399,12 +7408,10 @@ class TestMCPServerManager: manager._create_prefixed_tools( [MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice ) - for_bob = manager.get_listed_tool(server, "srv-turn", bob) + assert manager.get_listed_tool(server, "srv-turn", bob) is None - if shared: - assert for_bob is not None and for_bob.description == "alice view" - else: - assert for_bob is None + for_alice = manager.get_listed_tool(server, "srv-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 From de92a1103c6a466e3f4c69382be529f9954e0656 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 23:49:02 +0000 Subject: [PATCH 27/47] fix(mcp): keep oauth2 listing on the minted or signed credential, not the stored BYOK secret The listing helper that keys the per-caller catalog by the stored BYOK credential also handed that credential to the upstream client, which on an oauth2 server short-circuited the client_credentials mint and the MCPJWTSigner gate. Split the two: the catalog identity keeps the stored credential so tools/call finds the caller's slot, while an oauth2 server's tools/list sends only the per-request header, letting the M2M mint or signed JWT proceed as on main Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 19 ++++++- .../mcp_server/test_mcp_server_manager.py | 52 +++++++++++++++++++ 2 files changed, 69 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7509f62f7f4..274487af9f7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1361,6 +1361,18 @@ async def _byok_listing_auth_header( return await _get_byok_credential(mcp_server, user_api_key_auth) +def _listing_upstream_auth_header( + mcp_server: MCPServer, + mcp_auth_header: str | dict[str, str] | None, + byok_auth_header: str | dict[str, str] | None, +) -> str | dict[str, str] | None: + """The credential a tools/list sends upstream: an oauth2 server's is minted or signed, never the + stored BYOK one, the same rule ``_get_tools_from_mcp_servers`` applies to its own resolution.""" + if mcp_server.auth_type == MCPAuth.oauth2: + return mcp_auth_header + return byok_auth_header + + def _client_forwarded_authorization_headers( mcp_server: MCPServer, oauth2_headers: dict[str, str] | None, @@ -4500,6 +4512,9 @@ class MCPServerManager: client = None resolved_mcp_auth_header: Final = await _byok_listing_auth_header(server, user_api_key_auth, mcp_auth_header) + upstream_mcp_auth_header: Final = _listing_upstream_auth_header( + server, mcp_auth_header, resolved_mcp_auth_header + ) listed_caller: Final = ListedToolsCaller( user_api_key_auth=user_api_key_auth, mcp_auth_header=resolved_mcp_auth_header, @@ -4549,7 +4564,7 @@ class MCPServerManager: if ( get_mcp_jwt_signer() is not None and not has_static_authorization - and not resolved_mcp_auth_header + and not upstream_mcp_auth_header and not has_extra_authorization ): extra_headers = await inject_mcp_jwt_headers_for_upstream( @@ -4572,7 +4587,7 @@ class MCPServerManager: client = await self._create_mcp_client( server=server, - mcp_auth_header=resolved_mcp_auth_header, + mcp_auth_header=upstream_mcp_auth_header, extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index eab57b0aee4..28a1098b436 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7380,6 +7380,58 @@ class TestMCPServerManager: assert listed is not None and listed.description == "t" assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" + @pytest.mark.asyncio + async def test_oauth2_byok_listing_leaves_the_minted_token_and_signer_in_place(self): + """The stored BYOK secret keys the catalog slot tools/call reads, but never reaches the oauth2 + upstream on tools/list: the client mints its M2M token and MCPJWTSigner still signs.""" + 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", + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="csec", + token_url="http://cc1/token", + is_byok=True, + ) + 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="m2m 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) + 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() + call_side = ListedToolsCaller(user_api_key_auth=alice, mcp_auth_header="BYOK-ALICE-SECRET") + listed = manager.get_listed_tool(server, "echo", call_side) + assert listed is not None and listed.description == "m2m catalog" + @pytest.mark.parametrize( ("signer", "static_headers"), [ From 0e4de07c6510428c4437ff3a930033aded3bfe18 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 1 Oct 2026 00:40:49 +0000 Subject: [PATCH 28/47] test(mcp): oauth2 BYOK listing sends the minted token, not the stored secret, through the real proxy Integration cell for the listing fix: a client_credentials BYOK server with a stored user credential, one tools/list as that user, the peer must see a live minted bearer and one /token mint. Red at the pre-fix tip (zero mints, stored secret upstream), green at the fixed head and at the merge base Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/mcp/test_mcp_credentials.py | 57 +++++++++++++++++-- 1 file changed, 52 insertions(+), 5 deletions(-) diff --git a/tests/integration/mcp/test_mcp_credentials.py b/tests/integration/mcp/test_mcp_credentials.py index b3f5725040a..63e5d2da9ea 100644 --- a/tests/integration/mcp/test_mcp_credentials.py +++ b/tests/integration/mcp/test_mcp_credentials.py @@ -16,6 +16,7 @@ from integration._support.mcp import ( tool_calls, tool_names, ) +from integration._support.oauth_server import oauth_server ADD: Final = {"a": 2, "b": 3} STATIC_MODES: Final = ( @@ -194,6 +195,56 @@ def test_byok_server_uses_the_calling_users_stored_credential_and_fails_closed_w 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" + + 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] @@ -208,10 +259,6 @@ def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_us assert response.status_code == 200, response.text names: Final = {tool["name"] for tool in response.json()["tools"]} assert "add" in names, names - listings: Final = tuple( - item - for item in peer.drain() - if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list" - ) + listings: Final = _listings(peer) assert len(listings) == 1, listings assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr" From 9c68ca44124f6fae1c375a68c38380074141341a Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 1 Oct 2026 01:59:51 +0000 Subject: [PATCH 29/47] fix(mcp): keep the stored BYOK credential for catalog identity only on tools/list Listing used the resolved stored credential both to key the caller's catalog slot and as the upstream transport header, so REST api_key and bearer_token listings sent the user's secret instead of the server's static token and the MCPJWTSigner gate went quiet. The upstream client and the signer gate now read the caller-supplied mcp_auth_header for every auth type, exactly as before the catalog existed, and the stored credential only names the slot tools/call reads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 27 +++---------- tests/integration/mcp/test_mcp_credentials.py | 39 +++++++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 37 +++++++++++++----- 3 files changed, 72 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 274487af9f7..09714b65f86 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1346,13 +1346,12 @@ async def _resolve_byok_mcp_auth_header( return mcp_auth_header -async def _byok_listing_auth_header( +async def _byok_catalog_auth_header( mcp_server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, mcp_auth_header: str | dict[str, str] | None, ) -> str | dict[str, str] | None: - """The credential a tools/list may use: a supplied header forwards unchanged, and a missing one - falls back to the stored credential without the tool-call path's byok_auth_required raise.""" + """Keys the caller's catalog slot the way tools/call will look it up; never sent upstream.""" if not mcp_server.is_byok or mcp_auth_header is not None: return mcp_auth_header @@ -1361,18 +1360,6 @@ async def _byok_listing_auth_header( return await _get_byok_credential(mcp_server, user_api_key_auth) -def _listing_upstream_auth_header( - mcp_server: MCPServer, - mcp_auth_header: str | dict[str, str] | None, - byok_auth_header: str | dict[str, str] | None, -) -> str | dict[str, str] | None: - """The credential a tools/list sends upstream: an oauth2 server's is minted or signed, never the - stored BYOK one, the same rule ``_get_tools_from_mcp_servers`` applies to its own resolution.""" - if mcp_server.auth_type == MCPAuth.oauth2: - return mcp_auth_header - return byok_auth_header - - def _client_forwarded_authorization_headers( mcp_server: MCPServer, oauth2_headers: dict[str, str] | None, @@ -4511,13 +4498,9 @@ class MCPServerManager: verbose_logger.info("_get_tools_from_server for %s...", server.name) client = None - resolved_mcp_auth_header: Final = await _byok_listing_auth_header(server, user_api_key_auth, mcp_auth_header) - upstream_mcp_auth_header: Final = _listing_upstream_auth_header( - server, mcp_auth_header, resolved_mcp_auth_header - ) listed_caller: Final = ListedToolsCaller( user_api_key_auth=user_api_key_auth, - mcp_auth_header=resolved_mcp_auth_header, + mcp_auth_header=await _byok_catalog_auth_header(server, user_api_key_auth, mcp_auth_header), raw_headers=raw_headers, oauth2_headers=oauth2_headers, ) @@ -4564,7 +4547,7 @@ class MCPServerManager: if ( get_mcp_jwt_signer() is not None and not has_static_authorization - and not upstream_mcp_auth_header + and not mcp_auth_header and not has_extra_authorization ): extra_headers = await inject_mcp_jwt_headers_for_upstream( @@ -4587,7 +4570,7 @@ class MCPServerManager: client = await self._create_mcp_client( server=server, - mcp_auth_header=upstream_mcp_auth_header, + mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, diff --git a/tests/integration/mcp/test_mcp_credentials.py b/tests/integration/mcp/test_mcp_credentials.py index 63e5d2da9ea..e53707a38cc 100644 --- a/tests/integration/mcp/test_mcp_credentials.py +++ b/tests/integration/mcp/test_mcp_credentials.py @@ -245,6 +245,45 @@ def test_oauth2_byok_listing_sends_the_minted_token_not_the_users_stored_secret( 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] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 28a1098b436..f361b8bfd09 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7380,10 +7380,32 @@ class TestMCPServerManager: 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_oauth2_byok_listing_leaves_the_minted_token_and_signer_in_place(self): - """The stored BYOK secret keys the catalog slot tools/call reads, but never reaches the oauth2 - upstream on tools/list: the client mints its M2M token and MCPJWTSigner still signs.""" + async def test_byok_listing_keys_the_catalog_by_the_stored_secret_but_never_sends_it_upstream( + self, server_auth: dict[str, object] + ): + """The stored BYOK secret keys the catalog slot tools/call reads, but tools/list sends upstream + exactly what the caller supplied (nothing here), 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, @@ -7396,16 +7418,13 @@ class TestMCPServerManager: name="cc1", transport=MCPTransport.http, url="http://cc1", - auth_type=MCPAuth.oauth2, - client_id="cid", - client_secret="csec", - token_url="http://cc1/token", 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="m2m catalog", inputSchema={})] + 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") @@ -7430,7 +7449,7 @@ class TestMCPServerManager: signer_headers.assert_awaited_once() call_side = ListedToolsCaller(user_api_key_auth=alice, mcp_auth_header="BYOK-ALICE-SECRET") listed = manager.get_listed_tool(server, "echo", call_side) - assert listed is not None and listed.description == "m2m catalog" + assert listed is not None and listed.description == "listed catalog" @pytest.mark.parametrize( ("signer", "static_headers"), From a4a658b1f8d40f7e5a77d3c73fc53a3d2b69e50c Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 1 Oct 2026 17:55:41 +0000 Subject: [PATCH 30/47] fix(mcp): type the listed-tool metadata read from pre-call kwargs for the basedpyright gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/utils.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a914cfee03a..72a7e3ee434 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1153,6 +1153,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) @@ -1741,8 +1743,8 @@ class ProxyLogging: tool_name=kwargs.get("name", ""), arguments=kwargs.get("arguments", {}), server_name=kwargs.get("server_name"), - tool_description=kwargs.get("tool_description"), - tool_input_schema=kwargs.get("tool_input_schema"), + 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(), ) From 12947900b68a683ae1da322dd9f613769def414d Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 02:09:53 -0700 Subject: [PATCH 31/47] fix(mcp): hand never-listed tools/call hooks name and arguments only The local-registry call path fell back to the registry entry with the admin description override when no tools/list had been recorded for the caller, so a pre_mcp_call guardrail scanned a description the caller was never served and blocked OpenAPI calls that passed before, and base's own selected-guardrail REST test failed on the two-text redaction. _registered_tool_metadata now returns the listed entry or None, so a tools/call with no prior listing sends name and arguments only as promised, and that REST test double goes back to its base shape --- .../_experimental/mcp_server/operations.py | 18 +-- .../test_mcp_server_tool_calls_and_headers.py | 103 +++++++----------- .../mcp_server/test_rest_endpoints.py | 3 +- 3 files changed, 42 insertions(+), 82 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 891c37f4fe7..33b61ad310e 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -141,7 +141,6 @@ from litellm.types.mcp import ( without_header, ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer -from litellm.types.mcp_server.tool_registry import MCPTool as RegisteredTool from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall from litellm.utils import Rules, client, function_setup @@ -1609,17 +1608,10 @@ async def _list_mcp_resource_templates( return managed_resource_templates -def _registered_tool_metadata( - name: str, registered: RegisteredTool, server: MCPServer, caller: ListedToolsCaller -) -> MCPTool: - """The tool as ``tools/list`` served it to this caller (pinned, overridden, guardrail-masked) when a - listing was recorded, else the registry entry with the admin description override applied.""" - listed: Final = global_mcp_server_manager.get_listed_tool(server, name, caller) - if listed is not None: - return listed - overrides: Final = server.tool_name_to_description - description: Final = overrides.get(name, registered.description) if overrides else registered.description - return MCPTool(name=name, description=description, input_schema=registered.input_schema) +def _registered_tool_metadata(name: str, server: MCPServer, caller: ListedToolsCaller) -> MCPTool | None: + """The tool as ``tools/list`` served it to this caller (pinned, overridden, guardrail-masked), or None when + no listing was recorded so the call hands the hooks name and arguments only.""" + return global_mcp_server_manager.get_listed_tool(server, name, caller) def _resolve_display_name_to_original( @@ -2093,7 +2085,6 @@ async def _execute_mcp_tool( guardrail_context=guardrail_context, tool=_registered_tool_metadata( original_tool_name, - local_tool, mcp_server, listed_tools_caller_for( mcp_server, user_api_key_auth, mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers @@ -2212,7 +2203,6 @@ async def _execute_mcp_tool( guardrail_context=guardrail_context, tool=_registered_tool_metadata( original_tool_name, - registered_local_tool, prefix_server, listed_tools_caller_for( prefix_server, 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 afe0ef99419..fcf600d7b98 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 @@ -8070,10 +8070,13 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): @pytest.mark.asyncio -async def test_execute_mcp_tool_hands_openapi_registered_tool_metadata_to_pre_call_hooks(): - """OpenAPI-generated tools dispatch through the local registry, so the pre-call hooks must get the - registered description and input schema on that path too, even when no tools/list ran first.""" +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", @@ -8082,6 +8085,7 @@ async def test_execute_mcp_tool_hands_openapi_registered_tool_metadata_to_pre_ca 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( @@ -8089,71 +8093,42 @@ async def test_execute_mcp_tool_hands_openapi_registered_tool_metadata_to_pre_ca ) manager = mcp_module.global_mcp_server_manager manager._listed_tools_by_server_id.pop(petstore.server_id, None) - pre_call_tool_check = AsyncMock(return_value={}) + 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), ): - 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=UserAPIKeyAuth(api_key="sk-user", user_id="alice"), + 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) - handed_tool = pre_call_tool_check.call_args.kwargs["tool"] - assert (handed_tool.name, handed_tool.description, handed_tool.input_schema) == ( - "list_pets", - "List the pets", - schema, - ) - - -@pytest.mark.asyncio -async def test_execute_mcp_tool_hands_openapi_hooks_the_admin_description_clients_saw(): - """tools/list shows the admin's tool_name_to_description wording, so the local-registry call path - must hand the pre-call hooks that same wording rather than the generated 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": "ADMIN DESC"}, - ) - 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=schema, handler=lambda petId: "ok" - ) - manager = mcp_module.global_mcp_server_manager - manager._listed_tools_by_server_id.pop(petstore.server_id, None) - 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=UserAPIKeyAuth(api_key="sk-user", user_id="alice"), - ) - finally: - mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") - - handed_tool = pre_call_tool_check.call_args.kwargs["tool"] - assert (handed_tool.description, handed_tool.input_schema) == ("ADMIN DESC", schema) + 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 @@ -8267,9 +8242,9 @@ async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entr @pytest.mark.asyncio -async def test_execute_mcp_tool_hands_hooks_the_metadata_of_the_operation_it_runs_when_names_collide(): - """An OpenAPI operation whose name starts with its own server prefix must not be reported to the - pre-call hooks with the metadata of the shorter operation, since that is not the one that runs.""" +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( @@ -8306,11 +8281,7 @@ async def test_execute_mcp_tool_hands_hooks_the_metadata_of_the_operation_it_run finally: registry.unregister_tools_with_prefix("petstore-") - handed_tool = pre_call_tool_check.call_args.kwargs["tool"] - assert (handed_tool.description, handed_tool.input_schema) == ( - "long", - {"type": "object", "properties": {"petId": {"type": "integer"}}}, - ) + assert pre_call_tool_check.call_args.kwargs["tool"] is None assert result.content[0].text == "long" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index 1cf23e0b49b..759014b54c5 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -3326,8 +3326,7 @@ async def test_request_selected_tool_specific_guardrail_applies_to_virtual_execu default_on=False, custom_code="def apply_guardrail(inputs, request_data, input_type):\n" ' if inputs.get("tools", [{}])[0].get("function", {}).get("name") == "execute":\n' - ' texts = [t.replace("confidential", "redacted") for t in inputs.get("texts", [])]\n' - f' return {{"action": "{action}", "reason": "resolved tool blocked", "texts": texts}}\n' + f' return {{"action": "{action}", "reason": "resolved tool blocked", "texts": ["redacted"]}}\n' " return allow()\n", ) manager: Final = mcp_server_manager.MCPServerManager() From e9b8c0fd5c63df6b310cd61a18d586ed27eb969f Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 02:15:16 -0700 Subject: [PATCH 32/47] fix(mcp): keep during_mcp_call hooks on name and arguments only call_tool handed the caller's listed entry to the during-hook task as well, so during_mcp_call guardrails scanned the description line and schema leaves of any listed tool after the upstream call had already run, blocking calls that passed before whenever the policy matched the description, returned a fixed-length texts list, or hit the depth guard on a deep schema. The listed entry is only disclosed for pre_mcp_call, so the during task no longer receives it and its request object carries no description or schema, as before --- .../_experimental/mcp_server/mcp_server_manager.py | 1 - .../mcp_server/test_mcp_server_manager.py | 11 ++++++----- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 09714b65f86..de5281998ce 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -6694,7 +6694,6 @@ class MCPServerManager: start_time=start_time, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, - tool=self.get_listed_tool(mcp_server, name, listed_caller), ) tasks.append(during_hook_task) 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 f361b8bfd09..8fd931c8b7c 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 @@ -7085,7 +7085,9 @@ class TestMCPServerManager: assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ("Runs the test tool", schema) @pytest.mark.asyncio - async def test_call_tool_hands_listed_tool_metadata_to_during_call_hooks_through_real_conversion(self): + 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") @@ -7103,10 +7105,9 @@ class TestMCPServerManager: ) during_data = proxy_logging_obj.during_call_hook.call_args.kwargs["data"] - assert (during_data["mcp_tool_description"], during_data["mcp_input_schema"]) == ( - "Runs the test tool", - schema, - ) + 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): From 4b0b326dc2d036b9ca71e2db05dcd69ade3cd896 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 02:37:50 -0700 Subject: [PATCH 33/47] fix(mcp): key the BYOK catalog slot by the client's header, not the stored credential tools/list resolved the stored BYOK credential to pick the caller's catalog slot, which read the credential store before the classified try block. With Postgres down and a cold per-worker cache that made every REST tools/list on an is_byok server fail with tools=[] and no upstream call, and the read seeded the per-worker cache (including a negative entry), so a tools/call on another worker after a store, rotate or revoke on this one kept using the stale value. The slot is now keyed by what the client supplied plus the caller's hashed key, on both sides. _get_tools_from_server and call_tool take a keyword-only catalog_auth_header that defaults to mcp_auth_header as received (the default is the builtin Ellipsis so it survives a module reload). The /mcp fan-out and execute_mcp_tool, which swap the resolved credential into mcp_auth_header, pass the client's value explicitly. What goes upstream is unchanged. _byok_catalog_auth_header is gone. --- .../mcp_server/mcp_server_manager.py | 31 ++++--- .../_experimental/mcp_server/operations.py | 20 +++- .../mcp_server/test_mcp_server.py | 1 + .../mcp_server/test_mcp_server_manager.py | 92 ++++++++++++++++--- 4 files changed, 116 insertions(+), 28 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index de5281998ce..f7638846a81 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 @@ -1346,18 +1346,14 @@ async def _resolve_byok_mcp_auth_header( return mcp_auth_header -async def _byok_catalog_auth_header( - mcp_server: MCPServer, - user_api_key_auth: UserAPIKeyAuth | None, +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: - """Keys the caller's catalog slot the way tools/call will look it up; never sent upstream.""" - if not mcp_server.is_byok or mcp_auth_header is not None: - return mcp_auth_header - - from litellm.proxy._experimental.mcp_server.operations import _get_byok_credential - - return await _get_byok_credential(mcp_server, user_api_key_auth) + """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( @@ -4479,6 +4475,8 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None = None, client_ip: str | None = None, proxy_logging_obj: ProxyLogging | None = None, + *, + catalog_auth_header: str | dict[str, str] | None | EllipsisType = ..., ) -> Sequence[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -4486,6 +4484,8 @@ class MCPServerManager: 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`` Returns: List[MCPTool]: List of tools available on the server with prefixed names @@ -4500,7 +4500,7 @@ class MCPServerManager: client = None listed_caller: Final = ListedToolsCaller( user_api_key_auth=user_api_key_auth, - mcp_auth_header=await _byok_catalog_auth_header(server, user_api_key_auth, mcp_auth_header), + mcp_auth_header=_catalog_auth_header(mcp_auth_header, catalog_auth_header), raw_headers=raw_headers, oauth2_headers=oauth2_headers, ) @@ -6627,6 +6627,8 @@ 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 = ..., ) -> CallToolResult | InputRequiredResult: """ Call a tool with the given name and arguments @@ -6638,6 +6640,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 @@ -6649,6 +6653,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) @@ -6659,7 +6664,7 @@ class MCPServerManager: mcp_auth_header, ) listed_caller: Final = listed_tools_caller_for( - mcp_server, user_api_key_auth, mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers + mcp_server, user_api_key_auth, client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers ) ######################################################### diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 33b61ad310e..be2ca076de8 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1115,6 +1115,7 @@ 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) @@ -1131,6 +1132,7 @@ 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, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -2013,6 +2015,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: @@ -2087,7 +2090,12 @@ async def _execute_mcp_tool( original_tool_name, mcp_server, listed_tools_caller_for( - mcp_server, user_api_key_auth, mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers + mcp_server, + user_api_key_auth, + client_auth_header, + mcp_server_auth_headers, + raw_headers, + oauth2_headers, ), ), ) @@ -2139,6 +2147,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, @@ -2207,7 +2216,7 @@ async def _execute_mcp_tool( listed_tools_caller_for( prefix_server, user_api_key_auth, - mcp_auth_header, + client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers, @@ -2615,8 +2624,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 @@ -2626,6 +2639,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, 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 f8bf72428aa..00ea7c0c78c 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,7 @@ async def test_get_tools_from_mcp_servers(): user_api_key_auth=None, oauth2_headers=None, proxy_logging_obj=None, + catalog_auth_header=None, ): if server.server_id == "server1_id": return [mock_tool_1] 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 8fd931c8b7c..8756ec421ae 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 @@ -53,6 +53,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, @@ -7322,7 +7323,54 @@ class TestMCPServerManager: assert listed is not None and listed.description == "everyone" @pytest.mark.asyncio - async def test_byok_stored_credential_lists_into_the_slot_tools_call_reads(self): + 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) + + 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, @@ -7337,20 +7385,40 @@ class TestMCPServerManager: url="http://byok-catalog", is_byok=True, ) + manager.registry = {"byok-catalog": server} user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-user") - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + 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, user_api_key_auth=user) + await manager._get_tools_from_server(server=server, mcp_auth_header=list_header, user_api_key_auth=user) + 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")) - call_side = ListedToolsCaller(user_api_key_auth=user, mcp_auth_header="stored-secret") - listed = manager.get_listed_tool(server, "turn", call_side) - assert listed is not None and listed.description == "stored cred 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): @@ -7401,12 +7469,13 @@ class TestMCPServerManager: ], ) @pytest.mark.asyncio - async def test_byok_listing_keys_the_catalog_by_the_stored_secret_but_never_sends_it_upstream( + async def test_byok_listing_keys_the_catalog_by_the_caller_and_never_touches_the_stored_secret( self, server_auth: dict[str, object] ): - """The stored BYOK secret keys the catalog slot tools/call reads, but tools/list sends upstream - exactly what the caller supplied (nothing here), so the static token, the M2M mint and - MCPJWTSigner all behave as they did before the catalog existed, whatever the auth_type.""" + """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, @@ -7448,8 +7517,7 @@ class TestMCPServerManager: assert client_kwargs["mcp_auth_header"] is None, client_kwargs assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"} signer_headers.assert_awaited_once() - call_side = ListedToolsCaller(user_api_key_auth=alice, mcp_auth_header="BYOK-ALICE-SECRET") - listed = manager.get_listed_tool(server, "echo", call_side) + 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( From 9f46ec11b295fbc8b353145bba7fb92a5e439922 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 02:43:31 -0700 Subject: [PATCH 34/47] fix(mcp): drop a listed catalog recorded across a server save _record_listed_tools ran after the awaited upstream fetch, so a PUT /v1/mcp/server that landed mid-fetch had its invalidation undone when the fetch completed: hooks then saw the pre-save description next to the post-save definition until the next listing, instead of name and arguments only. The manager now keeps a per-server listed-tools generation, bumped by _invalidate_server_definition_caches. _get_tools_from_server reads it before the fetch and _record_listed_tools skips the write when it moved; the next listing records normally. --- .../mcp_server/mcp_server_manager.py | 23 +++++++++-- .../mcp_server/test_mcp_server_manager.py | 38 +++++++++++++++++++ 2 files changed, 57 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f7638846a81..eff17696cb0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2035,6 +2035,7 @@ class MCPServerManager: } """ self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list + self._listed_tools_generations: Mapping[str, int] = MappingProxyType({}) 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 @@ -4504,6 +4505,7 @@ class MCPServerManager: 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 — @@ -4602,7 +4604,7 @@ class MCPServerManager: # 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) + self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation) if not add_prefix: return unprefixed_tools return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] @@ -4618,7 +4620,7 @@ class MCPServerManager: raw_headers=raw_headers, ) prefixed_or_original_tools: Final = self._create_prefixed_tools( - guarded_tools, server, add_prefix=add_prefix, caller=listed_caller + guarded_tools, server, add_prefix=add_prefix, caller=listed_caller, generation=listed_generation ) return prefixed_or_original_tools @@ -4670,6 +4672,9 @@ class MCPServerManager: self._invalidate_discovery_lists(server_id) self._listed_tools_by_server_id.pop(server_id, None) + self._listed_tools_generations = MappingProxyType( + {**self._listed_tools_generations, server_id: self._listed_tools_generations.get(server_id, 0) + 1} + ) invalidate_oauth_metadata_cache(server_id) def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None: @@ -4715,8 +4720,17 @@ class MCPServerManager: ) def _record_listed_tools( - self, server: MCPServer, tools: Sequence[MCPTool], caller: ListedToolsCaller | None + self, + server: MCPServer, + tools: Sequence[MCPTool], + caller: ListedToolsCaller | None, + generation: int | None = None, ) -> None: + """Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation + read before the listing's upstream fetch; a server save that landed mid-fetch moved it, and the + pre-save catalog is then dropped rather than written over the invalidation.""" + 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({})) @@ -5638,6 +5652,7 @@ class MCPServerManager: server: MCPServer, add_prefix: bool = True, caller: ListedToolsCaller | None = None, + generation: int | None = None, ) -> list[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -5663,7 +5678,7 @@ class MCPServerManager: for spelling in iter_known_tool_name_spellings(original_name, server): self.tool_name_to_mcp_server_name_mapping[spelling] = prefix - self._record_listed_tools(server, tools, caller) + self._record_listed_tools(server, tools, caller, generation) verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name) return prefixed_tools 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 8756ec421ae..e79c855b21b 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 @@ -7194,6 +7194,44 @@ class TestMCPServerManager: 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) + + 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) + listed = manager.get_listed_tool(server, "turn", caller) + assert listed is not None and listed.description == "after save" + @pytest.mark.asyncio async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager() From 71be87b893ec03e3b7d14952788586ed1de73529 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 03:57:06 -0700 Subject: [PATCH 35/47] fix(mcp): drop the catalog again once a saved OpenAPI server's registry is rebuilt add_server and update_server publish the saved definition before the OpenAPI registry entries are rebuilt from the spec, so a listing recorded during that fetch held the pre-save entries under the new generation. The generation is bumped a second time after the registry refresh. The during-hook task no longer accepts a listed entry, the one-line wrapper over get_listed_tool is inlined at its two call sites, and the per-server generation map is a plain dict. --- .../mcp_server/mcp_server_manager.py | 16 ++++----- .../_experimental/mcp_server/operations.py | 15 +++----- .../mcp_server/test_mcp_server_manager.py | 35 +++++++++++++++++++ 3 files changed, 46 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index eff17696cb0..20d6128d3fb 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2035,7 +2035,7 @@ class MCPServerManager: } """ self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list - self._listed_tools_generations: Mapping[str, int] = MappingProxyType({}) + 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 @@ -3351,6 +3351,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._invalidate_server_definition_caches(mcp_server.server_id) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Added MCP Server: %s", new_server.name) @@ -3388,6 +3390,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._invalidate_server_definition_caches(mcp_server.server_id) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Updated MCP Server: %s", new_server.name) @@ -4672,9 +4676,7 @@ class MCPServerManager: self._invalidate_discovery_lists(server_id) self._listed_tools_by_server_id.pop(server_id, None) - self._listed_tools_generations = MappingProxyType( - {**self._listed_tools_generations, server_id: self._listed_tools_generations.get(server_id, 0) + 1} - ) + self._listed_tools_generations[server_id] = self._listed_tools_generations.get(server_id, 0) + 1 invalidate_oauth_metadata_cache(server_id) def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None: @@ -4727,8 +4729,7 @@ class MCPServerManager: generation: int | None = None, ) -> None: """Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation - read before the listing's upstream fetch; a server save that landed mid-fetch moved it, and the - pre-save catalog is then dropped rather than written over the invalidation.""" + read before the listing's upstream fetch; the record is skipped when it no longer matches.""" 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) @@ -6059,7 +6060,6 @@ class MCPServerManager: start_time: datetime.datetime, litellm_logging_obj: "LiteLLMLoggingObj | None" = None, guardrail_context: Mapping[str, object] | None = None, - tool: MCPTool | None = None, ): """Create and return a during hook task for MCP tool calls. @@ -6074,8 +6074,6 @@ class MCPServerManager: tool_name=name, arguments=arguments, server_name=server_name_from_prefix, - tool_description=tool.description if tool is not None else None, - tool_input_schema=tool.input_schema if tool is not None else None, start_time=start_time.timestamp() if start_time else None, hidden_params=HiddenParams(), ) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index be2ca076de8..151e3a87939 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -81,7 +81,6 @@ 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, @@ -1610,12 +1609,6 @@ async def _list_mcp_resource_templates( return managed_resource_templates -def _registered_tool_metadata(name: str, server: MCPServer, caller: ListedToolsCaller) -> MCPTool | None: - """The tool as ``tools/list`` served it to this caller (pinned, overridden, guardrail-masked), or None when - no listing was recorded so the call hands the hooks name and arguments only.""" - return global_mcp_server_manager.get_listed_tool(server, name, caller) - - def _resolve_display_name_to_original( name: str, allowed_mcp_servers: list[MCPServer], @@ -2086,9 +2079,9 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, - tool=_registered_tool_metadata( - original_tool_name, + tool=global_mcp_server_manager.get_listed_tool( mcp_server, + original_tool_name, listed_tools_caller_for( mcp_server, user_api_key_auth, @@ -2210,9 +2203,9 @@ async def _execute_mcp_tool( raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, - tool=_registered_tool_metadata( - original_tool_name, + tool=global_mcp_server_manager.get_listed_tool( prefix_server, + original_tool_name, listed_tools_caller_for( prefix_server, user_api_key_auth, 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 e79c855b21b..3ad7652fb8c 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 @@ -7232,6 +7232,41 @@ class TestMCPServerManager: 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 async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager() From 46639297ad7b52c4a2f2c5fb3704ac78864b1d07 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 12:52:57 -0700 Subject: [PATCH 36/47] fix(mcp): keep discovery and OAuth metadata caches across an OpenAPI spec re-read add_server and update_server ran the full server-definition invalidation a second time after the awaited OpenAPI spec fetch, which also dropped the prompts/resources/templates discovery entries and the OAuth protected-resource metadata filled under the already-published definition, so the next request went upstream again. Only the listed-tool catalog recorded during the fetch holds pre-save entries, so the post-fetch pass now drops just that catalog and bumps its generation via the new _drop_listed_tools helper, which the full invalidation also calls. --- .../mcp_server/mcp_server_manager.py | 9 ++- .../mcp_server/test_mcp_server_manager.py | 58 +++++++++++++++++++ 2 files changed, 64 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 20d6128d3fb..ce9a5f62bf2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3352,7 +3352,7 @@ class MCPServerManager: self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) if new_server.spec_path: - self._invalidate_server_definition_caches(mcp_server.server_id) + 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) @@ -3391,7 +3391,7 @@ class MCPServerManager: self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) if new_server.spec_path: - self._invalidate_server_definition_caches(mcp_server.server_id) + 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) @@ -4675,9 +4675,12 @@ 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 - invalidate_oauth_metadata_cache(server_id) 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. 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 3ad7652fb8c..632e3415245 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 @@ -5,6 +5,7 @@ import json import logging import os import sys +import time from collections.abc import AsyncIterator from datetime import datetime from pathlib import Path @@ -40,6 +41,7 @@ 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, @@ -7267,6 +7269,62 @@ class TestMCPServerManager: 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() From ed75c1aef7416df978ffe5074ae35ac0ef60077f Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 12:56:38 -0700 Subject: [PATCH 37/47] fix(mcp): look a called tool up in the listed catalog by its bare name only get_listed_tool stripped the server prefix a second time when the exact name was absent from the caller's listing, so a never-listed upstream tool whose bare name starts with the server prefix resolved to the listed sibling and that sibling's description and input schema reached the pre-call hooks for a call to a different tool. Every caller already passes the once-stripped bare name, so the lookup is now exact. Tests that looked the catalog up by a prefixed name now use the bare name the callers pass; two new tests pin the never-listed sibling case at the manager and at the tools/call path. --- .../mcp_server/mcp_server_manager.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 66 ++++++++++++------- .../test_mcp_server_tool_calls_and_headers.py | 47 +++++++++++++ 3 files changed, 90 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index ce9a5f62bf2..f37916cad3f 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5691,7 +5691,7 @@ class MCPServerManager: listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity) if not listed: return None - return listed.get(name) or listed.get(strip_known_server_prefix(name, server)) + return listed.get(name) def _create_prefixed_prompts( self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True 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 632e3415245..2b4b025f97b 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 @@ -7131,16 +7131,35 @@ class TestMCPServerManager: 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_prefixed_name_and_latest_listing(self): + 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._create_prefixed_tools([MCPTool(name="echo", description="v1", inputSchema={})], server) manager._create_prefixed_tools([MCPTool(name="echo", description="v2", inputSchema={})], server) - by_prefixed_name = manager.get_listed_tool(server, "srv-echo") - assert by_prefixed_name is not None and by_prefixed_name.description == "v2" + 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._create_prefixed_tools( + [ + MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}), + MCPTool(name="bar", description="Fetches bar records", inputSchema={"type": "object"}), + ], + server, + caller=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"}}} @@ -7157,7 +7176,7 @@ class TestMCPServerManager: ) await manager._get_tools_from_server(server, add_prefix=True) - overridden = manager.get_listed_tool(server, "srv-echo") + 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") @@ -7178,7 +7197,7 @@ class TestMCPServerManager: served = await manager._get_tools_from_server(server, add_prefix=True, proxy_logging_obj=proxy_logging_obj) assert [tool.description for tool in served] == ["Read a [MASKED] note"] - listed = manager.get_listed_tool(server, "notes-read_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" ) @@ -7360,15 +7379,15 @@ class TestMCPServerManager: caller=ListedToolsCaller(user_api_key_auth=bob), ) - alice_tool = manager.get_listed_tool(server, "srv-read", ListedToolsCaller(user_api_key_auth=alice)) - bob_tool = manager.get_listed_tool(server, "srv-read", 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, "srv-read", carol) is None + assert manager.get_listed_tool(server, "read", carol) is None shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") manager._create_prefixed_tools( @@ -7434,11 +7453,11 @@ class TestMCPServerManager: [MCPTool(name="turn", description="Catalog B", inputSchema={})], server, caller=caller_b ) - for_a = manager.get_listed_tool(server, "srv-turn", caller_a) - for_b = manager.get_listed_tool(server, "srv-turn", 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, "srv-turn", ListedToolsCaller()) is None + assert manager.get_listed_tool(server, "turn", ListedToolsCaller()) is None def test_shared_server_ignores_headers_it_never_forwards(self): manager = MCPServerManager() @@ -7676,9 +7695,9 @@ class TestMCPServerManager: manager._create_prefixed_tools( [MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice ) - assert manager.get_listed_tool(server, "srv-turn", bob) is None + assert manager.get_listed_tool(server, "turn", bob) is None - for_alice = manager.get_listed_tool(server, "srv-turn", alice) + 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): @@ -7696,10 +7715,10 @@ class TestMCPServerManager: manager._create_prefixed_tools( [MCPTool(name="turn", description="slot a", inputSchema={})], server, caller=alice ) - assert manager.get_listed_tool(server, "srv-turn", bob) is None + 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, "srv-turn", same_key) + listed = manager.get_listed_tool(server, "turn", same_key) assert listed is not None and listed.description == "slot a" @@ -7746,7 +7765,7 @@ class TestMCPServerManager: proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) await manager.call_tool( server_name="catalog", - name="catalog-turn", + 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, @@ -7785,13 +7804,13 @@ class TestMCPServerManager: [MCPTool(name="read", description="u1 again", inputSchema={})], server, caller=callers[1] ) - assert manager.get_listed_tool(server, "srv-read", callers[0]) is None - second = manager.get_listed_tool(server, "srv-read", 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, "srv-read", callers[-1]) + 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, "srv-read") + shared = manager.get_listed_tool(server, "read") assert shared is not None and shared.description == "shared" @pytest.mark.asyncio @@ -7826,10 +7845,9 @@ class TestMCPServerManager: 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"] - for name in ("list_pets", "petstore-list_pets"): - tool = manager.get_listed_tool(server, name) - assert tool is not None and tool.description == "List pets" - assert tool.input_schema["properties"] == {"limit": {"type": "integer"}} + 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): 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 fcf600d7b98..5aec9a9186f 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 @@ -8285,6 +8285,53 @@ async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_ho 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_rest_unresolved_prefixed_name_routes_to_requested_server(): """A prefixed REST name that resolves to no tool must still dispatch to the server_id. From ad746dadb7d4019966dae970a05a6de9da944a91 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 14:22:05 -0700 Subject: [PATCH 38/47] fix(mcp): record a listed-tool catalog only for a listing the caller is served _get_tools_from_server now records the catalog into the caller's listed-tools slot only when asked (record_listing=True), which the served listings pass: the /mcp and Responses API tools/list handlers via _get_tools_from_mcp_servers, MCPServerManager.list_tools, and the REST listing via _list_server_tools. Four internal listings stop recording, so a later tools/call hands pre_mcp_call hooks name and arguments only, as on main: - _list_tools_before_first_call, the implicit listing inside tools/call when this worker does not yet expose the tool - fetch_pinnable_tool_catalog, the admin pin snapshot listed without the catalog guard and without description overrides - _initialize_tool_name_to_mcp_server_name_mapping, the startup fill - get_tools_for_server, used by the semantic tool filter _create_prefixed_tools returns to its tool-name mapping job only; the record follows it in _get_tools_from_server. --- .../mcp_server/mcp_server_manager.py | 14 +- .../_experimental/mcp_server/operations.py | 6 + .../mcp_server/rest_endpoints.py | 13 +- .../mcp_server/test_mcp_server.py | 2 + .../mcp_server/test_mcp_server_manager.py | 166 +++++++++++++----- .../test_mcp_server_tool_calls_and_headers.py | 81 ++++++++- 6 files changed, 227 insertions(+), 55 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f37916cad3f..2d13dc0dcf3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3870,6 +3870,7 @@ class MCPServerManager: 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: @@ -4482,6 +4483,7 @@ class MCPServerManager: proxy_logging_obj: ProxyLogging | None = None, *, 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. @@ -4491,6 +4493,8 @@ class MCPServerManager: 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 @@ -4608,7 +4612,8 @@ class MCPServerManager: # 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) + if record_listing: + self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation) if not add_prefix: return unprefixed_tools return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] @@ -4624,8 +4629,10 @@ class MCPServerManager: raw_headers=raw_headers, ) prefixed_or_original_tools: Final = self._create_prefixed_tools( - guarded_tools, server, add_prefix=add_prefix, caller=listed_caller, generation=listed_generation + guarded_tools, server, add_prefix=add_prefix ) + if record_listing: + self._record_listed_tools(server, guarded_tools, listed_caller, listed_generation) return prefixed_or_original_tools @@ -5655,8 +5662,6 @@ class MCPServerManager: tools: Sequence[MCPTool], server: MCPServer, add_prefix: bool = True, - caller: ListedToolsCaller | None = None, - generation: int | None = None, ) -> list[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -5682,7 +5687,6 @@ class MCPServerManager: for spelling in iter_known_tool_name_spellings(original_name, server): self.tool_name_to_mcp_server_name_mapping[spelling] = prefix - self._record_listed_tools(server, tools, caller, generation) verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name) return prefixed_tools diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 151e3a87939..d8886e26da7 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -957,6 +957,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 = True, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -967,6 +969,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 @@ -1132,6 +1136,7 @@ async def _get_tools_from_mcp_servers( oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, catalog_auth_header=catalog_auth_header, + record_listing=record_listing, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1805,6 +1810,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) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 12cdab59e0f..b9eb17b1758 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -703,6 +703,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, @@ -713,6 +715,7 @@ 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( @@ -734,7 +737,14 @@ 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 + server, + server_auth_header, + raw_headers, + user_api_key_auth, + extra_headers, + client_ip, + proxy_logging_obj, + record_listing=True, ) if not apply_tool_filters: @@ -776,6 +786,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/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index 00ea7c0c78c..441cc57c292 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -964,6 +964,7 @@ async def test_get_tools_from_mcp_servers(): oauth2_headers=None, proxy_logging_obj=None, catalog_auth_header=None, + record_listing=True, ): if server.server_id == "server1_id": return [mock_tool_1] @@ -1999,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=True, ) # 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 2b4b025f97b..b0b40982ade 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 @@ -7050,7 +7050,8 @@ class TestMCPServerManager: 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, caller=caller) + 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) @@ -7134,8 +7135,8 @@ class TestMCPServerManager: 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._create_prefixed_tools([MCPTool(name="echo", description="v1", inputSchema={})], server) - manager._create_prefixed_tools([MCPTool(name="echo", description="v2", inputSchema={})], server) + 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" @@ -7147,13 +7148,13 @@ class TestMCPServerManager: 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._create_prefixed_tools( + 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"}), ], - server, - caller=caller, + caller, ) assert manager.get_listed_tool(server, "srv-foo", caller) is None @@ -7174,7 +7175,7 @@ class TestMCPServerManager: url="http://srv", tool_name_to_description={"echo": "Admin wording"}, ) - await manager._get_tools_from_server(server, add_prefix=True) + 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 @@ -7194,7 +7195,9 @@ class TestMCPServerManager: 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) + 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") @@ -7206,8 +7209,8 @@ class TestMCPServerManager: 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._create_prefixed_tools([MCPTool(name="echo", description="old", inputSchema={})], server) - manager._create_prefixed_tools([MCPTool(name="ping", description="kept", inputSchema={})], 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) @@ -7236,7 +7239,7 @@ class TestMCPServerManager: 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) + 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() @@ -7249,7 +7252,7 @@ class TestMCPServerManager: 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) + 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" @@ -7348,7 +7351,7 @@ class TestMCPServerManager: 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._create_prefixed_tools([MCPTool(name="echo", description="shared", inputSchema={})], server) + manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) await manager.invalidate_user_oauth_token_cache("alice", server.server_id) @@ -7368,15 +7371,15 @@ class TestMCPServerManager: 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._create_prefixed_tools( + manager._record_listed_tools( + server, [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], - server, - caller=ListedToolsCaller(user_api_key_auth=alice), + ListedToolsCaller(user_api_key_auth=alice), ) - manager._create_prefixed_tools( - [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], + manager._record_listed_tools( server, - caller=ListedToolsCaller(user_api_key_auth=bob), + [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)) @@ -7390,10 +7393,10 @@ class TestMCPServerManager: assert manager.get_listed_tool(server, "read", carol) is None shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") - manager._create_prefixed_tools( - [MCPTool(name="echo", description="everyone", inputSchema={})], + manager._record_listed_tools( shared, - caller=ListedToolsCaller(user_api_key_auth=alice), + [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" @@ -7446,12 +7449,8 @@ class TestMCPServerManager: server = MCPServer( **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} ) - manager._create_prefixed_tools( - [MCPTool(name="turn", description="Catalog A", inputSchema={})], server, caller=caller_a - ) - manager._create_prefixed_tools( - [MCPTool(name="turn", description="Catalog B", inputSchema={})], server, caller=caller_b - ) + 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) @@ -7462,10 +7461,10 @@ class TestMCPServerManager: 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._create_prefixed_tools( - [MCPTool(name="turn", description="everyone", inputSchema={})], + manager._record_listed_tools( server, - caller=ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), + [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"}) @@ -7499,7 +7498,7 @@ class TestMCPServerManager: AsyncMock(side_effect=RuntimeError("DB DOWN")), ), ): - await manager._get_tools_from_server(server=server, user_api_key_auth=user) + 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" @@ -7550,7 +7549,9 @@ class TestMCPServerManager: 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) + 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) ) @@ -7590,6 +7591,7 @@ class TestMCPServerManager: server=server, mcp_auth_header="Bearer hdr", user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), + record_listing=True, ) caller: Final = ListedToolsCaller( @@ -7659,7 +7661,7 @@ class TestMCPServerManager: signer_headers, ), ): - await manager._get_tools_from_server(server=server, user_api_key_auth=alice) + 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")) @@ -7692,8 +7694,8 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=signer, ): - manager._create_prefixed_tools( - [MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice + manager._record_listed_tools( + server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice ) assert manager.get_listed_tool(server, "turn", bob) is None @@ -7712,9 +7714,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=MagicMock(), ): - manager._create_prefixed_tools( - [MCPTool(name="turn", description="slot a", inputSchema={})], server, caller=alice - ) + 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")) @@ -7756,6 +7756,7 @@ class TestMCPServerManager: 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() @@ -7789,20 +7790,18 @@ class TestMCPServerManager: url="http://srv", auth_type=MCPAuth.oauth2_token_exchange, ) - manager._create_prefixed_tools([MCPTool(name="read", description="shared", inputSchema={})], server) + 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._create_prefixed_tools( - [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], + manager._record_listed_tools( server, - caller=caller, + [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], + caller, ) - manager._create_prefixed_tools( - [MCPTool(name="read", description="u1 again", inputSchema={})], server, caller=callers[1] - ) + 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]) @@ -7840,7 +7839,7 @@ class TestMCPServerManager: handler=_handler, ) try: - listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix) + 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-") @@ -7882,7 +7881,7 @@ class TestMCPServerManager: handler=_handler, ) try: - listed = await manager._get_tools_from_server(server=server, add_prefix=True) + 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) @@ -7892,6 +7891,77 @@ class TestMCPServerManager: 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): """ 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 5aec9a9186f..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 @@ -37,7 +37,7 @@ from litellm.proxy._types import ( 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(): @@ -8332,6 +8332,85 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation 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. From 021952155975fd80ef96edfb40cf548a28587510 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 15:49:47 -0700 Subject: [PATCH 39/47] fix(mcp): opt every listing out of catalog recording unless it is served The aggregate listing and _list_mcp_tools now default to record_listing=False, so a catalog fetched inside a tools/call no longer fills the caller's listed-tools slot. The /mcp/proxy meta-tools (call_tool, search_tools, get_tool_schema) and the tool-search virtual tool stop recording: /mcp/proxy serves only the meta-tools and the search serves only its hits, so a later pre_mcp_call hook was reading a description the caller never listed. The tools/list handler, the Responses MCP handler and the /v1/mcp/tools management listing opt in with record_listing=True, since each serves the catalog to the caller. --- .../_experimental/mcp_server/operations.py | 8 ++- .../mcp_management_endpoints.py | 1 + .../mcp/litellm_proxy_mcp_handler.py | 1 + .../mcp_server/test_mcp_proxy_mode.py | 55 +++++++++++++++++++ .../mcp_server/test_mcp_tool_search.py | 35 +++++++++++- .../mcp_server/test_operations.py | 30 ++++++++++ .../mcp/test_litellm_proxy_mcp_handler.py | 44 ++++++++++++++- 7 files changed, 171 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index d8886e26da7..3eb6265de15 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -958,7 +958,7 @@ async def _get_tools_from_mcp_servers( client_ip: str | None = None, mcp_proxy_mode: bool = False, *, - record_listing: bool = True, + record_listing: bool = False, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -1470,6 +1470,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. @@ -1480,6 +1482,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 @@ -1498,6 +1502,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 @@ -2757,6 +2762,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/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index e879b6daadd..02250c5f226 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1010,6 +1010,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/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 71f61079154..f1b9152ab3c 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -328,6 +328,7 @@ class LiteLLM_Proxy_MCP_Handler: litellm_trace_id=litellm_trace_id, request_tags=request_tags, raw_headers=raw_headers, + record_listing=True, ) tools: Final = listing.tools 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_tool_search.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 4575741aa8b..95903b13ba0 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 @@ -22,6 +22,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, @@ -31,12 +32,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, ...]: @@ -1268,3 +1271,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_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index bb900de4f98..47c146ae1bd 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -3,7 +3,10 @@ from unittest.mock import AsyncMock, patch import pytest from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult +from mcp.types import Tool as MCPTool +from litellm.proxy._experimental.mcp_server import operations +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._types import UserAPIKeyAuth from litellm.types.mcp import MCPAuth, MCPTransport @@ -665,3 +668,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/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 643944673a2..49c4a4a55b2 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -4,20 +4,25 @@ import sys import textwrap import types from typing import Any, Final, cast -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException from mcp.types import CallToolResult, TextContent, Tool as MCPTool from openai.types.responses.tool_param import Mcp +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 +from litellm.proxy._types import UserAPIKeyAuth 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.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 @@ -1311,3 +1316,40 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py logged: Final = setup.call_args.kwargs["metadata"]["headers"] assert logged == {"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a"} assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel" + + +@pytest.mark.asyncio +async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch: pytest.MonkeyPatch) -> None: + """The Responses bridge serves the listing to the model and its own tools/call reads the slot, so + this listing records the caller's catalog.""" + manager: Final = mcp_operations.global_mcp_server_manager + server: Final = MCPServer(server_id="responses-slot", name="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"})] + 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: + tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + user_api_key_auth=user, + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/responses-slot"}], + ) + listed: Final = 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 tools] == ["responses-slot-echo"] + assert listed is not None and listed.description == "Echo text back" From ddb672155f13a1f3be67538f34cb6d591004c91d Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 16:16:56 -0700 Subject: [PATCH 40/47] fix(mcp): key the listed-tool slot by the caller's admission identity and forwarded bearer The slot a tools/list records for a later tools/call was keyed by (user_id, api_key) only, so every team-only JWT caller shared one slot and one JWT user acting in two teams shared a slot; a tools/call then handed pre_mcp_call hooks a description another caller was served. The slot is now keyed by the hashed key, user, team and organization, plus the admission credential of a caller admitted with neither a key nor a user. The caller bearer split the slot only on client-forwarded-token and token-exchange servers; a legacy delegated oauth2 server (delegate_auth_to_upstream without client credentials) also forwards it upstream and served a different catalog per bearer into one slot. The bearer now splits the slot on every server whose egress forwards it (_consumes_caller_authorization) or exchanges it. --- .../mcp_server/mcp_server_manager.py | 47 ++++--- .../mcp_server/test_mcp_server_manager.py | 117 ++++++++++++++++++ 2 files changed, 149 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2d13dc0dcf3..49b2c2691c5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1221,6 +1221,21 @@ def listed_tools_caller_for( ) +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. @@ -4692,11 +4707,14 @@ class MCPServerManager: 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 own key (default-on guardrails, key or team - selections and opt-outs), so every keyed caller gets its own slot, on OpenAPI servers too. - Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or exchanged 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. + 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 @@ -4706,19 +4724,18 @@ class MCPServerManager: 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 server.is_client_forwarded_token or server.auth_type == MCPAuth.oauth2_token_exchange + if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange else None ) - _, digest = self._discovery_key( - server, - auth, - caller.mcp_auth_header, - forwarded, - stdio_env, - caller_bearer, - per_caller=auth is not 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 digest + return hashlib.sha256(material.encode()).hexdigest() @staticmethod def _forwarded_header_values( 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 b0b40982ade..d73028a6751 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 @@ -7722,6 +7722,123 @@ class TestMCPServerManager: 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() From b0e6f66190803425c634154c00aeed36883847c2 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 08:58:48 +0000 Subject: [PATCH 41/47] fix(mcp): record only tools served by the bridge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_experimental/mcp_server/operations.py | 62 ++++++++++++++++--- .../mcp/litellm_proxy_mcp_handler.py | 8 +++ .../mcp/test_litellm_proxy_mcp_handler.py | 52 +++++++++++++--- 3 files changed, 104 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index b60a6c99fd0..417d5280ef4 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -4,7 +4,7 @@ import asyncio import traceback import types import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from datetime import datetime from typing import Any, Final, NoReturn, TypeAlias, overload @@ -81,6 +81,7 @@ 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, @@ -957,6 +958,7 @@ async def _get_tools_from_mcp_servers( mcp_proxy_mode: bool = False, *, record_listing: bool = False, + served_tool_selector: Callable[[list[MCPTool]], list[MCPTool]] | None = None, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -969,6 +971,7 @@ async def _get_tools_from_mcp_servers( 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 + served_tool_selector: Select existing tool objects from the aggregate for deferred recording Returns: AggregateToolListing: Combined tools from filtered servers plus each server's @@ -1066,12 +1069,12 @@ async def _get_tools_from_mcp_servers( async def _fetch_and_filter_server_tools( server: MCPServer, - ) -> "tuple[list[MCPTool], ServerOutcome]": + ) -> tuple[list[MCPTool], ServerOutcome, Callable[[frozenset[int]], None] | None]: """Fetch and filter tools from a single server, classifying any failure into that server's outcome so the aggregate can keep serving the healthy subset without a broken server masquerading as an empty one.""" if server is None: - return [], ServerListOk(tool_count=0) + return [], ServerListOk(tool_count=0), None server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, @@ -1123,6 +1126,18 @@ async def _get_tools_from_mcp_servers( try: from litellm.proxy.proxy_server import proxy_logging_obj + defer_recording: Final = record_listing and served_tool_selector is not None + listed_generation: Final = ( + global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0) + if defer_recording + else None + ) + listed_caller: Final = ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=catalog_auth_header, + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ) tools: Final = await global_mcp_server_manager._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, @@ -1134,7 +1149,7 @@ async def _get_tools_from_mcp_servers( oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, catalog_auth_header=catalog_auth_header, - record_listing=record_listing, + record_listing=record_listing and served_tool_selector is None, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1144,6 +1159,14 @@ async def _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, ) + unprefixed_tools: Final = ( + tuple( + tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server) or tool.name}) + for tool in filtered_tools + ) + if defer_recording + else () + ) if mcp_proxy_mode: from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity @@ -1151,13 +1174,27 @@ async def _get_tools_from_mcp_servers( else: filtered_tools = apply_display_name_overrides(filtered_tools, server) + catalog_entries: Final = tuple(zip(filtered_tools, unprefixed_tools)) + + def record_served_tools(served_ids: frozenset[int]) -> None: + global_mcp_server_manager._record_listed_tools( + server, + [original for exposed, original in catalog_entries if id(exposed) in served_ids], + listed_caller, + listed_generation, + ) + verbose_logger.debug( "Successfully fetched %s tools from server %s, %s after filtering", len(tools), server.name, len(filtered_tools), ) - return filtered_tools, ServerListOk(tool_count=len(filtered_tools)) + return ( + filtered_tools, + ServerListOk(tool_count=len(filtered_tools)), + record_served_tools if defer_recording else None, + ) except MCPUpstreamAuthError as e: # Absorb so one unauthenticated server does not empty every other server's # tools. Surfacing the upstream 401 to the client as a re-auth challenge is @@ -1166,20 +1203,25 @@ async def _get_tools_from_mcp_servers( # error). Single-server routes surface it via the request-scope preemptive # check in _raise_preemptive_401_for_unauthenticated_servers instead. verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name) - return [], classify_list_exception(e) + return [], classify_list_exception(e), None except Exception as e: verbose_logger.exception("Error getting tools from server %s: %s", server.name, e) - return [], classify_list_exception(e) + return [], classify_list_exception(e), None # Fetch tools from all servers in parallel tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers] results: Final = await asyncio.gather(*tasks) # Flatten results into single list - all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools] + all_tools: Final[list[MCPTool]] = [tool for tools, _, _record in results for tool in tools] + if record_listing and served_tool_selector is not None: + served_ids: Final = frozenset(id(tool) for tool in served_tool_selector(all_tools)) + for _tools, _outcome, record in results: + if record is not None: + record(served_ids) server_outcomes: Final[dict[str, ServerOutcome]] = { _aggregate_server_key(server): outcome - for server, (_, outcome) in zip(allowed_mcp_servers, results) + for server, (_, outcome, _record) in zip(allowed_mcp_servers, results) if server is not None } @@ -1187,7 +1229,7 @@ async def _get_tools_from_mcp_servers( if litellm_logging_obj: per_server_tool_counts: Final[dict[str, int]] = { _aggregate_server_key(server): len(server_tools) - for server, (server_tools, _) in zip(allowed_mcp_servers, results) + for server, (server_tools, _, _record) in zip(allowed_mcp_servers, results) if server is not None } diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 2ca1c32f6dc..98de4ecd403 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -318,6 +318,13 @@ class LiteLLM_Proxy_MCP_Handler: # names), so use None and let the auth object's mcp_servers do the filtering. effective_server_filter: Final = None if resolved_toolset_ids else (resolved_mcp_servers or None) + def served_tools(tools: list[MCPTool]) -> list[MCPTool]: + filtered: Final = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools( + tools, mcp_tools_with_litellm_proxy + ) + deduplicated, _server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(filtered, []) + return deduplicated + listing: Final = await _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -329,6 +336,7 @@ class LiteLLM_Proxy_MCP_Handler: request_tags=request_tags, raw_headers=raw_headers, record_listing=True, + served_tool_selector=served_tools, ) tools: Final = listing.tools 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 44b7034ca34..ea540088b27 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1316,13 +1316,30 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py @pytest.mark.asyncio -async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch: pytest.MonkeyPatch) -> None: +@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_get_mcp_tools_from_manager_records_the_served_catalog( + monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str] +) -> None: """The Responses bridge serves the listing to the model and its own tools/call reads the slot, so this listing records the caller's catalog.""" manager: Final = mcp_operations.global_mcp_server_manager - server: Final = MCPServer(server_id="responses-slot", name="responses-slot", transport=MCPTransport.http) + 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"})] + 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=[]), @@ -1340,16 +1357,35 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), ): try: - tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + 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"}], + mcp_tools_with_litellm_proxy=[ + { + "type": "mcp", + "server_url": "litellm_proxy/mcp/responses-slot", + "allowed_tools": allowed_tools, + } + ], + ) + caller: Final = ListedToolsCaller(user_api_key_auth=user) + 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 == { + tool.name.removeprefix("responses_slot-"): (tool.description, tool.input_schema) for tool in tools + } + assert ( + manager.get_listed_tool( + server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller")) + ) + is None ) - listed: Final = 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 tools] == ["responses-slot-echo"] - assert listed is not None and listed.description == "Echo text back" + assert [tool.name for tool in tools] == expected_names def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace: From 68eec0b3c5dde7a29d13c1b99e222a45837b0c73 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 18:22:00 +0000 Subject: [PATCH 42/47] fix(mcp): keep bridge tool metadata request-local --- .../mcp_server/mcp_server_manager.py | 3 +- .../_experimental/mcp_server/operations.py | 62 ++-------- litellm/responses/main.py | 3 + .../responses/mcp/chat_completions_handler.py | 2 + .../mcp/litellm_proxy_mcp_handler.py | 17 ++- .../responses/mcp/mcp_streaming_iterator.py | 3 + .../mcp/test_litellm_proxy_mcp_handler.py | 117 ++++++++++++++++-- 7 files changed, 134 insertions(+), 73 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1c6ac52f422..985e1a635e7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -6664,6 +6664,7 @@ class MCPServerManager: 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 @@ -6717,7 +6718,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), + 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 417d5280ef4..b60a6c99fd0 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -4,7 +4,7 @@ import asyncio import traceback import types import uuid -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Mapping, Sequence from datetime import datetime from typing import Any, Final, NoReturn, TypeAlias, overload @@ -81,7 +81,6 @@ 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, @@ -958,7 +957,6 @@ async def _get_tools_from_mcp_servers( mcp_proxy_mode: bool = False, *, record_listing: bool = False, - served_tool_selector: Callable[[list[MCPTool]], list[MCPTool]] | None = None, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -971,7 +969,6 @@ async def _get_tools_from_mcp_servers( 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 - served_tool_selector: Select existing tool objects from the aggregate for deferred recording Returns: AggregateToolListing: Combined tools from filtered servers plus each server's @@ -1069,12 +1066,12 @@ async def _get_tools_from_mcp_servers( async def _fetch_and_filter_server_tools( server: MCPServer, - ) -> tuple[list[MCPTool], ServerOutcome, Callable[[frozenset[int]], None] | None]: + ) -> "tuple[list[MCPTool], ServerOutcome]": """Fetch and filter tools from a single server, classifying any failure into that server's outcome so the aggregate can keep serving the healthy subset without a broken server masquerading as an empty one.""" if server is None: - return [], ServerListOk(tool_count=0), None + return [], ServerListOk(tool_count=0) server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, @@ -1126,18 +1123,6 @@ async def _get_tools_from_mcp_servers( try: from litellm.proxy.proxy_server import proxy_logging_obj - defer_recording: Final = record_listing and served_tool_selector is not None - listed_generation: Final = ( - global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0) - if defer_recording - else None - ) - listed_caller: Final = ListedToolsCaller( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=catalog_auth_header, - raw_headers=raw_headers, - oauth2_headers=oauth2_headers, - ) tools: Final = await global_mcp_server_manager._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, @@ -1149,7 +1134,7 @@ async def _get_tools_from_mcp_servers( oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, catalog_auth_header=catalog_auth_header, - record_listing=record_listing and served_tool_selector is None, + record_listing=record_listing, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1159,14 +1144,6 @@ async def _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, ) - unprefixed_tools: Final = ( - tuple( - tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server) or tool.name}) - for tool in filtered_tools - ) - if defer_recording - else () - ) if mcp_proxy_mode: from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity @@ -1174,27 +1151,13 @@ async def _get_tools_from_mcp_servers( else: filtered_tools = apply_display_name_overrides(filtered_tools, server) - catalog_entries: Final = tuple(zip(filtered_tools, unprefixed_tools)) - - def record_served_tools(served_ids: frozenset[int]) -> None: - global_mcp_server_manager._record_listed_tools( - server, - [original for exposed, original in catalog_entries if id(exposed) in served_ids], - listed_caller, - listed_generation, - ) - verbose_logger.debug( "Successfully fetched %s tools from server %s, %s after filtering", len(tools), server.name, len(filtered_tools), ) - return ( - filtered_tools, - ServerListOk(tool_count=len(filtered_tools)), - record_served_tools if defer_recording else None, - ) + return filtered_tools, ServerListOk(tool_count=len(filtered_tools)) except MCPUpstreamAuthError as e: # Absorb so one unauthenticated server does not empty every other server's # tools. Surfacing the upstream 401 to the client as a re-auth challenge is @@ -1203,25 +1166,20 @@ async def _get_tools_from_mcp_servers( # error). Single-server routes surface it via the request-scope preemptive # check in _raise_preemptive_401_for_unauthenticated_servers instead. verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name) - return [], classify_list_exception(e), None + return [], classify_list_exception(e) except Exception as e: verbose_logger.exception("Error getting tools from server %s: %s", server.name, e) - return [], classify_list_exception(e), None + return [], classify_list_exception(e) # Fetch tools from all servers in parallel tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers] results: Final = await asyncio.gather(*tasks) # Flatten results into single list - all_tools: Final[list[MCPTool]] = [tool for tools, _, _record in results for tool in tools] - if record_listing and served_tool_selector is not None: - served_ids: Final = frozenset(id(tool) for tool in served_tool_selector(all_tools)) - for _tools, _outcome, record in results: - if record is not None: - record(served_ids) + all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools] server_outcomes: Final[dict[str, ServerOutcome]] = { _aggregate_server_key(server): outcome - for server, (_, outcome, _record) in zip(allowed_mcp_servers, results) + for server, (_, outcome) in zip(allowed_mcp_servers, results) if server is not None } @@ -1229,7 +1187,7 @@ async def _get_tools_from_mcp_servers( if litellm_logging_obj: per_server_tool_counts: Final[dict[str, int]] = { _aggregate_server_key(server): len(server_tools) - for server, (server_tools, _, _record) in zip(allowed_mcp_servers, results) + for server, (server_tools, _) in zip(allowed_mcp_servers, results) if server is not None } diff --git a/litellm/responses/main.py b/litellm/responses/main.py index d145f8cc6b8..68952d2e525 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 98de4ecd403..8e37466c6c9 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -318,13 +318,6 @@ class LiteLLM_Proxy_MCP_Handler: # names), so use None and let the auth object's mcp_servers do the filtering. effective_server_filter: Final = None if resolved_toolset_ids else (resolved_mcp_servers or None) - def served_tools(tools: list[MCPTool]) -> list[MCPTool]: - filtered: Final = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools( - tools, mcp_tools_with_litellm_proxy - ) - deduplicated, _server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(filtered, []) - return deduplicated - listing: Final = await _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -335,8 +328,6 @@ class LiteLLM_Proxy_MCP_Handler: litellm_trace_id=litellm_trace_id, request_tags=request_tags, raw_headers=raw_headers, - record_listing=True, - served_tool_selector=served_tools, ) tools: Final = listing.tools @@ -705,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 @@ -869,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: @@ -1161,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: """ @@ -1190,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/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index ea540088b27..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,9 +1,10 @@ +import asyncio import importlib import subprocess import sys import textwrap import types -from typing import Any, Final, cast +from typing import Any, Final, Literal, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -12,20 +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 +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: @@ -1316,6 +1323,7 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py @pytest.mark.asyncio +@pytest.mark.parametrize("real_listing", [False, True]) @pytest.mark.parametrize( ("allowed_tools", "expected_names"), [ @@ -1325,11 +1333,9 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py (["absent"], []), ], ) -async def test_get_mcp_tools_from_manager_records_the_served_catalog( - monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str] +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: - """The Responses bridge serves the listing to the model and its own tools/call reads the slot, so - this listing records the caller's catalog.""" manager: Final = mcp_operations.global_mcp_server_manager server: Final = MCPServer( server_id="responses-slot", name="responses_slot", alias="responses_slot", transport=MCPTransport.http @@ -1357,6 +1363,15 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog( 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=[ @@ -1367,15 +1382,12 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog( } ], ) - caller: Final = ListedToolsCaller(user_api_key_auth=user) 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 == { - tool.name.removeprefix("responses_slot-"): (tool.description, tool.input_schema) for tool in tools - } + assert recorded == before assert ( manager.get_listed_tool( server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller")) @@ -1388,6 +1400,89 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog( 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={}), From 7f7df081b9e0f0b196142a66a86d8b641ea1b6f4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 18:40:00 +0000 Subject: [PATCH 43/47] refactor(mcp): centralize listed catalog recording guard --- .../_experimental/mcp_server/mcp_server_manager.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 332800556ef..304f556721b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4638,8 +4638,9 @@ class MCPServerManager: # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". unprefixed_tools: Final = guarded_openapi - if record_listing: - self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation) + self._record_listed_tools( + server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing + ) if not add_prefix: return unprefixed_tools return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] @@ -4657,8 +4658,9 @@ class MCPServerManager: prefixed_or_original_tools: Final = self._create_prefixed_tools( guarded_tools, server, add_prefix=add_prefix ) - if record_listing: - self._record_listed_tools(server, guarded_tools, listed_caller, listed_generation) + self._record_listed_tools( + server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing + ) return prefixed_or_original_tools @@ -4765,9 +4767,13 @@ class MCPServerManager: 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) From 29baaf3f5e83561d682f23afeea4da77cc10113b Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 21:43:56 +0000 Subject: [PATCH 44/47] fix(mcp): preserve base TPM reservations for listed tool calls --- .../hooks/parallel_request_limiter_v3.py | 17 ++++- .../hooks/test_parallel_request_limiter_v3.py | 76 +++++++++++++++++++ 2 files changed, 91 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b33cea5742d..d979d56afa7 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1028,7 +1028,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _estimate_tokens_for_request( self, - data: dict, + data: dict[str, object], model: str | None = None, min_configured_tpm_limit: int | None = None, call_type: str | None = None, @@ -1054,8 +1054,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): floor entirely, so the reservation reflects what this tenant's model actually emits rather than one constant shared by every tenant. """ + reservation_data: Final = ( + { + **data, + "messages": [ + { + "role": "user", + "content": f"Tool: {data.get('mcp_tool_name')}\nArguments: {data.get('mcp_arguments')}", + } + ], + } + if call_type == CallTypes.call_mcp_tool.value and "mcp_tool_name" in data and "mcp_arguments" in data + else data + ) estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens( - data=data, + data=reservation_data, min_configured_tpm_limit=min_configured_tpm_limit, call_type=call_type, configured_output_tokens=configured_output_tokens, 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..1d8731945b8 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,79 @@ 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"] +) +async def test_mcp_description_does_not_change_admission_or_reserved_tokens(description: str | None) -> 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) + + 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 == 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": ( + f"Tool: echo\nDescription: {description}\nArguments: {{'q': 'hello'}}" + if description + else "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): From acc02ce8d4bca80dd7130a095f8cee841a2314e1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 00:00:46 +0000 Subject: [PATCH 45/47] fix(mcp): preserve project token reservations for listed calls --- .../hooks/parallel_request_limiter_v3.py | 39 ++++++---- .../hooks/test_parallel_request_limiter_v3.py | 77 +++++++++++++++++++ 2 files changed, 101 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index d979d56afa7..999f9ea7331 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1026,6 +1026,26 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if existing_cap is None or effective_cap < existing_cap: data[cap_field] = effective_cap # rebind-ok: downstream routing requires the bounded output cap + @staticmethod + def _mcp_token_reservation_data(data: object, call_type: str | None) -> object: + if ( + call_type != CallTypes.call_mcp_tool.value + or not isinstance(data, dict) + or "mcp_tool_name" not in data + or "mcp_arguments" not in data + ): + return data + mcp_data: Final = TypeAdapter(dict[str, object]).validate_python(data) + return { + **mcp_data, + "messages": [ + { + "role": "user", + "content": f"Tool: {mcp_data['mcp_tool_name']}\nArguments: {mcp_data['mcp_arguments']}", + } + ], + } + def _estimate_tokens_for_request( self, data: dict[str, object], @@ -1054,19 +1074,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): floor entirely, so the reservation reflects what this tenant's model actually emits rather than one constant shared by every tenant. """ - reservation_data: Final = ( - { - **data, - "messages": [ - { - "role": "user", - "content": f"Tool: {data.get('mcp_tool_name')}\nArguments: {data.get('mcp_arguments')}", - } - ], - } - if call_type == CallTypes.call_mcp_tool.value and "mcp_tool_name" in data and "mcp_arguments" in data - else data - ) + reservation_data: Final = self._mcp_token_reservation_data(data, call_type) estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens( data=reservation_data, min_configured_tpm_limit=min_configured_tpm_limit, @@ -3787,13 +3795,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if v is not None ] min_configured_otpm_limit: Final = min(configured_otpm_limits) if configured_otpm_limits else None + reservation_data: Final = self._mcp_token_reservation_data(data, call_type) _, raw_estimated_output_tokens = self._estimate_input_and_output_tokens( - data=data, + data=reservation_data, min_configured_tpm_limit=min_configured_otpm_limit, call_type=call_type, ) raw_estimated_input_tokens: Final = await offload_token_count(self._estimate_precise_input_tokens)( - data=data, model=requested_model, call_type=call_type + data=reservation_data, model=requested_model, call_type=call_type ) estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1) estimated_output_tokens: Final = ( 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 1d8731945b8..5682c73503e 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -153,6 +153,83 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc ] +@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)]) +async def test_mcp_description_preserves_project_input_and_output_reservations( + description: str | None, itpm_limit: int, otpm_limit: int +) -> 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}, + }, + ) + + 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, 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": ( + f"Tool: echo\nDescription: {description}\nArguments: {{'q': 'hello'}}" + if description + else "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]] = { From c669328ba10b6b2dac7cd79e2400ffec2af52be9 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 03:08:40 +0000 Subject: [PATCH 46/47] fix(mcp): record served catalogs and preserve call message bytes --- .../_experimental/mcp_server/operations.py | 19 ++++- .../mcp_server/rest_endpoints.py | 45 +++++----- .../hooks/parallel_request_limiter_v3.py | 30 +------ litellm/proxy/utils.py | 3 +- .../mcp_server/test_operations.py | 83 +++++++++++++++++++ .../hooks/test_parallel_request_limiter_v3.py | 31 +++---- 6 files changed, 148 insertions(+), 63 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index b60a6c99fd0..b0410e87103 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -81,6 +81,7 @@ 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, @@ -1123,6 +1124,7 @@ async def _get_tools_from_mcp_servers( try: from litellm.proxy.proxy_server import proxy_logging_obj + listed_generation: Final = global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0) tools: Final = await global_mcp_server_manager._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, @@ -1134,7 +1136,7 @@ async def _get_tools_from_mcp_servers( oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, catalog_auth_header=catalog_auth_header, - record_listing=record_listing, + record_listing=False, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1143,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 diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 3f31ec1a1dd..c8e12ce370a 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 ( @@ -719,8 +720,8 @@ if MCP_AVAILABLE: ) 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, @@ -736,7 +737,8 @@ if MCP_AVAILABLE: """ from litellm.proxy.proxy_server import proxy_logging_obj - tools = await _list_server_tools( + listed_generation: Final = global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0) + tools: Final = await _list_server_tools( server, server_auth_header, raw_headers, @@ -744,30 +746,31 @@ if MCP_AVAILABLE: extra_headers, client_ip, proxy_logging_obj, - record_listing=True, + 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 + ) + 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 diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 999f9ea7331..b33cea5742d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1026,29 +1026,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if existing_cap is None or effective_cap < existing_cap: data[cap_field] = effective_cap # rebind-ok: downstream routing requires the bounded output cap - @staticmethod - def _mcp_token_reservation_data(data: object, call_type: str | None) -> object: - if ( - call_type != CallTypes.call_mcp_tool.value - or not isinstance(data, dict) - or "mcp_tool_name" not in data - or "mcp_arguments" not in data - ): - return data - mcp_data: Final = TypeAdapter(dict[str, object]).validate_python(data) - return { - **mcp_data, - "messages": [ - { - "role": "user", - "content": f"Tool: {mcp_data['mcp_tool_name']}\nArguments: {mcp_data['mcp_arguments']}", - } - ], - } - def _estimate_tokens_for_request( self, - data: dict[str, object], + data: dict, model: str | None = None, min_configured_tpm_limit: int | None = None, call_type: str | None = None, @@ -1074,9 +1054,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): floor entirely, so the reservation reflects what this tenant's model actually emits rather than one constant shared by every tenant. """ - reservation_data: Final = self._mcp_token_reservation_data(data, call_type) estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens( - data=reservation_data, + data=data, min_configured_tpm_limit=min_configured_tpm_limit, call_type=call_type, configured_output_tokens=configured_output_tokens, @@ -3795,14 +3774,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if v is not None ] min_configured_otpm_limit: Final = min(configured_otpm_limits) if configured_otpm_limits else None - reservation_data: Final = self._mcp_token_reservation_data(data, call_type) _, raw_estimated_output_tokens = self._estimate_input_and_output_tokens( - data=reservation_data, + data=data, min_configured_tpm_limit=min_configured_otpm_limit, call_type=call_type, ) raw_estimated_input_tokens: Final = await offload_token_count(self._estimate_precise_input_tokens)( - data=reservation_data, model=requested_model, call_type=call_type + data=data, model=requested_model, call_type=call_type ) estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1) estimated_output_tokens: Final = ( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index cf73fc0b99c..df7d859b15a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1479,7 +1479,8 @@ class ProxyLogging: if request_obj.tool_input_schema is not None else kwargs.get("mcp_input_schema") ) - description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else "" + 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}" ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index 47c146ae1bd..fc1e23229f9 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -1,18 +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 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 5682c73503e..d3a723fe1d3 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -115,7 +115,10 @@ def test_api_key_descriptor_applies_budget_throttle( @pytest.mark.parametrize( "description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"] ) -async def test_mcp_description_does_not_change_admission_or_reserved_tokens(description: str | None) -> None: +@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()) @@ -127,7 +130,10 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc messages: Final = data["messages"] caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-description-reservation"), tpm_limit=64) - await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool") + 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 @@ -144,11 +150,7 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc assert messages == [ { "role": "user", - "content": ( - f"Tool: echo\nDescription: {description}\nArguments: {{'q': 'hello'}}" - if description - else "Tool: echo\nArguments: {'q': 'hello'}" - ), + "content": "Tool: echo\nArguments: {'q': 'hello'}", } ] @@ -158,8 +160,10 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc "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 + 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)) @@ -188,7 +192,10 @@ async def test_mcp_description_preserves_project_input_and_output_reservations( }, ) - await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool") + 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 @@ -221,11 +228,7 @@ async def test_mcp_description_preserves_project_input_and_output_reservations( assert messages == [ { "role": "user", - "content": ( - f"Tool: echo\nDescription: {description}\nArguments: {{'q': 'hello'}}" - if description - else "Tool: echo\nArguments: {'q': 'hello'}" - ), + "content": "Tool: echo\nArguments: {'q': 'hello'}", } ] From cbc12d1e1cbed6736008d0e29793d43d8757c3d7 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 03:18:39 +0000 Subject: [PATCH 47/47] test(mcp): align listing expectation with deferred recording --- tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 441cc57c292..362ffcca1fa 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -2000,7 +2000,7 @@ async def test_get_tools_for_single_server(): client_ip=None, user_api_key_auth=None, proxy_logging_obj=ANY, - record_listing=True, + record_listing=False, ) # Verify the result