diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b2128cb0553..0d28d4d26c4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -253,6 +253,76 @@ def _without_authorization( return filtered or None +def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str: + """Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.""" + if mcp_server.auth_type == MCPAuth.api_key: + return f"ApiKey {mcp_auth_header}" + if mcp_server.auth_type == MCPAuth.basic: + return f"Basic {mcp_auth_header}" + return f"Bearer {mcp_auth_header}" + + +def _openapi_forwarded_extra_headers( + mcp_server: MCPServer, + raw_headers: Optional[dict[str, str]], + user_api_key_auth: Optional[UserAPIKeyAuth], +) -> Optional[dict[str, str]]: + if not mcp_server.extra_headers or not raw_headers: + return None + normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} + skip_caller_authorization = _should_strip_caller_authorization( + mcp_server=mcp_server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + forwarded: dict[str, str] = {} + for header_name in mcp_server.extra_headers: + if not isinstance(header_name, str): + continue + if skip_caller_authorization and header_name.lower() == "authorization": + continue + value = normalized_raw.get(header_name.lower()) + if value is not None: + forwarded[header_name] = value + return forwarded or None + + +async def _resolve_byok_mcp_auth_header( + mcp_server: MCPServer, + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_auth_header: Optional[str], +) -> Optional[str]: + """Resolve BYOK credential for tool calls that bypass ``execute_mcp_tool``.""" + if not mcp_server.is_byok: + return mcp_auth_header + + from litellm.proxy._experimental.mcp_server.server import ( + _check_byok_credential, + _get_byok_credential, + ) + + if not mcp_auth_header: + byok_cred = await _get_byok_credential(mcp_server, user_api_key_auth) + if byok_cred is None: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": ( + "No stored credential found for this BYOK server. " + "Complete the OAuth authorization flow to provide your API key." + ), + }, + headers={"WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"'}, + ) + return byok_cred + + await _check_byok_credential(mcp_server, user_api_key_auth) + return mcp_auth_header + + def _extract_upstream_auth_failure( exc: BaseException, ) -> Optional[tuple[int, Optional[str]]]: @@ -3861,6 +3931,15 @@ class MCPServerManager: start_time = datetime.datetime.now() mcp_server = self._resolve_mcp_server_for_tool_call(server_name, name) + # Resolved before any hook runs so a missing BYOK credential (401) never + # leaves during-hook side effects (audit logging, rate-limit bookkeeping) + # recorded against a call that ultimately fails. + mcp_auth_header = await _resolve_byok_mcp_auth_header( + mcp_server, + user_api_key_auth, + mcp_auth_header, + ) + ######################################################### # Pre MCP Tool Call Hook # Allow validation and modification of tool calls before execution @@ -3907,9 +3986,25 @@ class MCPServerManager: server_name, ) + auth_header_value = ( + _format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None + ) + forwarded_headers = _openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth) + async def _call_openapi_via_handler(): - async with self._limit_outbound_concurrency(mcp_server): - return await self._call_openapi_tool_handler(mcp_server, name, arguments) + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_auth_header, + _request_extra_headers, + ) + + auth_token = _request_auth_header.set(auth_header_value) + extra_token = _request_extra_headers.set(forwarded_headers) + try: + async with self._limit_outbound_concurrency(mcp_server): + return await self._call_openapi_tool_handler(mcp_server, name, arguments) + finally: + _request_auth_header.reset(auth_token) + _request_extra_headers.reset(extra_token) tasks.append(asyncio.create_task(_call_openapi_via_handler())) else: @@ -4553,6 +4648,8 @@ class MCPServerManager: teams=[], mcp_access_groups=server.access_groups or [], allowed_tools=server.allowed_tools or [], + tool_name_to_display_name=server.tool_name_to_display_name, + tool_name_to_description=server.tool_name_to_description, extra_headers=server.extra_headers or [], mcp_info=server.mcp_info, static_headers=server.static_headers, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index ce0698cb7ac..d482e537c5d 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -77,6 +77,7 @@ if MCP_AVAILABLE: ListMCPToolsRestAPIResponseObject, MCPInfo, MCPServer, + _apply_toolset_scope, _fire_mcp_success_logging, _tool_name_matches, execute_mcp_tool, @@ -541,10 +542,37 @@ if MCP_AVAILABLE: "message": "Successfully retrieved tools", } + def _as_query_str(value: Any) -> Optional[str]: + """Coerce an Optional[str] Query param to str|None, dropping unresolved FastAPI defaults.""" + return value if isinstance(value, str) else None + + async def _resolve_toolset_scope( + toolset_name: Optional[str], + user_api_key_dict: UserAPIKeyAuth, + ) -> UserAPIKeyAuth: + """Resolve ``toolset_name`` to its scoped ``UserAPIKeyAuth``, or return unchanged.""" + if not toolset_name: + return user_api_key_dict + + from litellm.proxy.utils import get_prisma_client_or_throw + + prisma_client = get_prisma_client_or_throw("Database not available. Connect a database to your proxy") + toolset = await global_mcp_server_manager.get_toolset_by_name_cached(prisma_client, toolset_name) + if toolset is None: + raise HTTPException( + status_code=404, + detail=f"Toolset '{toolset_name}' not found", + ) + return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id) + @router.get("/tools/list", dependencies=[Depends(user_api_key_auth)]) async def list_tool_rest_api( request: Request, server_id: Optional[str] = Query(None, description="The server id to list tools for"), + mcp_server_name: Optional[str] = Query( + None, description="Filter tools to a single MCP server by name or alias" + ), + toolset_name: Optional[str] = Query(None, description="Filter tools to a single toolset by name"), include_disabled_tools: bool = Query( False, description=( @@ -582,16 +610,29 @@ if MCP_AVAILABLE: ) try: + mcp_server_name = _as_query_str(mcp_server_name) + toolset_name = _as_query_str(toolset_name) + # The full catalog (allowlist filter skipped) is admin-only so the # REST endpoint can't be used to enumerate deliberately-disabled tools. apply_tool_filters = not ( include_disabled_tools and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN ) - if apply_tool_filters and getattr( - getattr(user_api_key_dict, "object_permission", None), - "mcp_tool_search_enabled", - False, + user_api_key_dict = await _resolve_toolset_scope(toolset_name, user_api_key_dict) + + if server_id is None: + server_id = mcp_server_name + + if ( + apply_tool_filters + and server_id is None + and toolset_name is None + and getattr( + getattr(user_api_key_dict, "object_permission", None), + "mcp_tool_search_enabled", + False, + ) ): from litellm.proxy._experimental.mcp_server.tool_search import ( get_virtual_tool_definitions, @@ -719,6 +760,8 @@ if MCP_AVAILABLE: request_path=request.scope.get("_original_path") or request.url.path, ) except HTTPException as http_exc: + if http_exc.status_code == status.HTTP_404_NOT_FOUND: + raise # Internal access/IP 403s keep the legacy error-dict response shape # so the existing contract stays intact. verbose_logger.exception("HTTPException in list_tool_rest_api: %s", str(http_exc)) diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 9cb6d404b01..c9c60030dbc 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -214,12 +214,14 @@ def server_applies_tool_allowlist(mcp_server: Any) -> bool: def validate_and_normalize_mcp_server_payload(payload: Any) -> None: """ - Validate and normalize MCP server payload fields (server_name and alias). + Validate and normalize MCP server payload fields (server_name, alias, and + tool_name_to_display_name). This function: 1. Validates that server_name and alias don't contain the MCP_TOOL_PREFIX_SEPARATOR - 2. Normalizes alias by replacing spaces with underscores - 3. Sets default alias if not provided (using server_name as base) + 2. Validates that tool_name_to_display_name values satisfy Bedrock's tool-name pattern + 3. Normalizes alias by replacing spaces with underscores + 4. Sets default alias if not provided (using server_name as base) Args: payload: The payload object containing server_name and alias fields @@ -235,6 +237,10 @@ def validate_and_normalize_mcp_server_payload(payload: Any) -> None: if hasattr(payload, "alias") and payload.alias: validate_mcp_server_name(payload.alias, raise_http_exception=True) + # Tool display name validation: must satisfy Bedrock's tool-name pattern + if hasattr(payload, "tool_name_to_display_name") and payload.tool_name_to_display_name: + validate_tool_display_names(payload.tool_name_to_display_name) + # Alias normalization and defaulting alias = getattr(payload, "alias", None) server_name = getattr(payload, "server_name", None) @@ -409,6 +415,42 @@ def validate_mcp_server_name(server_name: str, raise_http_exception: bool = Fals raise Exception(error_message) +TOOL_DISPLAY_NAME_PATTERN = re.compile(r"^[a-zA-Z0-9_-]+$") + + +def validate_tool_display_names(tool_name_to_display_name: Optional[Mapping[str, str]]) -> None: + """ + Validate tool display name overrides against Bedrock's tool-name constraint. + + A display name replaces the tool name sent to the LLM provider, so it must + satisfy the strictest provider requirement in use (Bedrock's + ``[a-zA-Z0-9_-]+``); a name with spaces or other characters saves + successfully but fails every subsequent Bedrock tool call. + + Raises: + HTTPException: If any display name fails the pattern. + """ + if not tool_name_to_display_name: + return + + for original_name, display_name in tool_name_to_display_name.items(): + if display_name and not TOOL_DISPLAY_NAME_PATTERN.match(display_name): + from fastapi import HTTPException + from starlette import status + + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": ( + f"Invalid display name '{display_name}' for tool '{original_name}'. " + "Display names may only contain letters, digits, underscores, and " + "hyphens (no spaces or other special characters), since they replace " + "the tool name sent to the LLM provider." + ) + }, + ) + + class MCPMissingUserEnvVarsError(Exception): """Raised when an MCP request can't be built because the calling user has not supplied one or more required per-user environment variables. diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 999945b3823..c1cf2d967eb 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -16,7 +16,10 @@ from typing import ( from litellm._logging import verbose_logger from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.proxy._experimental.mcp_server.utils import split_server_prefix_from_name +from litellm.proxy._experimental.mcp_server.utils import ( + split_server_prefix_from_name, + strip_known_server_prefix, +) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.responses.main import aresponses from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator @@ -628,6 +631,9 @@ class LiteLLM_Proxy_MCP_Handler: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) + from litellm.proxy._experimental.mcp_server.server import ( + _resolve_display_name_to_original, + ) from litellm.proxy.proxy_server import proxy_logging_obj tool_results = [] @@ -654,11 +660,13 @@ class LiteLLM_Proxy_MCP_Handler: server_name = tool_server_map[tool_name] - # Remove the server name prefix if the tool name includes it. - sanitized_tool_name = tool_name - unprefixed_name, prefixed_server_name = split_server_prefix_from_name(tool_name) - if prefixed_server_name and prefixed_server_name == server_name and unprefixed_name: - sanitized_tool_name = unprefixed_name + mcp_server = global_mcp_server_manager.get_mcp_server_by_name( + server_name + ) or global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) + resolved_tool_name = ( + _resolve_display_name_to_original(tool_name, [mcp_server]) if mcp_server else tool_name + ) + sanitized_tool_name = strip_known_server_prefix(resolved_tool_name, mcp_server) start_time = datetime.now() logging_input = [ @@ -741,7 +749,6 @@ class LiteLLM_Proxy_MCP_Handler: "arguments": parsed_arguments, "namespaced_tool_name": tool_name, } - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) if mcp_server: mcp_info = mcp_server.mcp_info or {} standard_logging_mcp_tool_call["mcp_server_name"] = ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 363948ff4e6..73486fe0b6a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -43,17 +43,13 @@ class TestConvertMcpHookResponseToKwargs: def test_extracts_modified_arguments(self): original = {"arguments": {"old": "value"}} response = {"modified_arguments": {"new": "value"}} - result = self.proxy_logging._convert_mcp_hook_response_to_kwargs( - response, original - ) + result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(response, original) assert result["arguments"] == {"new": "value"} def test_extracts_extra_headers(self): original = {"arguments": {"key": "val"}} response = {"extra_headers": {"Authorization": "Bearer signed-jwt"}} - result = self.proxy_logging._convert_mcp_hook_response_to_kwargs( - response, original - ) + result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(response, original) assert result["extra_headers"] == {"Authorization": "Bearer signed-jwt"} def test_extracts_both_arguments_and_headers(self): @@ -62,9 +58,7 @@ class TestConvertMcpHookResponseToKwargs: "modified_arguments": {"new": "value"}, "extra_headers": {"X-Custom": "header-val"}, } - result = self.proxy_logging._convert_mcp_hook_response_to_kwargs( - response, original - ) + result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(response, original) assert result["arguments"] == {"new": "value"} assert result["extra_headers"] == {"X-Custom": "header-val"} @@ -72,9 +66,7 @@ class TestConvertMcpHookResponseToKwargs: """Backward compat: hooks that only return modified_arguments still work.""" original = {"arguments": {"key": "val"}} response = {"modified_arguments": {"key": "new_val"}} - result = self.proxy_logging._convert_mcp_hook_response_to_kwargs( - response, original - ) + result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(response, original) assert "extra_headers" not in result assert result["arguments"] == {"key": "new_val"} @@ -82,9 +74,7 @@ class TestConvertMcpHookResponseToKwargs: """Empty dict for extra_headers is falsy and should not be set.""" original = {"arguments": {"key": "val"}} response = {"extra_headers": {}} - result = self.proxy_logging._convert_mcp_hook_response_to_kwargs( - response, original - ) + result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(response, original) assert "extra_headers" not in result @@ -107,18 +97,10 @@ class TestPreCallToolCheckReturnsHeaders: server = self._make_server() proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( - return_value=MagicMock() - ) - proxy_logging._convert_mcp_to_llm_format = MagicMock( - return_value={"model": "fake"} - ) - proxy_logging.pre_call_hook = AsyncMock( - return_value={"modified_arguments": {"key": "val"}} - ) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( - return_value={"arguments": {"key": "val"}} - ) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.pre_call_hook = AsyncMock(return_value={"modified_arguments": {"key": "val"}}) + proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {"key": "val"}}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -146,15 +128,9 @@ class TestPreCallToolCheckReturnsHeaders: hook_headers = {"Authorization": "Bearer signed-jwt", "X-Trace-Id": "abc123"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( - return_value=MagicMock() - ) - proxy_logging._convert_mcp_to_llm_format = MagicMock( - return_value={"model": "fake"} - ) - proxy_logging.pre_call_hook = AsyncMock( - return_value={"extra_headers": hook_headers} - ) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.pre_call_hook = AsyncMock(return_value={"extra_headers": hook_headers}) proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( return_value={"arguments": {"key": "val"}, "extra_headers": hook_headers} ) @@ -183,12 +159,8 @@ class TestPreCallToolCheckReturnsHeaders: server = self._make_server() proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( - return_value=MagicMock() - ) - proxy_logging._convert_mcp_to_llm_format = MagicMock( - return_value={"model": "fake"} - ) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): @@ -219,18 +191,10 @@ class TestPreCallToolCheckReturnsHeaders: modified_args = {"key": "modified", "extra": "added"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( - return_value=MagicMock() - ) - proxy_logging._convert_mcp_to_llm_format = MagicMock( - return_value={"model": "fake"} - ) - proxy_logging.pre_call_hook = AsyncMock( - return_value={"modified_arguments": modified_args} - ) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( - return_value={"arguments": modified_args} - ) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.pre_call_hook = AsyncMock(return_value={"modified_arguments": modified_args}) + proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": modified_args}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -260,12 +224,8 @@ class TestPreCallToolCheckReturnsHeaders: hook_headers = {"Authorization": "Bearer jwt"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( - return_value=MagicMock() - ) - proxy_logging._convert_mcp_to_llm_format = MagicMock( - return_value={"model": "fake"} - ) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"dummy": True}) proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( return_value={"arguments": modified_args, "extra_headers": hook_headers} @@ -345,9 +305,7 @@ class TestCallToolFlowsHookHeaders: mock_call.assert_called_once() call_kwargs = mock_call.call_args - assert ( - call_kwargs.kwargs.get("hook_extra_headers") == hook_headers - ) + assert call_kwargs.kwargs.get("hook_extra_headers") == hook_headers @pytest.mark.asyncio async def test_no_hook_headers_when_no_proxy_logging(self): @@ -434,9 +392,7 @@ class TestCallToolFlowsHookHeaders: spec_path="/path/to/spec.yaml", ) - with patch.object( - manager, "_get_mcp_server_from_tool_name", return_value=server - ): + with patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server): with patch.object( manager, "pre_call_tool_check", @@ -467,10 +423,7 @@ class TestCallToolFlowsHookHeaders: proxy_logging_obj=proxy_logging, ) mock_logger.warning.assert_called_once() - assert ( - "header injection is not supported" - in mock_logger.warning.call_args[0][0] - ) + assert "header injection is not supported" in mock_logger.warning.call_args[0][0] @pytest.mark.asyncio async def test_openapi_server_no_error_without_hook_headers(self): @@ -486,9 +439,7 @@ class TestCallToolFlowsHookHeaders: spec_path="/path/to/spec.yaml", ) - with patch.object( - manager, "_get_mcp_server_from_tool_name", return_value=server - ): + with patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server): with patch.object( manager, "pre_call_tool_check", @@ -539,25 +490,19 @@ class TestHookHeaderMergePriority: async def test_hook_headers_override_static_headers(self): """Hook headers should take precedence over static_headers.""" manager = MCPServerManager() - server = self._make_server( - static_headers={"Authorization": "Bearer static-token", "X-Static": "yes"} - ) + server = self._make_server(static_headers={"Authorization": "Bearer static-token", "X-Static": "yes"}) hook_headers = {"Authorization": "Bearer hook-signed-jwt"} captured_extra_headers: Dict[str, Any] = {} - async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs - ): + async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs): captured_extra_headers["value"] = extra_headers mock_client = MagicMock() mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object( - manager, "_create_mcp_client", side_effect=fake_create_mcp_client - ): + with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): with patch.object(manager, "_build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( @@ -587,17 +532,13 @@ class TestHookHeaderMergePriority: captured_extra_headers: Dict[str, Any] = {} - async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs - ): + async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs): captured_extra_headers["value"] = extra_headers mock_client = MagicMock() mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object( - manager, "_create_mcp_client", side_effect=fake_create_mcp_client - ): + with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): with patch.object(manager, "_build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( @@ -633,17 +574,13 @@ class TestHookHeaderMergePriority: captured_extra_headers: Dict[str, Any] = {} - async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs - ): + async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs): captured_extra_headers["value"] = extra_headers mock_client = MagicMock() mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object( - manager, "_create_mcp_client", side_effect=fake_create_mcp_client - ): + with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): with patch.object(manager, "_build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( @@ -689,17 +626,13 @@ class TestHookHeaderMergePriority: captured_extra_headers: Dict[str, Any] = {} - async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs - ): + async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs): captured_extra_headers["value"] = extra_headers mock_client = MagicMock() mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object( - manager, "_create_mcp_client", side_effect=fake_create_mcp_client - ): + with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): with patch.object(manager, "_build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( @@ -737,17 +670,13 @@ class TestHookHeaderMergePriority: captured_extra_headers: Dict[str, Any] = {} - async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs - ): + async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs): captured_extra_headers["value"] = extra_headers mock_client = MagicMock() mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object( - manager, "_create_mcp_client", side_effect=fake_create_mcp_client - ): + with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): with patch.object(manager, "_build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( @@ -822,9 +751,7 @@ class TestMcpRateLimitServerNameSurfacing: request_obj.tool_name = "list_repos" request_obj.arguments = {"org": "acme"} - result = self.proxy_logging._convert_mcp_to_llm_format( - request_obj, {"mcp_rate_limit_server_name": "github"} - ) + result = self.proxy_logging._convert_mcp_to_llm_format(request_obj, {"mcp_rate_limit_server_name": "github"}) assert result["mcp_server_name"] == "github" @@ -861,16 +788,10 @@ class TestMcpRateLimitServerNameSurfacing: return {"model": "fake"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( - return_value=MagicMock() - ) - proxy_logging._convert_mcp_to_llm_format = MagicMock( - side_effect=capture_convert - ) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging._convert_mcp_to_llm_format = MagicMock(side_effect=capture_convert) proxy_logging.pre_call_hook = AsyncMock(return_value=None) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( - return_value={"arguments": {}} - ) + proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {}}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -889,3 +810,226 @@ class TestMcpRateLimitServerNameSurfacing: ) assert captured["kwargs"]["mcp_rate_limit_server_name"] == "gh" + + +class TestOpenApiByokCallTool: + @pytest.mark.asyncio + async def test_call_tool_openapi_byok_injects_request_auth_contextvar(self): + """Playground/responses call call_tool directly; BYOK must reach OpenAPI handlers.""" + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_auth_header, + ) + + manager = MCPServerManager() + server = MCPServer( + server_id="byok-openapi", + name="firecrawl_byok_test", + server_name="firecrawl_byok_test", + url="https://api.firecrawl.dev", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + spec_path="https://example.com/openapi.json", + is_byok=True, + ) + user_auth = UserAPIKeyAuth(user_id="default_user_id", api_key="sk-dashboard") + captured_auth: dict[str, Optional[str]] = {} + + async def fake_openapi_handler(_server, _name, _arguments): + captured_auth["value"] = _request_auth_header.get() + return MagicMock() + + with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager._resolve_byok_mcp_auth_header", + new=AsyncMock(return_value="fc-test-key"), + ): + with patch.object( + manager, + "_call_openapi_tool_handler", + side_effect=fake_openapi_handler, + ): + await manager.call_tool( + server_name=server.server_name, + name="scrapeandextractfromurl", + arguments={"body": {"url": "https://example.com"}}, + user_api_key_auth=user_auth, + ) + + assert captured_auth["value"] == "ApiKey fc-test-key" + + +class TestFormatByokOpenapiAuthHeader: + def _server(self, auth_type): + return MCPServer( + server_id="s1", + name="s1", + server_name="s1", + url="https://example.com", + transport=MCPTransport.http, + auth_type=auth_type, + ) + + def test_api_key_auth_type(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _format_byok_openapi_auth_header, + ) + + assert _format_byok_openapi_auth_header(self._server(MCPAuth.api_key), "secret") == "ApiKey secret" + + def test_basic_auth_type(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _format_byok_openapi_auth_header, + ) + + assert _format_byok_openapi_auth_header(self._server(MCPAuth.basic), "secret") == "Basic secret" + + def test_defaults_to_bearer(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _format_byok_openapi_auth_header, + ) + + assert _format_byok_openapi_auth_header(self._server(MCPAuth.oauth2), "secret") == "Bearer secret" + + +class TestOpenapiForwardedExtraHeaders: + def _server(self, extra_headers): + return MCPServer( + server_id="s1", + name="s1", + server_name="s1", + url="https://example.com", + transport=MCPTransport.http, + extra_headers=extra_headers, + ) + + def test_returns_none_without_extra_headers_config(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _openapi_forwarded_extra_headers, + ) + + server = self._server(None) + assert _openapi_forwarded_extra_headers(server, {"X-Custom": "v"}, None) is None + + def test_returns_none_without_raw_headers(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _openapi_forwarded_extra_headers, + ) + + server = self._server(["X-Custom"]) + assert _openapi_forwarded_extra_headers(server, None, None) is None + + def test_forwards_configured_header_case_insensitively(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _openapi_forwarded_extra_headers, + ) + + server = self._server(["X-Custom"]) + result = _openapi_forwarded_extra_headers(server, {"x-custom": "v"}, None) + assert result == {"X-Custom": "v"} + + def test_returns_none_when_no_configured_header_is_present(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _openapi_forwarded_extra_headers, + ) + + server = self._server(["X-Missing"]) + assert _openapi_forwarded_extra_headers(server, {"x-custom": "v"}, None) is None + + def test_skips_authorization_when_caller_header_must_be_stripped(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _openapi_forwarded_extra_headers, + ) + + server = self._server(["Authorization"]) + server.auth_type = MCPAuth.oauth2_token_exchange + result = _openapi_forwarded_extra_headers(server, {"authorization": "Bearer caller-token"}, None) + assert result is None + + def test_skips_non_string_header_entries(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _openapi_forwarded_extra_headers, + ) + + server = self._server(["X-Custom"]) + server.extra_headers = [123, "X-Custom"] # simulate malformed legacy config data + result = _openapi_forwarded_extra_headers(server, {"x-custom": "v"}, None) + assert result == {"X-Custom": "v"} + + +class TestResolveByokMcpAuthHeader: + def _server(self, is_byok): + return MCPServer( + server_id="s1", + name="s1", + server_name="s1", + url="https://example.com", + transport=MCPTransport.http, + is_byok=is_byok, + ) + + @pytest.mark.asyncio + async def test_non_byok_server_passes_header_through_unchanged(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _resolve_byok_mcp_auth_header, + ) + + server = self._server(is_byok=False) + result = await _resolve_byok_mcp_auth_header(server, None, "caller-header") + assert result == "caller-header" + + @pytest.mark.asyncio + async def test_byok_server_uses_stored_credential_when_no_header_supplied(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _resolve_byok_mcp_auth_header, + ) + + server = self._server(is_byok=True) + user_auth = UserAPIKeyAuth(user_id="user-1", api_key="sk-dashboard") + + with patch( + "litellm.proxy._experimental.mcp_server.server._get_byok_credential", + new=AsyncMock(return_value="stored-cred"), + ): + result = await _resolve_byok_mcp_auth_header(server, user_auth, None) + + assert result == "stored-cred" + + @pytest.mark.asyncio + async def test_byok_server_raises_401_when_no_credential_stored(self): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _resolve_byok_mcp_auth_header, + ) + + server = self._server(is_byok=True) + user_auth = UserAPIKeyAuth(user_id="user-1", api_key="sk-dashboard") + + with patch( + "litellm.proxy._experimental.mcp_server.server._get_byok_credential", + new=AsyncMock(return_value=None), + ): + with pytest.raises(HTTPException) as exc_info: + await _resolve_byok_mcp_auth_header(server, user_auth, None) + + assert exc_info.value.status_code == 401 + assert exc_info.value.detail["error"] == "byok_auth_required" + + @pytest.mark.asyncio + async def test_byok_server_checks_credential_and_keeps_caller_header_when_supplied(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _resolve_byok_mcp_auth_header, + ) + + server = self._server(is_byok=True) + user_auth = UserAPIKeyAuth(user_id="user-1", api_key="sk-dashboard") + check_mock = AsyncMock(return_value=None) + + with patch( + "litellm.proxy._experimental.mcp_server.server._check_byok_credential", + new=check_mock, + ): + result = await _resolve_byok_mcp_auth_header(server, user_auth, "caller-header") + + check_mock.assert_awaited_once_with(server, user_auth) + assert result == "caller-header" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 6d3e09c6b75..358c0409db4 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 @@ -4074,6 +4074,24 @@ class TestMCPServerTimestamps: assert table.created_at is None assert table.updated_at is None + def test_build_mcp_server_table_preserves_tool_overrides(self): + """Tool display/description overrides must survive registry -> API table conversion.""" + manager = MCPServerManager() + server = MCPServer( + server_id="override-server", + name="deepwiki", + server_name="deepwiki_mcp", + url="https://example.com/mcp", + transport=MCPTransport.http, + tool_name_to_display_name={"read_wiki_structure": "browse_docs"}, + tool_name_to_description={"read_wiki_structure": "Browse repository documentation"}, + ) + + table = manager._build_mcp_server_table(server) + + assert table.tool_name_to_display_name == {"read_wiki_structure": "browse_docs"} + assert table.tool_name_to_description == {"read_wiki_structure": "Browse repository documentation"} + @pytest.mark.asyncio async def test_round_trip_timestamps_preserved(self): """Timestamps survive the full round-trip: LiteLLM_MCPServerTable -> MCPServer -> LiteLLM_MCPServerTable.""" 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 fb4eee54e1d..e114f46e866 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 @@ -1050,6 +1050,292 @@ class TestListToolsRestAPI: assert result["error"] == "unexpected_error" assert "access_denied" in result["message"] + async def test_mcp_server_name_query_param_resolves_to_server(self, monkeypatch): + """mcp_server_name is a name-based alias for server_id: it should + resolve to the matching server and scope the response to it.""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + stub_server = MCPServer( + server_id="uuid-abc-123", + name="my-server", + transport=MCPTransport.sse, + ) + stub_server.alias = "my-server" + stub_server.server_name = "my-server" + stub_server.available_on_public_internet = True + stub_server.allowed_tools = None + stub_server.mcp_info = {"server_name": "my-server"} + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["uuid-abc-123"] + + captured = {"called": False, "server_arg": None} + + async def fake_get_tools( + server, + server_auth_header, + raw_headers=None, + user_api_key_auth=None, + extra_headers=None, + apply_tool_filters=True, + ): + captured["called"] = True + captured["server_arg"] = server + return ["tool-x"] + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_name", + lambda name: stub_server if name == "my-server" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda sid: stub_server if sid == "uuid-abc-123" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + result = await rest_endpoints.list_tool_rest_api( + request, + server_id=None, + mcp_server_name="my-server", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert captured["called"] is True + assert captured["server_arg"] is stub_server + assert result["tools"] == ["tool-x"] + assert result["error"] is None + + async def test_mcp_server_name_filter_uses_real_catalog_with_tool_search(self, monkeypatch): + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + stub_server = MCPServer( + server_id="uuid-search-123", + name="search-server", + transport=MCPTransport.sse, + ) + stub_server.alias = "search-server" + stub_server.server_name = "search-server" + stub_server.available_on_public_internet = True + stub_server.allowed_tools = None + stub_server.mcp_info = {"server_name": "search-server"} + user_api_key_dict = UserAPIKeyAuth( + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="search-scope", + mcp_tool_search_enabled=True, + mcp_servers=["uuid-search-123"], + ) + ) + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["uuid-search-123"] + + async def fake_get_tools( + server, + server_auth_header, + raw_headers=None, + user_api_key_auth=None, + extra_headers=None, + apply_tool_filters=True, + ): + return ["scoped-tool"] + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_name", + lambda name: stub_server if name == "search-server" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "uuid-search-123" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + result = await rest_endpoints.list_tool_rest_api( + request, + server_id=None, + mcp_server_name="search-server", + user_api_key_dict=user_api_key_dict, + ) + + assert result["tools"] == ["scoped-tool"] + assert result["error"] is None + + async def test_toolset_name_query_param_scopes_to_toolset_servers(self, monkeypatch): + """toolset_name should resolve the toolset, apply its scope to the + caller's UserAPIKeyAuth via _apply_toolset_scope, and only list tools + from servers the scoped auth is allowed to see.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + scoped_auth = UserAPIKeyAuth( + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="toolset-scope", + mcp_tool_search_enabled=True, + mcp_servers=["toolset-server-1"], + ) + ) + + class StubToolset: + toolset_id = "toolset-1" + + class StubServer: + alias = "toolset-server-1" + server_name = "toolset-server-1" + name = "toolset-server-1" + allowed_tools = None + mcp_info = {"server_name": "toolset-server-1"} + available_on_public_internet = True + + stub_server = StubServer() + + async def fake_get_toolset_by_name_cached(prisma_client, toolset_name): + assert toolset_name == "research_tools" + return StubToolset() + + async def fake_apply_toolset_scope(user_api_key_auth, toolset_id): + assert toolset_id == "toolset-1" + return scoped_auth + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(**kwargs): + assert kwargs["user_api_key_auth"] is scoped_auth + return ["toolset-server-1"] + + async def fake_get_tools(server, server_auth_header, *args, **kwargs): + return ["toolset-tool-1"] + + monkeypatch.setattr( + "litellm.proxy.utils.get_prisma_client_or_throw", + lambda *args, **kwargs: MagicMock(), + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_toolset_by_name_cached", + fake_get_toolset_by_name_cached, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_apply_toolset_scope", + fake_apply_toolset_scope, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "toolset-server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + result = await rest_endpoints.list_tool_rest_api( + request, + server_id=None, + toolset_name="research_tools", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert result["tools"] == ["toolset-tool-1"] + assert result["error"] is None + + async def test_toolset_name_not_found_returns_error(self, monkeypatch): + async def fake_get_toolset_by_name_cached(prisma_client, toolset_name): + return None + + monkeypatch.setattr( + "litellm.proxy.utils.get_prisma_client_or_throw", + lambda *args, **kwargs: MagicMock(), + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_toolset_by_name_cached", + fake_get_toolset_by_name_cached, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.list_tool_rest_api( + request, + server_id=None, + toolset_name="does-not-exist", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert exc_info.value.status_code == 404 + assert "does-not-exist" in str(exc_info.value.detail) + async def test_oauth2_user_token_injected_for_single_server(self, monkeypatch): """For a single-server OAuth2 request, _get_user_oauth_extra_headers is called and the returned headers are forwarded to _get_tools_for_single_server.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_utils.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_utils.py new file mode 100644 index 00000000000..73fdee9cde3 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_utils.py @@ -0,0 +1,49 @@ +import pytest +from fastapi import HTTPException + +from litellm.proxy._experimental.mcp_server.utils import ( + validate_and_normalize_mcp_server_payload, + validate_tool_display_names, +) +from litellm.proxy._types import NewMCPServerRequest + + +class TestValidateToolDisplayNames: + def test_allows_none_and_empty(self): + validate_tool_display_names(None) + validate_tool_display_names({}) + + @pytest.mark.parametrize( + "display_name", + ["browse_repo_docs", "browse-repo-docs", "BrowseRepoDocs123"], + ) + def test_allows_bedrock_safe_names(self, display_name): + validate_tool_display_names({"read_wiki_structure": display_name}) + + @pytest.mark.parametrize( + "display_name", + ["Browse Repo Docs", "browse.repo.docs", "browse/repo", "browse@docs"], + ) + def test_rejects_names_bedrock_would_reject(self, display_name): + with pytest.raises(HTTPException) as exc_info: + validate_tool_display_names({"read_wiki_structure": display_name}) + assert exc_info.value.status_code == 400 + assert display_name in str(exc_info.value.detail) + + +class TestValidateAndNormalizeMcpServerPayload: + def test_rejects_invalid_tool_display_name_on_create(self): + payload = NewMCPServerRequest( + server_name="deepwiki_mcp", + tool_name_to_display_name={"read_wiki_structure": "Browse Repo Docs"}, + ) + with pytest.raises(HTTPException) as exc_info: + validate_and_normalize_mcp_server_payload(payload) + assert exc_info.value.status_code == 400 + + def test_accepts_valid_tool_display_name_on_create(self): + payload = NewMCPServerRequest( + server_name="deepwiki_mcp", + tool_name_to_display_name={"read_wiki_structure": "browse_repo_docs"}, + ) + validate_and_normalize_mcp_server_payload(payload) diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index 35bfa6ea9e4..80e2eb48f62 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -28,6 +28,7 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: call_tool=AsyncMock(return_value=_DummyMCPResult()), # Newer logging path calls this to enrich spend logs metadata _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -279,6 +280,87 @@ async def test_execute_tool_calls_keeps_tool_name_when_equal_to_server(monkeypat assert call_tool_mock.await_args.kwargs["name"] == tool_name +@pytest.mark.asyncio +async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_name( + monkeypatch, +): + call_tool_mock = _setup_mcp_call_environment(monkeypatch) + fake_server = types.SimpleNamespace( + alias="my_deepwiki", + server_name="deepwiki_test", + server_id="test-server-id", + short_prefix=None, + mcp_info=None, + tool_name_to_display_name=None, + ) + from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm + + _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock( + return_value=fake_server + ) + + tool_name = "my_deepwiki-read_wiki_structure" + tool_calls = [ + { + "id": "call-4", + "function": {"name": tool_name, "arguments": "{}"}, + } + ] + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={tool_name: "deepwiki_test"}, + tool_calls=tool_calls, + user_api_key_auth=None, + ) + + assert call_tool_mock.await_count == 1 + assert call_tool_mock.await_args is not None + assert call_tool_mock.await_args.kwargs["name"] == "read_wiki_structure" + + +@pytest.mark.asyncio +async def test_execute_tool_calls_reverse_maps_display_name(monkeypatch): + call_tool_mock = _setup_mcp_call_environment(monkeypatch) + colliding_server = types.SimpleNamespace( + alias=None, + server_name="other_mcp", + server_id="other-server-id", + short_prefix=None, + mcp_info=None, + tool_name_to_display_name={"search": "search_docs"}, + ) + fake_server = types.SimpleNamespace( + alias=None, + server_name="deepwiki_mcp", + server_id="test-server-id", + short_prefix=None, + mcp_info=None, + tool_name_to_display_name={"read_wiki_structure": "browse_repo_docs"}, + ) + from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm + + _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=colliding_server) + _msm.global_mcp_server_manager.get_mcp_server_by_name = MagicMock(return_value=fake_server) + + tool_name = "browse_repo_docs" + tool_calls = [ + { + "id": "call-5", + "function": {"name": tool_name, "arguments": "{}"}, + } + ] + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={tool_name: "deepwiki_mcp"}, + tool_calls=tool_calls, + user_api_key_auth=None, + ) + + assert call_tool_mock.await_count == 1 + assert call_tool_mock.await_args is not None + assert call_tool_mock.await_args.kwargs["name"] == "read_wiki_structure" + + @pytest.mark.asyncio async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkeypatch): """ diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py index ecf89ce18d4..cdace5f6327 100644 --- a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py @@ -66,6 +66,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: fake_manager = types.SimpleNamespace( call_tool=call_tool, _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index b2bf16abe39..9cac103d1a7 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -25,7 +25,7 @@ import OpenAPIFormSection, { OpenAPIKeyTool } from "./OpenAPIFormSection"; import MCPLogoSelector from "./MCPLogoSelector"; import EnvVarsSection from "./EnvVarsSection"; import { isAdminRole } from "@/utils/roles"; -import { validateMCPServerUrl, validateMCPServerName, normalizeEnvVars } from "./utils"; +import { validateMCPServerUrl, validateMCPServerName, normalizeEnvVars, TOOL_DISPLAY_NAME_PATTERN } from "./utils"; import NotificationsManager from "../molecules/notifications_manager"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; import { useTestMCPConnection } from "@/hooks/useTestMCPConnection"; @@ -326,6 +326,15 @@ const CreateMCPServer: React.FC = ({ }, [isModalVisible, prefillData, form]); const handleCreate = async (values: Record) => { + const invalidDisplayName = Object.entries(toolNameToDisplayName).find( + ([, displayName]) => displayName && !TOOL_DISPLAY_NAME_PATTERN.test(displayName), + ); + if (invalidDisplayName) { + NotificationsManager.fromBackend( + `Tool display name "${invalidDisplayName[1]}" is invalid. Only letters, digits, underscores, and hyphens are allowed (no spaces).`, + ); + return; + } setIsLoading(true); try { const { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx index 44b2ba25f11..7671f29b5e6 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -65,12 +65,20 @@ vi.mock("./mcp_tool_configuration", () => ({ + ), })); @@ -382,7 +390,7 @@ describe("MCPServerEdit (tool allowlist)", () => { it("saves tool overrides for legacy unrestricted servers", async () => { vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...interactiveOAuthServer, - tool_name_to_display_name: { read_user: "Read User" }, + tool_name_to_display_name: { read_user: "ReadUser" }, tool_name_to_description: { read_user: "Reads users" }, }); @@ -416,9 +424,36 @@ describe("MCPServerEdit (tool allowlist)", () => { const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; expect(payload.mcp_info.tool_allowlist_enforced).toBe(false); expect(payload.allowed_tools).toBeUndefined(); - expect(payload.tool_name_to_display_name).toEqual({ read_user: "Read User" }); + expect(payload.tool_name_to_display_name).toEqual({ read_user: "ReadUser" }); expect(payload.tool_name_to_description).toEqual({ read_user: "Reads users" }); }); + + it("blocks save and does not call the API when a tool display name contains a space", async () => { + render( + , + ); + + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Set invalid tool override" })); + }); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + }); }); describe("MCPServerEdit (interactive OAuth)", () => { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index bc9c3cfea07..413c7f1e11b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -21,7 +21,13 @@ import StdioConfiguration from "./StdioConfiguration"; import MCPLogoSelector from "./MCPLogoSelector"; import EnvVarsSection from "./EnvVarsSection"; import TokenEndpointAuthMethodField from "./TokenEndpointAuthMethodField"; -import { validateMCPServerUrl, validateMCPServerName, normalizeEnvVars } from "./utils"; +import { + validateMCPServerUrl, + validateMCPServerName, + normalizeEnvVars, + normalizeToolOverrideMap, + TOOL_DISPLAY_NAME_PATTERN, +} from "./utils"; import NotificationsManager from "../molecules/notifications_manager"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; @@ -257,8 +263,8 @@ const MCPServerEdit: React.FC = ({ if (hasExistingToolAllowlist) { setAllowedTools(mcpServer.allowed_tools ?? []); } - setToolNameToDisplayName(mcpServer.tool_name_to_display_name ?? {}); - setToolNameToDescription(mcpServer.tool_name_to_description ?? {}); + setToolNameToDisplayName(normalizeToolOverrideMap(mcpServer.tool_name_to_display_name)); + setToolNameToDescription(normalizeToolOverrideMap(mcpServer.tool_name_to_description)); }, [mcpServer, hasExistingToolAllowlist]); useEffect(() => { @@ -449,6 +455,15 @@ const MCPServerEdit: React.FC = ({ const handleSave = async (values: Record) => { if (!accessToken) return; + const invalidDisplayName = Object.entries(toolNameToDisplayName).find( + ([, displayName]) => displayName && !TOOL_DISPLAY_NAME_PATTERN.test(displayName), + ); + if (invalidDisplayName) { + NotificationsManager.fromBackend( + `Tool display name "${invalidDisplayName[1]}" is invalid. Only letters, digits, underscores, and hyphens are allowed (no spaces).`, + ); + return; + } try { // Ensure access groups is always a string array const { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx index 064e3dea614..cd8282d29ed 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx @@ -109,4 +109,78 @@ describe("MCPToolConfiguration", () => { expect(screen.getAllByText("Disabled")).toHaveLength(2); }); }); + + it("shows a validation error for a display name containing a space", async () => { + const Wrapper = () => { + const [toolNameToDisplayName, setToolNameToDisplayName] = useState>({}); + + return ( + + ); + }; + + render(); + + fireEvent.click(screen.getByText("Flat List")); + fireEvent.click(screen.getAllByTitle("Edit display name and description")[0]); + + const input = screen.getByPlaceholderText("read_user"); + fireEvent.change(input, { target: { value: "Browse Repo Docs" } }); + + await waitFor(() => { + expect( + screen.getByText("Only letters, digits, underscores, and hyphens are allowed (no spaces)."), + ).toBeInTheDocument(); + }); + }); + + it("accepts a Bedrock-safe display name without showing a validation error", async () => { + const Wrapper = () => { + const [toolNameToDisplayName, setToolNameToDisplayName] = useState>({}); + + return ( + + ); + }; + + render(); + + fireEvent.click(screen.getByText("Flat List")); + fireEvent.click(screen.getAllByTitle("Edit display name and description")[0]); + + const input = screen.getByPlaceholderText("read_user"); + fireEvent.change(input, { target: { value: "browse_repo_docs" } }); + + await waitFor(() => { + expect( + screen.queryByText("Only letters, digits, underscores, and hyphens are allowed (no spaces)."), + ).not.toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx index a1061c10515..1ebc07eac86 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx @@ -3,6 +3,7 @@ import { Card, Title, Text } from "@tremor/react"; import { ToolOutlined, CheckCircleOutlined, SearchOutlined, EditOutlined } from "@ant-design/icons"; import { Badge, Spin, Checkbox, Input, Radio } from "antd"; import McpCrudPermissionPanel from "./McpCrudPermissionPanel"; +import { TOOL_DISPLAY_NAME_PATTERN } from "./utils"; interface KeyTool { name: string; @@ -60,86 +61,98 @@ const ToolRow: React.FC = ({ onToggleExpand, onDisplayNameChange, onDescriptionChange, -}) => ( -
-
onToggle(tool.name)}> -
- onToggle(tool.name)} /> -
-
- {toolNameToDisplayName[tool.name] || tool.name} - - {isEnabled ? "Enabled" : "Disabled"} - - {toolNameToDisplayName[tool.name] && ( - - Custom name +}) => { + const displayNameValue = toolNameToDisplayName[tool.name] || ""; + const isDisplayNameInvalid = displayNameValue !== "" && !TOOL_DISPLAY_NAME_PATTERN.test(displayNameValue); + + return ( +
+
onToggle(tool.name)}> +
+ onToggle(tool.name)} /> +
+
+ {toolNameToDisplayName[tool.name] || tool.name} + + {isEnabled ? "Enabled" : "Disabled"} + {toolNameToDisplayName[tool.name] && ( + + Custom name + + )} +
+ {(toolNameToDescription[tool.name] || tool.description) && ( + + {toolNameToDescription[tool.name] || tool.description} + + )} + + {isEnabled ? "✓ Users can call this tool" : "✗ Users cannot call this tool"} + +
+ +
+
+ {isEditExpanded && ( +
e.stopPropagation()} + > +
+ Display Name + onDisplayNameChange(tool.name, e.target.value)} + status={isDisplayNameInvalid ? "error" : undefined} + /> + {isDisplayNameInvalid ? ( + + Only letters, digits, underscores, and hyphens are allowed (no spaces). + + ) : ( + + Override how this tool's name appears to users. Leave blank to use original. + )}
- {(toolNameToDescription[tool.name] || tool.description) && ( - - {toolNameToDescription[tool.name] || tool.description} +
+ Description + onDescriptionChange(tool.name, e.target.value)} + rows={2} + /> + + Override the tool description shown to users. Leave blank to use original. - )} - - {isEnabled ? "✓ Users can call this tool" : "✗ Users cannot call this tool"} - +
- -
+ )}
- {isEditExpanded && ( -
e.stopPropagation()} - > -
- Display Name - onDisplayNameChange(tool.name, e.target.value)} - /> - - Override how this tool's name appears to users. Leave blank to use original. - -
-
- Description - onDescriptionChange(tool.name, e.target.value)} - rows={2} - /> - - Override the tool description shown to users. Leave blank to use original. - -
-
- )} -
-); + ); +}; const MCPToolConfiguration: React.FC = ({ accessToken, diff --git a/ui/litellm-dashboard/src/components/mcp_tools/utils.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/utils.test.tsx index c4c8c52888c..3b4fda400c2 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/utils.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/utils.test.tsx @@ -1,5 +1,12 @@ import { describe, it, expect } from "vitest"; -import { extractMCPToken, maskUrl, getMaskedAndFullUrl, validateMCPServerUrl, validateMCPServerName } from "./utils"; +import { + extractMCPToken, + maskUrl, + getMaskedAndFullUrl, + validateMCPServerUrl, + validateMCPServerName, + normalizeToolOverrideMap, +} from "./utils"; describe("extractMCPToken", () => { it("should extract token after /mcp/", () => { @@ -67,3 +74,21 @@ describe("validateMCPServerName", () => { await expect(validateMCPServerName("my server")).rejects.toBeDefined(); }); }); + +describe("normalizeToolOverrideMap", () => { + it("returns empty object for nullish input", () => { + expect(normalizeToolOverrideMap(null)).toEqual({}); + expect(normalizeToolOverrideMap(undefined)).toEqual({}); + }); + + it("parses JSON string maps from legacy API responses", () => { + expect(normalizeToolOverrideMap('{"read_wiki_structure":"browse_docs"}')).toEqual({ + read_wiki_structure: "browse_docs", + }); + }); + + it("passes through object maps unchanged", () => { + const map = { read_user: "Read User" }; + expect(normalizeToolOverrideMap(map)).toBe(map); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/utils.tsx b/ui/litellm-dashboard/src/components/mcp_tools/utils.tsx index 6d9479a13c3..7d6e24fc480 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/utils.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/utils.tsx @@ -54,6 +54,14 @@ export const validateMCPServerName = (value: string) => { : Promise.resolve(); }; +export const TOOL_DISPLAY_NAME_PATTERN = /^[a-zA-Z0-9_-]+$/; + +export const validateToolDisplayName = (value: string) => { + return value && !TOOL_DISPLAY_NAME_PATTERN.test(value) + ? Promise.reject("Only letters, digits, underscores, and hyphens are allowed (no spaces).") + : Promise.resolve(); +}; + // Normalize the env_vars form list into the payload shape the backend expects. // Drops empty rows, invalid identifiers, and duplicate names; user-scoped entries never carry a value. export const normalizeEnvVars = (list: unknown): MCPEnvVar[] => { @@ -77,3 +85,22 @@ export const normalizeEnvVars = (list: unknown): MCPEnvVar[] => { } return out; }; + +/** Normalize tool override maps from API/DB (dict or JSON string) for form state. */ +export const normalizeToolOverrideMap = ( + value: Record | string | null | undefined, +): Record => { + if (!value) return {}; + if (typeof value === "string") { + try { + const parsed = JSON.parse(value) as unknown; + if (parsed && typeof parsed === "object" && !Array.isArray(parsed)) { + return parsed as Record; + } + } catch { + return {}; + } + return {}; + } + return value; +};