mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(mcp): alias/display-name tool routing, REST filters, BYOK auth (#32320)
* fix(mcp): resolve tool name prefix via known server prefixes, not string match When an MCP server's alias differs from its server_name, tool names are listed with the alias prefix but _execute_tool_calls compared that prefix against the server_name stored in tool_server_map. The mismatch silently skipped prefix stripping, forwarding the fully-prefixed tool name upstream and causing "Unknown tool" failures. Resolve the actual MCPServer object and strip using its known prefix forms (alias, server_name, server_id) instead. * fix(mcp): preserve tool overrides and scope REST tool listing Return saved tool display/description overrides from the server table API so the edit UI reloads them, resolve display names before prefix stripping on tool calls, and honor mcp_server_name and toolset_name filters on the REST tools list endpoint. Co-authored-by: Cursor <cursoragent@cursor.com> * fix(mcp): inject BYOK credentials on Playground OpenAPI tool calls Playground and Responses API route MCP execution through call_tool, which skipped BYOK lookup and never set the OpenAPI auth ContextVar, so upstream calls went out unauthenticated despite a stored user credential. Co-authored-by: Cursor <cursoragent@cursor.com> * test(mcp): cover alias-mismatch prefix stripping and display-name reverse mapping Regression tests for _execute_tool_calls: an MCP server whose alias differs from its server_name must still have its tool-name prefix stripped correctly, and a tool called by its configured display name must resolve back to the original tool name before dispatch. * fix(mcp): validate tool display names against Bedrock's tool-name pattern A display name replaces the tool name sent to the LLM provider, so a value with spaces or other special characters saves successfully but fails every subsequent Bedrock tool call. Validate tool_name_to_display_name server-side (create/update payload) against Bedrock's [a-zA-Z0-9_-]+ constraint, and add matching inline validation plus a save-blocking guard in the Admin UI's create and edit MCP server forms. * style(mcp): fix ruff/prettier formatting on CI No logic changes; satisfies the format checks flagged on PR #32320. * fix(mcp): fix CI failures on PR - complexity budget and stale test mock Extract toolset-scope resolution and query-param normalization out of list_tool_rest_api into helpers to bring it back under the C901 complexity budget (was 18, now within the 15 threshold). Add the missing get_mcp_server_by_name stub to the streaming iterator test's mock manager; the alias-fallback resolution added for tool-name-prefix stripping calls it unconditionally when _get_mcp_server_from_tool_name misses. * test(mcp): cover BYOK OpenAPI auth-header helpers to close codecov patch gap _format_byok_openapi_auth_header, _openapi_forwarded_extra_headers, and _resolve_byok_mcp_auth_header were only exercised indirectly via a mocked call_tool test, leaving their branches (auth-type formatting, header forwarding/stripping, missing-credential 401) uncovered. * fix(mcp): resolve BYOK auth before queuing the during-hook task _resolve_byok_mcp_auth_header can raise a 401 when no credential is stored. Resolving it after during_hook_task was already queued meant a hook's side effects (audit logging, rate-limit bookkeeping) could run and record success for a tool call that then fails on the missing credential. * fix: correct mcp alias routing regressions --------- Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
42f5b0bd34
commit
a78dc69a09
17 changed files with 1183 additions and 216 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"] = (
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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<CreateMCPServerProps> = ({
|
|||
}, [isModalVisible, prefillData, form]);
|
||||
|
||||
const handleCreate = async (values: Record<string, any>) => {
|
||||
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 {
|
||||
|
|
|
|||
|
|
@ -65,12 +65,20 @@ vi.mock("./mcp_tool_configuration", () => ({
|
|||
<button
|
||||
type="button"
|
||||
onClick={() => {
|
||||
onToolNameToDisplayNameChange({ read_user: "Read User" });
|
||||
onToolNameToDisplayNameChange({ read_user: "ReadUser" });
|
||||
onToolNameToDescriptionChange({ read_user: "Reads users" });
|
||||
}}
|
||||
>
|
||||
Set tool overrides
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => {
|
||||
onToolNameToDisplayNameChange({ read_user: "Read User" });
|
||||
}}
|
||||
>
|
||||
Set invalid tool override
|
||||
</button>
|
||||
</div>
|
||||
),
|
||||
}));
|
||||
|
|
@ -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(
|
||||
<MCPServerEdit
|
||||
mcpServer={{
|
||||
...interactiveOAuthServer,
|
||||
allowed_tools: [],
|
||||
mcp_info: { server_name: "OAuthServer" },
|
||||
}}
|
||||
accessToken="access-token"
|
||||
onCancel={vi.fn()}
|
||||
onSuccess={vi.fn()}
|
||||
availableAccessGroups={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
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)", () => {
|
||||
|
|
|
|||
|
|
@ -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<MCPServerEditProps> = ({
|
|||
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<MCPServerEditProps> = ({
|
|||
|
||||
const handleSave = async (values: Record<string, any>) => {
|
||||
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 {
|
||||
|
|
|
|||
|
|
@ -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<Record<string, string>>({});
|
||||
|
||||
return (
|
||||
<MCPToolConfiguration
|
||||
accessToken="token"
|
||||
formValues={{ url: "https://example.com/mcp", transport: "http", auth_type: "none" }}
|
||||
allowedTools={[]}
|
||||
existingAllowedTools={null}
|
||||
onAllowedToolsChange={vi.fn()}
|
||||
toolNameToDisplayName={toolNameToDisplayName}
|
||||
toolNameToDescription={{}}
|
||||
onToolNameToDisplayNameChange={setToolNameToDisplayName}
|
||||
onToolNameToDescriptionChange={vi.fn()}
|
||||
externalTools={tools}
|
||||
externalCanFetch
|
||||
isEditMode
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
||||
render(<Wrapper />);
|
||||
|
||||
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<Record<string, string>>({});
|
||||
|
||||
return (
|
||||
<MCPToolConfiguration
|
||||
accessToken="token"
|
||||
formValues={{ url: "https://example.com/mcp", transport: "http", auth_type: "none" }}
|
||||
allowedTools={[]}
|
||||
existingAllowedTools={null}
|
||||
onAllowedToolsChange={vi.fn()}
|
||||
toolNameToDisplayName={toolNameToDisplayName}
|
||||
toolNameToDescription={{}}
|
||||
onToolNameToDisplayNameChange={setToolNameToDisplayName}
|
||||
onToolNameToDescriptionChange={vi.fn()}
|
||||
externalTools={tools}
|
||||
externalCanFetch
|
||||
isEditMode
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
||||
render(<Wrapper />);
|
||||
|
||||
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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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<ToolRowProps> = ({
|
|||
onToggleExpand,
|
||||
onDisplayNameChange,
|
||||
onDescriptionChange,
|
||||
}) => (
|
||||
<div
|
||||
className={`rounded-lg border transition-colors ${
|
||||
isEnabled
|
||||
? "bg-blue-50 border-blue-300 hover:border-blue-400"
|
||||
: "bg-gray-50 border-gray-200 hover:border-gray-300"
|
||||
}`}
|
||||
>
|
||||
<div className="p-4 cursor-pointer" onClick={() => onToggle(tool.name)}>
|
||||
<div className="flex items-start gap-3">
|
||||
<Checkbox checked={isEnabled} onChange={() => onToggle(tool.name)} />
|
||||
<div className="flex-1">
|
||||
<div className="flex items-center gap-2">
|
||||
<Text className="font-medium text-gray-900">{toolNameToDisplayName[tool.name] || tool.name}</Text>
|
||||
<span
|
||||
className={`px-2 py-0.5 text-xs rounded-full font-medium ${
|
||||
isEnabled ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800"
|
||||
}`}
|
||||
>
|
||||
{isEnabled ? "Enabled" : "Disabled"}
|
||||
</span>
|
||||
{toolNameToDisplayName[tool.name] && (
|
||||
<span className="px-2 py-0.5 text-xs rounded-full font-medium bg-purple-100 text-purple-800">
|
||||
Custom name
|
||||
}) => {
|
||||
const displayNameValue = toolNameToDisplayName[tool.name] || "";
|
||||
const isDisplayNameInvalid = displayNameValue !== "" && !TOOL_DISPLAY_NAME_PATTERN.test(displayNameValue);
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`rounded-lg border transition-colors ${
|
||||
isEnabled
|
||||
? "bg-blue-50 border-blue-300 hover:border-blue-400"
|
||||
: "bg-gray-50 border-gray-200 hover:border-gray-300"
|
||||
}`}
|
||||
>
|
||||
<div className="p-4 cursor-pointer" onClick={() => onToggle(tool.name)}>
|
||||
<div className="flex items-start gap-3">
|
||||
<Checkbox checked={isEnabled} onChange={() => onToggle(tool.name)} />
|
||||
<div className="flex-1">
|
||||
<div className="flex items-center gap-2">
|
||||
<Text className="font-medium text-gray-900">{toolNameToDisplayName[tool.name] || tool.name}</Text>
|
||||
<span
|
||||
className={`px-2 py-0.5 text-xs rounded-full font-medium ${
|
||||
isEnabled ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800"
|
||||
}`}
|
||||
>
|
||||
{isEnabled ? "Enabled" : "Disabled"}
|
||||
</span>
|
||||
{toolNameToDisplayName[tool.name] && (
|
||||
<span className="px-2 py-0.5 text-xs rounded-full font-medium bg-purple-100 text-purple-800">
|
||||
Custom name
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
{(toolNameToDescription[tool.name] || tool.description) && (
|
||||
<Text className="text-gray-500 text-sm block mt-1">
|
||||
{toolNameToDescription[tool.name] || tool.description}
|
||||
</Text>
|
||||
)}
|
||||
<Text className="text-gray-400 text-xs block mt-1">
|
||||
{isEnabled ? "✓ Users can call this tool" : "✗ Users cannot call this tool"}
|
||||
</Text>
|
||||
</div>
|
||||
<button
|
||||
type="button"
|
||||
onClick={(e) => onToggleExpand(tool.name, e)}
|
||||
className={`p-1.5 rounded-md transition-colors ${
|
||||
isEditExpanded ? "bg-blue-100 text-blue-600" : "text-gray-400 hover:text-gray-600 hover:bg-gray-100"
|
||||
}`}
|
||||
title="Edit display name and description"
|
||||
>
|
||||
<EditOutlined />
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
{isEditExpanded && (
|
||||
<div
|
||||
className="px-4 pb-4 pt-3 border-t border-gray-200 space-y-3 bg-gray-50 rounded-b-lg"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
>
|
||||
<div>
|
||||
<Text className="text-xs font-medium text-gray-600 mb-1 block">Display Name</Text>
|
||||
<Input
|
||||
placeholder={tool.name}
|
||||
value={toolNameToDisplayName[tool.name] || ""}
|
||||
onChange={(e) => onDisplayNameChange(tool.name, e.target.value)}
|
||||
status={isDisplayNameInvalid ? "error" : undefined}
|
||||
/>
|
||||
{isDisplayNameInvalid ? (
|
||||
<Text className="text-xs text-red-500 mt-1 block">
|
||||
Only letters, digits, underscores, and hyphens are allowed (no spaces).
|
||||
</Text>
|
||||
) : (
|
||||
<Text className="text-xs text-gray-400 mt-1 block">
|
||||
Override how this tool's name appears to users. Leave blank to use original.
|
||||
</Text>
|
||||
)}
|
||||
</div>
|
||||
{(toolNameToDescription[tool.name] || tool.description) && (
|
||||
<Text className="text-gray-500 text-sm block mt-1">
|
||||
{toolNameToDescription[tool.name] || tool.description}
|
||||
<div>
|
||||
<Text className="text-xs font-medium text-gray-600 mb-1 block">Description</Text>
|
||||
<Input.TextArea
|
||||
placeholder={tool.description || "No description"}
|
||||
value={toolNameToDescription[tool.name] || ""}
|
||||
onChange={(e) => onDescriptionChange(tool.name, e.target.value)}
|
||||
rows={2}
|
||||
/>
|
||||
<Text className="text-xs text-gray-400 mt-1 block">
|
||||
Override the tool description shown to users. Leave blank to use original.
|
||||
</Text>
|
||||
)}
|
||||
<Text className="text-gray-400 text-xs block mt-1">
|
||||
{isEnabled ? "✓ Users can call this tool" : "✗ Users cannot call this tool"}
|
||||
</Text>
|
||||
</div>
|
||||
</div>
|
||||
<button
|
||||
type="button"
|
||||
onClick={(e) => onToggleExpand(tool.name, e)}
|
||||
className={`p-1.5 rounded-md transition-colors ${
|
||||
isEditExpanded ? "bg-blue-100 text-blue-600" : "text-gray-400 hover:text-gray-600 hover:bg-gray-100"
|
||||
}`}
|
||||
title="Edit display name and description"
|
||||
>
|
||||
<EditOutlined />
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
{isEditExpanded && (
|
||||
<div
|
||||
className="px-4 pb-4 pt-3 border-t border-gray-200 space-y-3 bg-gray-50 rounded-b-lg"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
>
|
||||
<div>
|
||||
<Text className="text-xs font-medium text-gray-600 mb-1 block">Display Name</Text>
|
||||
<Input
|
||||
placeholder={tool.name}
|
||||
value={toolNameToDisplayName[tool.name] || ""}
|
||||
onChange={(e) => onDisplayNameChange(tool.name, e.target.value)}
|
||||
/>
|
||||
<Text className="text-xs text-gray-400 mt-1 block">
|
||||
Override how this tool's name appears to users. Leave blank to use original.
|
||||
</Text>
|
||||
</div>
|
||||
<div>
|
||||
<Text className="text-xs font-medium text-gray-600 mb-1 block">Description</Text>
|
||||
<Input.TextArea
|
||||
placeholder={tool.description || "No description"}
|
||||
value={toolNameToDescription[tool.name] || ""}
|
||||
onChange={(e) => onDescriptionChange(tool.name, e.target.value)}
|
||||
rows={2}
|
||||
/>
|
||||
<Text className="text-xs text-gray-400 mt-1 block">
|
||||
Override the tool description shown to users. Leave blank to use original.
|
||||
</Text>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
);
|
||||
};
|
||||
|
||||
const MCPToolConfiguration: React.FC<MCPToolConfigurationProps> = ({
|
||||
accessToken,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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, string> | string | null | undefined,
|
||||
): Record<string, string> => {
|
||||
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<string, string>;
|
||||
}
|
||||
} catch {
|
||||
return {};
|
||||
}
|
||||
return {};
|
||||
}
|
||||
return value;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue