From 3c13e3b457ffa555d504056094452a60b17bcc33 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 15 Sep 2026 00:42:16 +0000 Subject: [PATCH 01/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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/23] 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(