Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
101 changes: 99 additions & 2 deletions litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]]]:
Expand Down Expand Up @@ -3861,6 +3931,15 @@ async def call_tool(
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
Expand Down Expand Up @@ -3907,9 +3986,25 @@ async def call_tool(
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:
Expand Down Expand Up @@ -4553,6 +4648,8 @@ def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable:
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,
Expand Down
51 changes: 47 additions & 4 deletions litellm/proxy/_experimental/mcp_server/rest_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,7 @@ def _connection_error_message(exc: BaseException) -> str:
ListMCPToolsRestAPIResponseObject,
MCPInfo,
MCPServer,
_apply_toolset_scope,
_fire_mcp_success_logging,
_tool_name_matches,
execute_mcp_tool,
Expand Down Expand Up @@ -541,10 +542,37 @@ async def _list_tools_for_single_server(
"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",
)
Comment thread
Sameerlite marked this conversation as resolved.
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=(
Expand Down Expand Up @@ -582,16 +610,29 @@ async def list_tool_rest_api(
)

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

Comment thread
Sameerlite marked this conversation as resolved.
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,
Expand Down Expand Up @@ -719,6 +760,8 @@ async def list_tool_rest_api(
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))
Expand Down
48 changes: 45 additions & 3 deletions litellm/proxy/_experimental/mcp_server/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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.
Expand Down
21 changes: 14 additions & 7 deletions litellm/responses/mcp/litellm_proxy_mcp_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,10 @@
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
Expand Down Expand Up @@ -628,6 +631,9 @@ async def _execute_tool_calls(
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 = []
Expand All @@ -654,11 +660,13 @@ async def _execute_tool_calls(

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
)
Comment thread
Sameerlite marked this conversation as resolved.
sanitized_tool_name = strip_known_server_prefix(resolved_tool_name, mcp_server)

start_time = datetime.now()
logging_input = [
Expand Down Expand Up @@ -741,7 +749,6 @@ async def _execute_tool_calls(
"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"] = (
Expand Down
Loading
Loading