diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c0792c32de2..02d7255700d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -257,6 +257,20 @@ _user_env_vars_cache: Final[dict[tuple[str, str], tuple[dict[str, str], float]]] _USER_ENV_VARS_CACHE_TTL: Final = 60 # seconds _USER_ENV_VARS_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth +_ListedToolsByCaller: TypeAlias = Mapping[str | None, Mapping[str, MCPTool]] +_LISTED_TOOLS_CALLERS_PER_SERVER: Final = 256 + + +@dataclass(frozen=True, slots=True) +class ListedToolsCaller: + """Request inputs that select which upstream catalog a caller was shown by tools/list.""" + + user_api_key_auth: UserAPIKeyAuth | None = None + mcp_auth_header: str | dict[str, str] | None = None + raw_headers: Mapping[str, str] | None = None + oauth2_headers: Mapping[str, str] | None = None + + # Auth types whose upstream OAuth endpoints (protected-resource + authorization-server metadata) the # gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes. # OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the @@ -1170,6 +1184,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. @@ -1295,6 +1328,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, @@ -1973,6 +2021,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 @@ -3777,19 +3826,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( @@ -3875,7 +3912,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.""" @@ -4406,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. @@ -4425,6 +4462,13 @@ class MCPServerManager: verbose_logger.info("_get_tools_from_server for %s...", server.name) client = None + resolved_mcp_auth_header: Final = await _byok_listing_auth_header(server, user_api_key_auth, mcp_auth_header) + listed_caller: Final = ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=resolved_mcp_auth_header, + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ) try: # Tool *listing* must not be blocked by missing per-user env vars — @@ -4468,7 +4512,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( @@ -4491,7 +4535,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, @@ -4522,8 +4566,10 @@ class MCPServerManager: # applied (e.g. "test_petstore-getinventory"). Do NOT pass them # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". + unprefixed_tools: Final = guarded_openapi + self._record_listed_tools(server, unprefixed_tools, listed_caller) if not add_prefix: - return list(guarded_openapi) + return unprefixed_tools return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] else: tools = await self._fetch_tools_with_timeout(client, server.name) @@ -4537,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 + guarded_tools, server, add_prefix=add_prefix, caller=listed_caller ) return prefixed_or_original_tools @@ -4588,8 +4634,77 @@ class MCPServerManager: ) self._invalidate_discovery_lists(server_id) + self._listed_tools_by_server_id.pop(server_id, None) invalidate_oauth_metadata_cache(server_id) + def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None: + """Key the listed-tool cache by every request input that can change the upstream catalog. + + Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or + exchanged as the OBO subject), the server-specific auth header, and the per-caller JWT + MCPJWTSigner mints for tools/list all reach upstream, so two callers differing in any of + them may be shown different tools. Shared servers with none of those stay on the shared + (``None``) slot. OpenAPI servers list from the process-wide registry. + """ + if server.spec_path or caller is None: + return None + auth: Final = caller.user_api_key_auth + forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None + header_env: Final = self._build_stdio_env(server, caller.raw_headers) + stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env + caller_bearer: Final = ( + self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) + if 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, + per_caller=self._signs_caller_identity_upstream(server), + ) + return digest + + @staticmethod + def _signs_caller_identity_upstream(server: MCPServer) -> bool: + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server + get_mcp_jwt_signer, + ) + + if get_mcp_jwt_signer() is None: + return False + return server.static_headers is None or not any(k.lower() == "authorization" for k in server.static_headers) + + @staticmethod + def _forwarded_header_values( + server: MCPServer, raw_headers: Mapping[str, str] | None + ) -> 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, MappingProxyType({})) + shared: Final = existing.get(None) + callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity)) + evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0) + entries: Final = ( + *(() if shared is None else ((None, shared),)), + *callers[evicted:], + (identity, listing), + ) + self._listed_tools_by_server_id[server.server_id] = MappingProxyType(dict(entries)) + def _discovery_key( self, server: MCPServer, @@ -4599,9 +4714,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) @@ -5490,7 +5607,13 @@ class MCPServerManager: {seen: kept for seen, kept in self._catalog_alert_signatures.items() if seen != key} ) - def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]: + def _create_prefixed_tools( + self, + tools: Sequence[MCPTool], + server: MCPServer, + add_prefix: bool = True, + caller: ListedToolsCaller | None = None, + ) -> list[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -5515,9 +5638,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, MappingProxyType({})).get(identity) + if not listed: + return None + tool: Final = listed.get(name) or listed.get(strip_known_server_prefix(name, server)) + if tool is None: + return None + description: Final = server.tool_name_to_description.get(tool.name) if server.tool_name_to_description else None + return tool if description is None else tool.model_copy(update={"description": description}) + def _create_prefixed_prompts( self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True ) -> list[Prompt]: @@ -5751,6 +5886,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. @@ -5764,6 +5900,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 @@ -5827,6 +5966,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 @@ -5882,6 +6023,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. @@ -5896,6 +6038,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(), ) @@ -6042,21 +6186,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 @@ -6507,6 +6637,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 @@ -6523,6 +6659,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"] @@ -6539,6 +6676,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) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 8bb2772c760..cccdf4c59c0 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 @@ -1606,6 +1607,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], @@ -2075,6 +2082,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. @@ -2143,7 +2151,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 @@ -2185,6 +2194,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 52f3eeb4ce7..94de0b2df7b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn import httpx from fastapi import HTTPException -from pydantic import TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger @@ -61,6 +61,7 @@ _GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset( _INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027" _AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...]) _MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool") +_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object]) _OBO_CACHE_MAX_ENTRIES: Final = 1000 _DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0 _TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0 @@ -82,6 +83,13 @@ def _parse_aadsts_codes(raw: object) -> tuple[int, ...]: return () +def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None: + try: + return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw) + except ValidationError: + return None + + def entra_assertion(value: object) -> str | None: """``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion. A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``.""" @@ -100,6 +108,14 @@ class _EvaluateResponse(TypedDict, total=False): correlationId: ReadOnly[str] +class _ToolReference(BaseModel): + model_config = ConfigDict(frozen=True) + + name: str + description: str | None = None + input_schema: Mapping[str, object] | None = Field(default=None, serialization_alias="inputSchema") + + class _UnavailableDetail(TypedDict): error: ReadOnly[str] message: ReadOnly[str] @@ -392,8 +408,14 @@ class Agent365Guardrail(CustomGuardrail): arguments: Final = data.get("mcp_arguments") server_name: Final = str(data.get("mcp_server_name") or "litellm") agent_id: Final = user_api_key_dict.key_alias + description: Final = data.get("mcp_tool_description") + tool_reference: Final = _ToolReference( + name=tool_name, + description=description if isinstance(description, str) and description else None, + input_schema=_parse_tool_input_schema(data.get("mcp_input_schema")), + ) payload: Final[dict[str, object]] = { # mutable-ok: JSON body with optional fields added below - "tool": {"name": tool_name}, + "tool": tool_reference.model_dump(by_alias=True, exclude_none=True), "serverName": server_name, "conversationId": self._resolve_conversation_id(data), } diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ea294b76e92..f980d1e2428 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1476,8 +1476,12 @@ class ProxyLogging: TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({})) ) - mcp_tool_description: Final = kwargs.get("mcp_tool_description") - mcp_input_schema: Final = kwargs.get("mcp_input_schema") + mcp_tool_description: Final = request_obj.tool_description or kwargs.get("mcp_tool_description") + mcp_input_schema: Final = ( + request_obj.tool_input_schema + if request_obj.tool_input_schema is not None + else kwargs.get("mcp_input_schema") + ) description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else "" tool_call_content: Final = ( f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}" @@ -1736,6 +1740,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/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" 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 e1e4cd3d161..8e87837611a 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 a7e56f3f84a..62a7830a85d 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 @@ -6877,7 +6878,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 @@ -7989,9 +7989,12 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): fake_server.server_name = "openapi-petstore" fake_server.alias = None fake_server.short_prefix = None + fake_server.tool_name_to_description = None fake_tool = MagicMock() fake_tool.name = "list_pets" + fake_tool.description = "test tool" + fake_tool.input_schema = {"type": "object"} start_time = datetime.now(timezone.utc) litellm_logging_obj, _ = function_setup( @@ -8042,6 +8045,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 a1dc0e779da..4cfc9d36714 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 @@ -42,6 +42,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, @@ -115,7 +116,6 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte assert sampling.await_args.kwargs["raw_headers"] == {"x-test-caller": "sampling-caller"} - @pytest.mark.asyncio async def test_sampling_callback_keeps_creation_context_after_caller_switch(): from mcp.server.auth.middleware.auth_context import auth_context_var @@ -223,8 +223,6 @@ def _reload_mcp_manager_module(): return reloaded - - @pytest.fixture(autouse=True) def enable_eager_mcp_oauth_discovery(monkeypatch): monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1") @@ -1398,7 +1396,9 @@ class TestMCPServerManager: assert not any("oauth2_id_jag" in message for message in caplog.messages) @pytest.mark.asyncio - async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, config_only_mcp_manager_factory, monkeypatch, caplog): + async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso( + self, config_only_mcp_manager_factory, monkeypatch, caplog + ): self._clear_sso_env(monkeypatch) monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid") manager = config_only_mcp_manager_factory() @@ -4684,7 +4684,9 @@ class TestMCPServerManager: @pytest.mark.parametrize("auth_type", [MCPAuth.none, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.oauth2]) @pytest.mark.parametrize("is_byok", [False, True]) @pytest.mark.parametrize("scheme", ["http", "https"]) - async def test_openapi_health_loads_spec_without_mcp_handshake(self, respx_mock, monkeypatch, auth_type, is_byok, scheme): + async def test_openapi_health_loads_spec_without_mcp_handshake( + self, respx_mock, monkeypatch, auth_type, is_byok, scheme + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4734,14 +4736,28 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("failure", "expected_status", "expected_error"), [ - (httpx.Response(401, text="secret response content"), "unhealthy", "OpenAPI specification request failed (HTTP 401)"), + ( + httpx.Response(401, text="secret response content"), + "unhealthy", + "OpenAPI specification request failed (HTTP 401)", + ), (httpx.Response(404), "unhealthy", "OpenAPI specification request failed (HTTP 404)"), (httpx.Response(500), "unhealthy", "OpenAPI specification request failed (HTTP 500)"), - (httpx.ConnectError("secret network details"), "unhealthy", "OpenAPI specification could not be loaded (ConnectError)"), - (httpx.Response(200, text="secret invalid JSON body"), "unhealthy", "OpenAPI specification could not be loaded (JSONDecodeError)"), + ( + httpx.ConnectError("secret network details"), + "unhealthy", + "OpenAPI specification could not be loaded (ConnectError)", + ), + ( + httpx.Response(200, text="secret invalid JSON body"), + "unhealthy", + "OpenAPI specification could not be loaded (JSONDecodeError)", + ), ], ) - async def test_openapi_health_reports_safe_failures(self, respx_mock, monkeypatch, failure, expected_status, expected_error): + async def test_openapi_health_reports_safe_failures( + self, respx_mock, monkeypatch, failure, expected_status, expected_error + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4899,7 +4915,10 @@ class TestMCPServerManager: @pytest.mark.asyncio @pytest.mark.parametrize("oauth2_flow", [None, "authorization_code", "client_credentials"]) async def test_health_check_server_oauth2_reports_reachability( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, oauth2_flow: Literal["authorization_code", "client_credentials"] | None + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + oauth2_flow: Literal["authorization_code", "client_credentials"] | None, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() @@ -4928,14 +4947,28 @@ class TestMCPServerManager: assert not {"authorization", "x-api-key", "cookie"}.intersection(route.calls[0].request.headers) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type", [ - MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token, - MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, MCPAuth.true_passthrough, MCPAuth.oauth_delegate, - ]) + @pytest.mark.parametrize( + "auth_type", + [ + MCPAuth.bearer_token, + MCPAuth.api_key, + MCPAuth.basic, + MCPAuth.authorization, + MCPAuth.token, + MCPAuth.oauth2_token_exchange, + MCPAuth.oauth2_id_jag, + MCPAuth.true_passthrough, + MCPAuth.oauth_delegate, + ], + ) @pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse]) @pytest.mark.parametrize("response_code", [200, 204, 302, 401, 403, 405, 503]) async def test_health_check_without_credentials_accepts_any_http_response( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, auth_type: MCPAuthType, transport: Literal[MCPTransport.http, MCPTransport.sse], + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + auth_type: MCPAuthType, + transport: Literal[MCPTransport.http, MCPTransport.sse], response_code: int, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") @@ -4966,6 +4999,7 @@ class TestMCPServerManager: self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, response_code: int ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + class UnreadBody(httpx.AsyncByteStream): def __init__(self) -> None: self.read = False @@ -4980,17 +5014,28 @@ class TestMCPServerManager: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="streaming-health", name="streaming-health", transport=MCPTransport.sse, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/events", + server_id="streaming-health", + name="streaming-health", + transport=MCPTransport.sse, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/events", ) manager.registry[server.server_id] = server bodies: Final = (UnreadBody(), UnreadBody()) - route: Final = respx_mock.get(server.url).mock(side_effect=[ - httpx.Response(response_code, stream=body, headers={ - "Content-Type": "text/event-stream", "Set-Cookie": "health=secret; Path=/", - "Location": "http://127.0.0.1/private", - }) for body in bodies - ]) + route: Final = respx_mock.get(server.url).mock( + side_effect=[ + httpx.Response( + response_code, + stream=body, + headers={ + "Content-Type": "text/event-stream", + "Set-Cookie": "health=secret; Path=/", + "Location": "http://127.0.0.1/private", + }, + ) + for body in bodies + ] + ) first: Final = await manager.health_check_server(server.server_id) second: Final = await manager.health_check_server(server.server_id) @@ -5001,19 +5046,28 @@ class TestMCPServerManager: assert all("cookie" not in call.request.headers for call in route.calls) @pytest.mark.asyncio - @pytest.mark.parametrize(("transport", "url"), [ - (MCPTransport.stdio, "https://mcp.example.test"), - (MCPTransport.http, None), (MCPTransport.http, ""), (MCPTransport.http, "not-a-url"), - (MCPTransport.http, "ftp://mcp.example.test"), - (MCPTransport.http, "https://user:secret@mcp.example.test"), - (MCPTransport.http, "https://mcp.example.test:bad/mcp"), - ]) + @pytest.mark.parametrize( + ("transport", "url"), + [ + (MCPTransport.stdio, "https://mcp.example.test"), + (MCPTransport.http, None), + (MCPTransport.http, ""), + (MCPTransport.http, "not-a-url"), + (MCPTransport.http, "ftp://mcp.example.test"), + (MCPTransport.http, "https://user:secret@mcp.example.test"), + (MCPTransport.http, "https://mcp.example.test:bad/mcp"), + ], + ) async def test_health_reachability_rejects_unprobeable_urls_without_requests( self, respx_mock: MockRouter, transport: Literal[MCPTransport.http, MCPTransport.stdio], url: str | None ) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="unprobeable", name="unprobeable", transport=transport, auth_type=MCPAuth.oauth2, url=url, + server_id="unprobeable", + name="unprobeable", + transport=transport, + auth_type=MCPAuth.oauth2, + url=url, ) manager.registry[server.server_id] = server @@ -5024,19 +5078,26 @@ class TestMCPServerManager: assert not respx_mock.calls @pytest.mark.asyncio - @pytest.mark.parametrize("failure", [ - httpx.ConnectError("TLS/connection failure with secret details"), - httpx.ReadTimeout("secret timeout details"), - httpx.RemoteProtocolError("secret malformed response"), - ]) + @pytest.mark.parametrize( + "failure", + [ + httpx.ConnectError("TLS/connection failure with secret details"), + httpx.ReadTimeout("secret timeout details"), + httpx.RemoteProtocolError("secret malformed response"), + ], + ) async def test_health_reachability_reports_no_response_without_secret_details( self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, failure: httpx.RequestError ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="failed-health", name="failed-health", transport=MCPTransport.http, - auth_type=MCPAuth.bearer_token, is_byok=True, url="https://mcp.example.test/secret?token=secret", + server_id="failed-health", + name="failed-health", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + is_byok=True, + url="https://mcp.example.test/secret?token=secret", ) manager.registry[server.server_id] = server route: Final = respx_mock.get(server.url).mock(side_effect=failure) @@ -5052,8 +5113,11 @@ class TestMCPServerManager: monkeypatch.setenv("SSL_SECURITY_LEVEL", "invalid-secret-cipher") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="bad-tls", name="bad-tls", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test", + server_id="bad-tls", + name="bad-tls", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test", ) manager.registry[server.server_id] = server @@ -5071,8 +5135,11 @@ class TestMCPServerManager: monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_HEALTH_CHECK_TIMEOUT", 0.1) manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="slow-health", name="slow-health", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/slow", + server_id="slow-health", + name="slow-health", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/slow", ) manager.registry[server.server_id] = server started: Final = asyncio.Event() @@ -5125,8 +5192,11 @@ class TestMCPServerManager: server_ids: Final = [f"health-{index}" for index in range(server_count)] manager.registry = { server_id: MCPServer( - server_id=server_id, name=server_id, transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url=f"https://health.example.test/{server_id}", + server_id=server_id, + name=server_id, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url=f"https://health.example.test/{server_id}", ) for server_id in server_ids } @@ -5453,8 +5523,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers captured["server_label"] = server_label @@ -5539,8 +5616,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers @@ -6952,6 +7036,551 @@ 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_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): + 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", 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( + [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", 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") + 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( + {"auth_type": MCPAuth.oauth2_token_exchange}, + ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", token="hashed-shared"), + raw_headers={"x-litellm-api-key": "sk-shared", "authorization": "Bearer entra-alice"}, + ), + ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", token="hashed-shared"), + raw_headers={"x-litellm-api-key": "sk-shared", "authorization": "Bearer entra-bob"}, + ), + id="shared-key-different-obo-subjects", + ), + pytest.param( + {"transport": MCPTransport.stdio, "command": "srv", "env": {"WS": "${X-WS}"}}, + ListedToolsCaller(raw_headers={"X-WS": "A"}), + ListedToolsCaller(raw_headers={"X-WS": "B"}), + id="header-driven-stdio-env", + ), + ], + ) + def test_upstream_identity_inputs_keep_listed_tools_apart(self, server_kwargs, caller_a, caller_b): + manager = MCPServerManager() + server = MCPServer( + **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} + ) + manager._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.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.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"), + ) + + 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( + ("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): + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers + ) + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="alice", token="hashed-alice")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="bob", token="hashed-bob")) + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=signer, + ): + manager._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 + + def test_signed_server_slot_splits_on_the_callers_key_not_only_the_user(self): + """Two keys sharing a user_id get different signed JWTs, so they split; the same key + presented again lands on its own slot.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-beta")) + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=MagicMock(), + ): + manager._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_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + listed = manager.get_listed_tool(server, "srv-turn", same_key) + + 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() + 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): """ @@ -9597,15 +10226,11 @@ class TestGetPublicMCPServers: if registered_in == "both" else server ) - manager.config_mcp_servers = ( - {server.server_id: config_server} if registered_in in ("config", "both") else {} - ) + manager.config_mcp_servers = {server.server_id: config_server} if registered_in in ("config", "both") else {} manager.registry = {server.server_id: server} if registered_in in ("database", "both") else {} original_server: Final = server.model_dump() original_config_server: Final = config_server.model_dump() - expected_public: Final = registered_in != "neither" and ( - public_ids == [server.server_id] or implicitly_public - ) + expected_public: Final = registered_in != "neither" and (public_ids == [server.server_id] or implicitly_public) with ( patch("litellm.public_mcp_servers", public_ids), @@ -9613,9 +10238,7 @@ class TestGetPublicMCPServers: ): public_servers: Final = manager.get_public_mcp_servers() assert manager.is_mcp_server_public(server.server_id) is expected_public - assert [item.server_id for item in public_servers] == ( - [server.server_id] if expected_public else [] - ) + assert [item.server_id for item in public_servers] == ([server.server_id] if expected_public else []) assert manager.is_mcp_server_public("server-alias") is False assert manager.is_mcp_server_public("missing-server") is False assert manager.is_mcp_server_public(server.server_id, public_ids=frozenset()) is ( @@ -10514,7 +11137,9 @@ class TestOBOConcurrencyLimit: inflight = {"current": 0, "peak": 0} class _ConcurrencyRecordingClient: - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False): + async def call_tool( + self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False + ): inflight["current"] += 1 inflight["peak"] = max(inflight["peak"], inflight["current"]) try: @@ -12787,7 +13412,9 @@ class TestConfigServerIdPinning: @pytest.mark.asyncio @pytest.mark.parametrize("aliasing_entry_first", [True, False]) - async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected(self, config_only_mcp_manager_factory, aliasing_entry_first: bool): + async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected( + self, config_only_mcp_manager_factory, aliasing_entry_first: bool + ): """A grant naming 'docs_server' reaches both servers unpinned; the pin would narrow it to one.""" manager = config_only_mcp_manager_factory() wiki = ( @@ -12803,7 +13430,9 @@ class TestConfigServerIdPinning: await manager.load_servers_from_config(dict((wiki, docs) if aliasing_entry_first else (docs, wiki))) @pytest.mark.asyncio - async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected(self, config_only_mcp_manager_factory): + async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected( + self, config_only_mcp_manager_factory + ): manager = config_only_mcp_manager_factory() with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"): @@ -12891,7 +13520,9 @@ class TestConfigServerIdPinning: assert second_round == first_round @pytest.mark.asyncio - async def test_shadow_warning_fires_again_when_the_shadowed_set_changes(self, config_only_mcp_manager_factory, caplog): + async def test_shadow_warning_fires_again_when_the_shadowed_set_changes( + self, config_only_mcp_manager_factory, caplog + ): manager = config_only_mcp_manager_factory() await manager.load_servers_from_config(self._config(server_id="docs-prod-1")) @@ -13062,7 +13693,9 @@ class TestConfigServerIdPinning: assert manager.config_mcp_servers["wiki"].url == "https://example.com/mcp" @pytest.mark.asyncio - async def test_a_row_that_shadows_one_id_still_reports_capturing_another(self, config_only_mcp_manager_factory, caplog): + async def test_a_row_that_shadows_one_id_still_reports_capturing_another( + self, config_only_mcp_manager_factory, caplog + ): """Skipping is per identifier, not per row, so the second collision is not lost.""" manager = config_only_mcp_manager_factory() await manager.load_servers_from_config( @@ -13434,7 +14067,8 @@ async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, ("none", {"Authorization": "Bearer injected"}, "extra-headers", "Bearer injected"), ], ) -async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_request_ctx, +async def test_debug_resolution_matches_final_header_conflict_winner( + _mcp_request_ctx, config: Literal["stored", "static", "none"], extra_headers: dict[str, str] | None, expected_source: str, @@ -13505,7 +14139,9 @@ async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_reques @pytest.mark.asyncio @pytest.mark.parametrize("transport", ["http", "stdio"]) -async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ctx, transport: Literal["http", "stdio"]) -> None: +async def test_debug_reports_legacy_signing_and_non_http_transport( + _mcp_request_ctx, transport: Literal["http", "stdio"] +) -> None: from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from starlette.requests import Request @@ -13549,12 +14185,16 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + server_id="temporary-oauth-discovery", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, ) manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", ) with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: @@ -13574,13 +14214,18 @@ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publi async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", + server_id="repeated-stale", + name="stale", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + oauth2_flow="authorization_code", ) manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) with ( patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, @@ -13600,13 +14245,20 @@ async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + server_id="resolved-replacement", + name="replacement", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy( + update={ + "url": "https://new.example.com/mcp", + "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + } ) - replacement: Final = original.model_copy(update={ - "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", - "token_url": "https://new.example.com/token", - }) manager.registry[original.server_id] = replacement assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement @@ -13614,8 +14266,11 @@ async def test_stale_discovery_falls_back_to_resolved_registered_server() -> Non def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="stale-publication", name="publication", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + server_id="stale-publication", + name="publication", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) original_slot: Final = manager._oauth_discovery_slot(original.server_id) @@ -13631,9 +14286,13 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + server_id="expiring-session", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) @@ -13734,7 +14393,9 @@ async def test_openapi_health_reports_size_limit_as_unknown_and_caches_failure(r result = await manager.health_check_server(server.server_id) cached = await manager.health_check_server(server.server_id) assert result.status == "unknown" - assert result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + assert ( + result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + ) assert cached.health_check_error == result.health_check_error assert cached.last_health_check == result.last_health_check assert route.call_count == 1 @@ -13746,8 +14407,11 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( - server_id="cancelled-cache", name="cancelled-cache", transport=MCPTransport.http, - spec_path="https://93.184.216.34/cancelled-cache.json", auth_type=MCPAuth.none, + server_id="cancelled-cache", + name="cancelled-cache", + transport=MCPTransport.http, + spec_path="https://93.184.216.34/cancelled-cache.json", + auth_type=MCPAuth.none, ) manager.registry = {server.server_id: server} started = asyncio.Event() @@ -13876,7 +14540,9 @@ class _DiscoveryUpstream: def _discovery_server() -> MCPServer: - return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) + return MCPServer( + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http + ) @pytest.mark.asyncio @@ -14044,7 +14710,9 @@ async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) assert upstream.initializes == 2 -@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5))) +@pytest.mark.parametrize( + "value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)) +) def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl @@ -14361,26 +15029,45 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them( class TestProtectedCredentialPreparation: @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,credential", [ - (MCPAuth.bearer_token, None), - (MCPAuth.bearer_token, "Bearer"), - (MCPAuth.api_key, None), - (MCPAuth.basic, "Basic"), - ]) + @pytest.mark.parametrize( + "auth_type,credential", + [ + (MCPAuth.bearer_token, None), + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.api_key, None), + (MCPAuth.basic, "Basic"), + ], + ) @pytest.mark.parametrize("dispatch", ["managed", "local"]) async def test_openapi_dispatch_rejects_unusable_effective_credentials( - self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - auth_type: MCPAuthType, credential: str | None, dispatch: str, + self, + tmp_path: Path, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + auth_type: MCPAuthType, + credential: str | None, + dispatch: str, ) -> None: from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix spec_path: Final = tmp_path / "openapi.json" - spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"}, - "paths": {"/echo": {"get": {"operationId": "echo"}}}})) + spec_path.write_text( + json.dumps( + { + "openapi": "3.0.0", + "info": {"title": "Auth", "version": "1"}, + "paths": {"/echo": {"get": {"operationId": "echo"}}}, + } + ) + ) server: Final = MCPServer( - server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential, + server_id="dispatch-auth", + name="dispatch-auth", + url="https://upstream.example", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=credential, ) manager: Final = MCPServerManager() await manager._register_openapi_tools(str(spec_path), server, server.url) @@ -14403,14 +15090,21 @@ class TestProtectedCredentialPreparation: self, transport: MCPTransport, client_secret: str | None, subject: str | None ) -> None: server = MCPServer( - server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp", - transport=transport, auth_type=MCPAuth.oauth2_token_exchange, - client_id="gateway", client_secret=client_secret, - token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback", + server_id="incomplete-obo", + name="incomplete-obo", + url="https://upstream.example/mcp", + transport=transport, + auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway", + client_secret=client_secret, + token_exchange_endpoint="https://idp.example/token", + authentication_token="static-fallback", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header="Bearer override", subject_token=subject, + server, + mcp_auth_header="Bearer override", + subject_token=subject, ) assert exc.value.status_code == (401 if subject is None else 500) assert "static-fallback" not in str(exc.value.detail) @@ -14423,8 +15117,11 @@ class TestProtectedCredentialPreparation: self, auth_type: MCPAuthType, credential: str | dict[str, str] | None ) -> None: server = MCPServer( - server_id="empty-static", name="empty-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-static", + name="empty-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential) @@ -14432,16 +15129,22 @@ class TestProtectedCredentialPreparation: assert "credential" in str(exc.value.detail).lower() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,headers", [ - (MCPAuth.api_key, {"X-API-Key": "key"}), - (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), - ]) + @pytest.mark.parametrize( + "auth_type,headers", + [ + (MCPAuth.api_key, {"X-API-Key": "key"}), + (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), + ], + ) async def test_static_auth_accepts_actual_forwarded_credential( self, auth_type: MCPAuthType, headers: dict[str, str] ) -> None: server = MCPServer( - server_id="header-static", name="header-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="header-static", + name="header-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers) assert client._get_auth_headers() == headers @@ -14450,29 +15153,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange]) async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None: server = MCPServer( - server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="openapi-empty", + name="openapi-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, token_exchange_endpoint="https://idp.example/token", ) with pytest.raises(HTTPException) as exc: await MCPServerManager().resolve_openapi_upstream_auth( - mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None, - user_api_key_auth=None, forwarded_headers=None, + mcp_server=server, + oauth2_headers=None, + raw_headers=None, + mcp_auth_header=None, + user_api_key_auth=None, + forwarded_headers=None, ) assert exc.value.status_code in (401, 500) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,slot,value", [ - (MCPAuth.api_key, "X-API-Key", "token"), - (MCPAuth.authorization, "Authorization", "opaque-secret-value"), - (MCPAuth.authorization, "Authorization", "Bearer abc"), - (MCPAuth.authorization, "Authorization", "Custom abc"), - ]) + @pytest.mark.parametrize( + "auth_type,slot,value", + [ + (MCPAuth.api_key, "X-API-Key", "token"), + (MCPAuth.authorization, "Authorization", "opaque-secret-value"), + (MCPAuth.authorization, "Authorization", "Bearer abc"), + (MCPAuth.authorization, "Authorization", "Custom abc"), + ], + ) async def test_raw_static_credentials_are_forwarded_unchanged( - self, auth_type: MCPAuthType, slot: str, value: str, + self, + auth_type: MCPAuthType, + slot: str, + value: str, ) -> None: - server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value) + server = MCPServer( + server_id="raw-key", + name="raw-key", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, + ) client = await MCPServerManager()._create_mcp_client(server) assert client._resolved_auth is not None request = httpx.Request("GET", server.url) @@ -14486,17 +15208,24 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"]) @pytest.mark.parametrize("source", ["configured", "caller", "forwarded"]) async def test_raw_authorization_rejects_bare_schemes_before_dispatch( - self, respx_mock: MockRouter, value: str, source: str, + self, + respx_mock: MockRouter, + value: str, + source: str, ) -> None: server: Final = MCPServer( - server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.authorization, + server_id="raw-empty", + name="raw-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.authorization, authentication_token=value if source == "configured" else None, ) destination: Final = respx_mock.route().respond(200) with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, + server, + mcp_auth_header=value if source == "caller" else None, extra_headers={"Authorization": value} if source == "forwarded" else None, ) assert exc.value.status_code == 500 @@ -14504,9 +15233,15 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None: - server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True, - token_exchange_endpoint="https://idp.example/token") + server = MCPServer( + server_id="obo-byok", + name="obo-byok", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + is_byok=True, + token_exchange_endpoint="https://idp.example/token", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override") assert exc.value.status_code == 401 @@ -14514,41 +15249,66 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")]) async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None: - server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured) + server = MCPServer( + server_id="override", + name="override", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=configured, + ) client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override) assert client._get_auth_headers()["Authorization"] == override @pytest.mark.asyncio @pytest.mark.parametrize("token", [None, "shared"]) async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None: - server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token) + server = MCPServer( + server_id="empty-header", + name="empty-header", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=token, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "}) assert exc.value.status_code == 500 @pytest.mark.asyncio async def test_custom_slot_uses_its_actual_credential(self) -> None: - server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, - upstream_token_header="X-Custom", authentication_token="key") + server = MCPServer( + server_id="custom", + name="custom", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", + authentication_token="key", + ) client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"}) assert client._credential_slot == "X-Custom" assert await client.discovery_auth_fingerprint() @pytest.mark.asyncio - @pytest.mark.parametrize("static_headers,accepted", [ - ({"apikey": "static-key"}, True), - ({"apikey": ""}, False), - ({"X-Tenant": "tenant"}, True), - ]) + @pytest.mark.parametrize( + "static_headers,accepted", + [ + ({"apikey": "static-key"}, True), + ({"apikey": ""}, False), + ({"X-Tenant": "tenant"}, True), + ], + ) async def test_api_key_carried_by_static_header_passes_fail_closed_check( self, static_headers: dict[str, str], accepted: bool ) -> None: server: Final = MCPServer( - server_id="static-slot", name="static-slot", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, static_headers=static_headers, + server_id="static-slot", + name="static-slot", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + static_headers=static_headers, ) if not accepted: with pytest.raises(HTTPException) as exc: @@ -14560,21 +15320,36 @@ class TestProtectedCredentialPreparation: assert all(request.headers[name] == value for name, value in static_headers.items()) @pytest.mark.asyncio - @pytest.mark.parametrize("static,forwarded,caller", [ - ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), - ({}, {"X-API-Key": "forwarded"}, None), - ({}, None, "ApiKey caller"), - ({"X-API-Key": "static"}, {"Authorization": ""}, None), - ]) + @pytest.mark.parametrize( + "static,forwarded,caller", + [ + ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), + ({}, {"X-API-Key": "forwarded"}, None), + ({}, None, "ApiKey caller"), + ({"X-API-Key": "static"}, {"Authorization": ""}, None), + ], + ) async def test_openapi_static_credentials_remain_supported( - self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None + self, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + static: dict[str, str], + forwarded: dict[str, str] | None, + caller: str | None, ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, _request_extra_headers, create_tool_function, + _request_auth_header, + _request_extra_headers, + create_tool_function, ) + tool: Final = create_tool_function( - "/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key, + "/echo", + "get", + {}, + "https://upstream.example", + headers=static, + auth_type=MCPAuth.api_key, ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") @@ -14608,8 +15383,13 @@ class TestProtectedCredentialPreparation: self.closed = True auth = CancelledAuth() - server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key) + server = MCPServer( + server_id="cancel", + name="cancel", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + ) client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth) with pytest.raises(asyncio.CancelledError): await prepare_mcp_client(server, client) @@ -14618,8 +15398,14 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization]) async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None: - server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ") + server = MCPServer( + server_id="blank-static", + name="blank-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=" ", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server) assert exc.value.status_code == 500 @@ -14627,8 +15413,13 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="]) async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None: - server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic) + server = MCPServer( + server_id="bad-basic", + name="bad-basic", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header}) assert exc.value.status_code == 500 @@ -14637,34 +15428,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None: - server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, - authentication_token=value if source == "configured" else None) + server = MCPServer( + server_id="basic-scheme", + name="basic-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value if source == "configured" else None, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None) assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,default_slot", [ - (MCPAuth.api_key, "fixture-key", "X-API-Key"), - (MCPAuth.bearer_token, "fixture-key", "Authorization"), - (MCPAuth.basic, "user:pass", "Authorization"), - (MCPAuth.token, "fixture-key", "Authorization"), - (MCPAuth.authorization, "fixture-key", "Authorization"), - ]) + @pytest.mark.parametrize( + "auth_type,value,default_slot", + [ + (MCPAuth.api_key, "fixture-key", "X-API-Key"), + (MCPAuth.bearer_token, "fixture-key", "Authorization"), + (MCPAuth.basic, "user:pass", "Authorization"), + (MCPAuth.token, "fixture-key", "Authorization"), + (MCPAuth.authorization, "fixture-key", "Authorization"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_usable_credential_survives_an_empty_alternate_header( self, auth_type: MCPAuthType, value: str, default_slot: str, source: str ) -> None: server: Final = MCPServer( - server_id="alternate", name="alternate", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom", + server_id="alternate", + name="alternate", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + upstream_token_header="X-Custom", authentication_token=value if source == "configured" else None, ) empty_slot: Final = default_slot if source == "configured" else "X-Custom" selected_slot: Final = "X-Custom" if source == "configured" else default_slot client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""}, + server, + mcp_auth_header=value if source == "caller" else None, + extra_headers={empty_slot: ""}, ) request: Final = await client.prepare_request_auth() assert request.headers[selected_slot] @@ -14673,8 +15478,12 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None: server: Final = MCPServer( - server_id="both-empty", name="both-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom", + server_id="both-empty", + name="both-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""}) @@ -14687,12 +15496,17 @@ class TestProtectedCredentialPreparation: self, custom_slot: str | None, source: str ) -> None: server: Final = MCPServer( - server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot, + server_id="caller-auth", + name="caller-auth", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header=custom_slot, ) headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""} client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=headers if source == "caller" else None, + server, + mcp_auth_header=headers if source == "caller" else None, extra_headers=headers if source == "forwarded" else None, ) request: Final = await client.prepare_request_auth() @@ -14701,14 +15515,29 @@ class TestProtectedCredentialPreparation: assert custom_slot is None or custom_slot not in request.headers @pytest.mark.asyncio - @pytest.mark.parametrize("value", [ - "", " ", "Bearer", "Basic", "token", "ApiKey", - "Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY", - ]) + @pytest.mark.parametrize( + "value", + [ + "", + " ", + "Bearer", + "Basic", + "token", + "ApiKey", + "Bearer Bearer", + "ApiKey ApiKey", + "token token", + "bEaReR BEARER", + "aPiKeY\tAPIKEY", + ], + ) async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None: server: Final = MCPServer( - server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, + server_id="caller-empty", + name="caller-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value}) @@ -14719,8 +15548,11 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None: server: Final = MCPServer( - server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, + server_id="basic-pair", + name="basic-pair", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14733,8 +15565,12 @@ class TestProtectedCredentialPreparation: import base64 server: Final = MCPServer( - server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value, + server_id="basic-valid", + name="basic-valid", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14743,17 +15579,27 @@ class TestProtectedCredentialPreparation: assert base64.b64decode(encoded) == value.encode() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value", [ - (MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"), - (MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"), - ]) + @pytest.mark.parametrize( + "auth_type,value", + [ + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.bearer_token, "Bearer "), + (MCPAuth.bearer_token, "bearer"), + (MCPAuth.token, "token"), + (MCPAuth.token, "token "), + (MCPAuth.token, "TOKEN"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix( self, auth_type: MCPAuthType, value: str, source: str ) -> None: server: Final = MCPServer( - server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-scheme", + name="empty-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14761,17 +15607,24 @@ class TestProtectedCredentialPreparation: assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,expected", [ - (MCPAuth.bearer_token, "token", "Bearer token"), - (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), - (MCPAuth.token, "tokenish", "token tokenish"), - ]) + @pytest.mark.parametrize( + "auth_type,value,expected", + [ + (MCPAuth.bearer_token, "token", "Bearer token"), + (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), + (MCPAuth.token, "tokenish", "token tokenish"), + ], + ) async def test_static_credentials_that_resemble_schemes_remain_usable( self, auth_type: MCPAuthType, value: str, expected: str ) -> None: server: Final = MCPServer( - server_id="real-token", name="real-token", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value, + server_id="real-token", + name="real-token", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14810,16 +15663,31 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon registry.register_tool("observer-execute", "Execute", {"type": "object"}, upstream) monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) manager = MCPServerManager() - manager.registry = {"observer": MCPServer( - server_id="observer", name="observer", server_name="observer", transport="http", - url="https://observer.example/mcp", spec_path="observer.json", auth_type="none", - )} + manager.registry = { + "observer": MCPServer( + server_id="observer", + name="observer", + server_name="observer", + transport="http", + url="https://observer.example/mcp", + spec_path="observer.json", + auth_type="none", + ) + } manager.tool_name_to_mcp_server_name_mapping = {"observer-execute": "observer"} - result = await asyncio.wait_for(manager.call_tool( - server_name="observer", name="execute", arguments={"text": "hello"}, - user_api_key_auth=UserAPIKeyAuth(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), - guardrail_context=MCPRequestContext.resolve_guardrail_context({"metadata": {"guardrails": ["observe"] if selected else []}}), - ), timeout=5) + result = await asyncio.wait_for( + manager.call_tool( + server_name="observer", + name="execute", + arguments={"text": "hello"}, + user_api_key_auth=UserAPIKeyAuth(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + guardrail_context=MCPRequestContext.resolve_guardrail_context( + {"metadata": {"guardrails": ["observe"] if selected else []}} + ), + ), + timeout=5, + ) assert tool_started.is_set() assert guardrail_started.is_set() is selected assert result.is_error is False @@ -14848,11 +15716,21 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback - upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) + upstream = MCPServer( + server_id="explicit-empty", + name="explicit_empty", + url="https://example.invalid/mcp", + transport=MCPTransport.http, + allow_sampling=True, + ) token = auth_context_var.set(None) sampling = AsyncMock() try: - legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99") + legacy_server.set_auth_context( + UserAPIKeyAuth(user_id="unrelated"), + raw_headers={"authorization": "unrelated-credential"}, + client_ip="192.0.2.99", + ) with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory, patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), @@ -14860,7 +15738,9 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie if legacy_factory: callback = _create_sampling_callback(user_api_key_auth=UserAPIKeyAuth(user_id="explicit")) else: - await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) + await MCPServerManager()._create_mcp_client( + upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None + ) callback = factory.call_args.kwargs["sampling_callback"] await callback(None, None) captured = sampling.await_args.kwargs @@ -14883,16 +15763,28 @@ class TestSharedIdentifierPrefixWarning: manager = MCPServerManager() rows = [ LiteLLM_MCPServerTable( - server_id="srv-a", server_name="alpha", alias="shared", url="https://a.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-a", + server_name="alpha", + alias="shared", + url="https://a.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-b", server_name="beta", alias="Shared", url="https://b.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-b", + server_name="beta", + alias="Shared", + url="https://b.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-c", server_name="gamma", alias="lonely", url="https://c.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-c", + server_name="gamma", + alias="lonely", + url="https://c.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), ] raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] @@ -14937,7 +15829,9 @@ class TestSharedIdentifierPrefixWarning: @pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]) async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision): manager = config_only_mcp_manager_factory() - await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}}) + await manager.load_servers_from_config( + {"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}} + ) server = next(iter(manager.config_mcp_servers.values())) client = await manager._create_mcp_client(server) assert server.protocol_version == revision @@ -14949,11 +15843,15 @@ async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_m def test_runtime_protocol_metadata_preserves_explicit_precedence( revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None ) -> None: - server: Final = MCPServer.model_validate({ - "server_id": "preview", "name": "preview", "transport": "http", - "mcp_info": {"protocol_version": revision}, - **({"protocol_version": explicit} if explicit is not None else {}), - }) + server: Final = MCPServer.model_validate( + { + "server_id": "preview", + "name": "preview", + "transport": "http", + "mcp_info": {"protocol_version": revision}, + **({"protocol_version": explicit} if explicit is not None else {}), + } + ) assert server.protocol_version == (explicit if explicit is not None else revision) 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/_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() 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 23131321938..e72b716665c 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 @@ -346,6 +346,29 @@ class TestAllowFlow: assert evaluate_call.json["conversationId"] == "sess-123" assert evaluate_call.json["agentId"] == "my-agent-key" + @pytest.mark.asyncio + async def test_evaluate_payload_includes_listed_tool_metadata(self): + handler: Final = FakeHandler([_token_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler) + schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]} + await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema)) + assert handler.calls[1].json["tool"] == { + "name": "send_email", + "description": "Send an email", + "inputSchema": schema, + } + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("description", "schema"), + [(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")], + ) + async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema): + handler: Final = FakeHandler([_token_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler) + await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema)) + assert handler.calls[1].json["tool"] == {"name": "send_email"} + @pytest.mark.asyncio async def test_non_mcp_call_type_skipped(self): handler: Final = FakeHandler([]) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py index c5bd89645b4..25be3b5de6b 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_input_schema"]) == ("Adds numbers", schema) + + +def test_mcp_tool_metadata_absent_when_tool_was_never_listed(proxy_logging): + obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={"name": "calc", "arguments": {}}) + out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={}) + assert "mcp_tool_description" not in out and "mcp_input_schema" not in out + + def test_create_mcp_request_object_from_kwargs_empty(proxy_logging): obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={}) snapshot = {