[Feat] mcp resources support (#16800)

* feat: mcp prompts support

* feat: mcp resources support
This commit is contained in:
YutaSaito
2025-11-20 14:53:44 -08:00
committed by GitHub
parent 0d812f98bc
commit 93affcb732
11 changed files with 4222 additions and 1529 deletions
+224 -2
View File
@@ -8,14 +8,21 @@ from datetime import timedelta
from typing import Awaitable, Callable, Dict, List, Optional, TypeVar, Union
import httpx
from mcp import ClientSession, StdioServerParameters
from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.client.streamable_http import streamablehttp_client
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import (
CallToolRequestParams as MCPCallToolRequestParams,
GetPromptRequestParams,
GetPromptResult,
Prompt,
ResourceTemplate,
)
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import TextContent
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
@@ -289,3 +296,218 @@ class MCPClient:
], # Empty content for error case
isError=True,
)
async def list_prompts(self) -> List[Prompt]:
"""List available prompts from the server."""
verbose_logger.debug(
f"MCP client listing tools from {self.server_url or 'stdio'}"
)
async def _list_prompts_operation(session: ClientSession):
return await session.list_prompts()
try:
result = await self.run_with_session(_list_prompts_operation)
prompt_count = len(result.prompts)
prompt_names = [prompt.name for prompt in result.prompts]
verbose_logger.info(
f"MCP client listed {prompt_count} tools from {self.server_url or 'stdio'}: {prompt_names}"
)
return result.prompts
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_prompts was cancelled")
raise
except Exception as e:
error_type = type(e).__name__
verbose_logger.error(
f"MCP client list_prompts failed - "
f"Error Type: {error_type}, "
f"Error: {str(e)}, "
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during list_tools - "
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
async def get_prompt(
self, get_prompt_request_params: GetPromptRequestParams
) -> GetPromptResult:
"""Fetch a prompt definition from the MCP server."""
verbose_logger.info(
f"MCP client fetching prompt '{get_prompt_request_params.name}' with arguments: {get_prompt_request_params.arguments}"
)
async def _get_prompt_operation(session: ClientSession):
verbose_logger.debug("MCP client sending get_prompt request to session")
return await session.get_prompt(
name=get_prompt_request_params.name,
arguments=get_prompt_request_params.arguments,
)
try:
get_prompt_result = await self.run_with_session(_get_prompt_operation)
verbose_logger.info(
f"MCP client get_prompt '{get_prompt_request_params.name}' completed successfully"
)
return get_prompt_result
except asyncio.CancelledError:
verbose_logger.warning("MCP client get_prompt was cancelled")
raise
except Exception as e:
import traceback
error_trace = traceback.format_exc()
verbose_logger.debug(f"MCP client get_prompt traceback:\n{error_trace}")
# Log detailed error information
error_type = type(e).__name__
verbose_logger.error(
f"MCP client get_prompt failed - "
f"Error Type: {error_type}, "
f"Error: {str(e)}, "
f"Prompt: {get_prompt_request_params.name}, "
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during get_prompt - "
"the MCP server may have crashed, disconnected, or timed out."
)
raise
async def list_resources(self) -> list[Resource]:
"""List available resources from the server."""
verbose_logger.debug(
f"MCP client listing resources from {self.server_url or 'stdio'}"
)
async def _list_resources_operation(session: ClientSession):
return await session.list_resources()
try:
result = await self.run_with_session(_list_resources_operation)
resource_count = len(result.resources)
resource_names = [resource.name for resource in result.resources]
verbose_logger.info(
f"MCP client listed {resource_count} resources from {self.server_url or 'stdio'}: {resource_names}"
)
return result.resources
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_resources was cancelled")
raise
except Exception as e:
error_type = type(e).__name__
verbose_logger.error(
f"MCP client list_resources failed - "
f"Error Type: {error_type}, "
f"Error: {str(e)}, "
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during list_resources - "
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
async def list_resource_templates(self) -> list[ResourceTemplate]:
"""List available resource templates from the server."""
verbose_logger.debug(
f"MCP client listing resource templates from {self.server_url or 'stdio'}"
)
async def _list_resource_templates_operation(session: ClientSession):
return await session.list_resource_templates()
try:
result = await self.run_with_session(_list_resource_templates_operation)
resource_template_count = len(result.resourceTemplates)
resource_template_names = [
resourceTemplate.name for resourceTemplate in result.resourceTemplates
]
verbose_logger.info(
f"MCP client listed {resource_template_count} resource templates from {self.server_url or 'stdio'}: {resource_template_names}"
)
return result.resourceTemplates
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_resource_templates was cancelled")
raise
except Exception as e:
error_type = type(e).__name__
verbose_logger.error(
f"MCP client list_resource_templates failed - "
f"Error Type: {error_type}, "
f"Error: {str(e)}, "
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during list_resource_templates - "
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
async def read_resource(self, url: AnyUrl) -> ReadResourceResult:
"""Fetch resource contents from the MCP server."""
verbose_logger.info(f"MCP client fetching resource '{url}'")
async def _read_resource_operation(session: ClientSession):
verbose_logger.debug("MCP client sending read_resource request to session")
return await session.read_resource(url)
try:
read_resource_result = await self.run_with_session(_read_resource_operation)
verbose_logger.info(
f"MCP client read_resource '{url}' completed successfully"
)
return read_resource_result
except asyncio.CancelledError:
verbose_logger.warning("MCP client read_resource was cancelled")
raise
except Exception as e:
import traceback
error_trace = traceback.format_exc()
verbose_logger.debug(f"MCP client read_resource traceback:\n{error_trace}")
# Log detailed error information
error_type = type(e).__name__
verbose_logger.error(
f"MCP client read_resource failed - "
f"Error Type: {error_type}, "
f"Error: {str(e)}, "
f"Url: {url}, "
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during read_resource - "
"the MCP server may have crashed, disconnected, or timed out."
)
raise
@@ -16,10 +16,19 @@ from urllib.parse import urlparse
from fastapi import HTTPException
from httpx import HTTPStatusError
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp import ReadResourceResult, Resource
from mcp.types import (
CallToolRequestParams as MCPCallToolRequestParams,
GetPromptRequestParams,
GetPromptResult,
Prompt,
ResourceTemplate,
)
from mcp.types import CallToolResult
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
import litellm
from litellm._logging import verbose_logger
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
@@ -29,11 +38,11 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.utils import (
add_server_prefix_to_tool_name,
get_server_name_prefix_tool_mcp,
add_server_prefix_to_name,
get_server_prefix,
is_tool_name_prefixed,
normalize_server_name,
split_server_prefix_from_name,
validate_mcp_server_name,
)
from litellm.proxy._types import (
@@ -357,7 +366,7 @@ class MCPServerManager:
base_tool_name = operation_id.replace(" ", "_").lower()
# Add server prefix to tool name
prefixed_tool_name = add_server_prefix_to_tool_name(
prefixed_tool_name = add_server_prefix_to_name(
base_tool_name, server_prefix
)
@@ -716,6 +725,190 @@ class MCPServerManager:
)
return []
async def get_prompts_from_server(
self,
server: MCPServer,
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
extra_headers: Optional[Dict[str, str]] = None,
add_prefix: bool = True,
) -> List[Prompt]:
"""
Helper method to get prompts from a single MCP server with prefixed names.
Args:
server (MCPServer): The server to query prompts from
mcp_auth_header: Optional auth header for MCP server
Returns:
List[Prompt]: List of prompts available on the server with prefixed names
"""
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"get_prompts_from_server for {server.name}...")
client = None
try:
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
client = self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
)
prompts = await client.list_prompts()
prefixed_or_original_prompts = self._create_prefixed_prompts(
prompts, server, add_prefix=add_prefix
)
return prefixed_or_original_prompts
except Exception as e:
verbose_logger.warning(
f"Failed to get prompts from server {server.name}: {str(e)}"
)
return []
async def get_resources_from_server(
self,
server: MCPServer,
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
extra_headers: Optional[Dict[str, str]] = None,
add_prefix: bool = True,
) -> List[Resource]:
"""Fetch available resources from a single MCP server."""
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"get_resources_from_server for {server.name}...")
client = None
try:
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
client = self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
)
resources = await client.list_resources()
prefixed_resources = self._create_prefixed_resources(
resources, server, add_prefix=add_prefix
)
return prefixed_resources
except Exception as e:
verbose_logger.warning(
f"Failed to get resources from server {server.name}: {str(e)}"
)
return []
async def get_resource_templates_from_server(
self,
server: MCPServer,
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
extra_headers: Optional[Dict[str, str]] = None,
add_prefix: bool = True,
) -> List[ResourceTemplate]:
"""Fetch available resource templates from a single MCP server."""
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"get_resource_templates_from_server for {server.name}...")
client = None
try:
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
client = self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
)
resource_templates = await client.list_resource_templates()
prefixed_templates = self._create_prefixed_resource_templates(
resource_templates, server, add_prefix=add_prefix
)
return prefixed_templates
except Exception as e:
verbose_logger.warning(
f"Failed to get resource templates from server {server.name}: {str(e)}"
)
return []
async def read_resource_from_server(
self,
server: MCPServer,
url: AnyUrl,
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
extra_headers: Optional[Dict[str, str]] = None,
) -> ReadResourceResult:
"""Read resource contents from a specific MCP server."""
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"read_resource_from_server for {server.name}...")
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
client = self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
)
return await client.read_resource(url)
async def get_prompt_from_server(
self,
server: MCPServer,
prompt_name: str,
arguments: Optional[Dict[str, Any]] = None,
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
extra_headers: Optional[Dict[str, str]] = None,
) -> GetPromptResult:
"""Fetch a specific prompt definition from a single MCP server."""
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"get_prompt_from_server for {server.name}...")
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
client = self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
)
get_prompt_request_params = GetPromptRequestParams(
name=prompt_name,
arguments=arguments,
)
return await client.get_prompt(get_prompt_request_params)
async def _descovery_metadata(
self,
server_url: str,
@@ -1026,7 +1219,7 @@ class MCPServerManager:
prefix = get_server_prefix(server)
for tool in tools:
prefixed_name = add_server_prefix_to_tool_name(tool.name, prefix)
prefixed_name = add_server_prefix_to_name(tool.name, prefix)
name_to_use = prefixed_name if add_prefix else tool.name
@@ -1046,6 +1239,82 @@ class MCPServerManager:
)
return prefixed_tools
def _create_prefixed_prompts(
self, prompts: List[Prompt], server: MCPServer, add_prefix: bool = True
) -> List[Prompt]:
"""
Create prefixed prompts and update prompt mapping.
Args:
prompts: List of original prompts from server
server: Server instance
Returns:
List of prompts with prefixed names
"""
prefixed_prompts = []
prefix = get_server_prefix(server)
for prompt in prompts:
prefixed_name = add_server_prefix_to_name(prompt.name, prefix)
name_to_use = prefixed_name if add_prefix else prompt.name
prompt.name = name_to_use
prefixed_prompts.append(prompt)
verbose_logger.info(
f"Successfully fetched {len(prefixed_prompts)} prompts from server {server.name}"
)
return prefixed_prompts
def _create_prefixed_resources(
self, resources: List[Resource], server: MCPServer, add_prefix: bool = True
) -> List[Resource]:
"""Prefix resource names and track origin server for read requests."""
prefixed_resources: List[Resource] = []
prefix = get_server_prefix(server)
for resource in resources:
name_to_use = (
add_server_prefix_to_name(resource.name, prefix)
if add_prefix
else resource.name
)
resource.name = name_to_use
prefixed_resources.append(resource)
verbose_logger.info(
f"Successfully fetched {len(prefixed_resources)} resources from server {server.name}"
)
return prefixed_resources
def _create_prefixed_resource_templates(
self,
resource_templates: List[ResourceTemplate],
server: MCPServer,
add_prefix: bool = True,
) -> List[ResourceTemplate]:
"""Prefix resource template names for multi-server scenarios."""
prefixed_templates: List[ResourceTemplate] = []
prefix = get_server_prefix(server)
for resource_template in resource_templates:
name_to_use = (
add_server_prefix_to_name(resource_template.name, prefix)
if add_prefix
else resource_template.name
)
resource_template.name = name_to_use
prefixed_templates.append(resource_template)
verbose_logger.info(
f"Successfully fetched {len(prefixed_templates)} resource templates from server {server.name}"
)
return prefixed_templates
def check_allowed_or_banned_tools(self, tool_name: str, server: MCPServer) -> bool:
"""
Check if the tool is allowed or banned for the given server
@@ -1080,7 +1349,7 @@ class MCPServerManager:
HTTPException: If allowed_params is configured for this tool but arguments contain disallowed params
"""
from litellm.proxy._experimental.mcp_server.utils import (
get_server_name_prefix_tool_mcp,
split_server_prefix_from_name,
)
# If no allowed_params configured, return all arguments
@@ -1088,7 +1357,7 @@ class MCPServerManager:
return
# Get the unprefixed tool name to match against config
unprefixed_tool_name, _ = get_server_name_prefix_tool_mcp(tool_name)
unprefixed_tool_name, _ = split_server_prefix_from_name(tool_name)
# Check both prefixed and unprefixed tool names
allowed_params_list = server.allowed_params.get(
@@ -1489,7 +1758,7 @@ class MCPServerManager:
start_time = datetime.datetime.now()
# Get the MCP server
prefixed_tool_name = add_server_prefix_to_tool_name(name, server_name)
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
if mcp_server is None:
raise ValueError(f"Tool {name} not found")
@@ -1595,7 +1864,7 @@ class MCPServerManager:
for tool in tools:
# The tool.name here is already prefixed from _get_tools_from_server
# Extract original name for mapping
original_name, _ = get_server_name_prefix_tool_mcp(tool.name)
original_name, _ = split_server_prefix_from_name(tool.name)
self.tool_name_to_mcp_server_name_mapping[original_name] = server.name
self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name
@@ -1623,7 +1892,7 @@ class MCPServerManager:
(
original_tool_name,
server_name_from_prefix,
) = get_server_name_prefix_tool_mcp(tool_name)
) = split_server_prefix_from_name(tool_name)
if original_tool_name in self.tool_name_to_mcp_server_name_mapping:
for server in self.get_registry().values():
if normalize_server_name(server.name) == normalize_server_name(
+682 -39
View File
@@ -8,7 +8,15 @@ from datetime import datetime
from typing import Any, AsyncIterator, Dict, List, Optional, Tuple, Union
from fastapi import FastAPI, HTTPException
from pydantic import ConfigDict
from mcp import ReadResourceResult, Resource
from mcp.server.lowlevel.helper_types import ReadResourceContents
from mcp.types import (
BlobResourceContents,
GetPromptResult,
ResourceTemplate,
TextResourceContents,
)
from pydantic import AnyUrl, ConfigDict
from starlette.types import Receive, Scope, Send
from litellm._logging import verbose_logger
@@ -54,6 +62,7 @@ if MCP_AVAILABLE:
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
from mcp.types import EmbeddedResource, ImageContent, TextContent
from mcp.types import Tool as MCPTool
from mcp.types import Prompt
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
MCPAuthenticatedUser,
@@ -66,7 +75,7 @@ if MCP_AVAILABLE:
global_mcp_tool_registry,
)
from litellm.proxy._experimental.mcp_server.utils import (
get_server_name_prefix_tool_mcp,
split_server_prefix_from_name,
)
######################################################
@@ -303,6 +312,208 @@ if MCP_AVAILABLE:
return response
@server.list_prompts()
async def list_prompts() -> List[Prompt]:
"""
List all available prompts
"""
try:
# Get user authentication from context variable
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = get_auth_context()
verbose_logger.debug(
f"MCP list_prompts - User API Key Auth from context: {user_api_key_auth}"
)
verbose_logger.debug(
f"MCP list_prompts - MCP servers from context: {mcp_servers}"
)
verbose_logger.debug(
f"MCP list_prompts - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
# Get mcp_servers from context variable
verbose_logger.debug("MCP list_prompts - Calling _list_prompts")
prompts = await _list_mcp_prompts(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
verbose_logger.info(
f"MCP list_prompts - Successfully returned {len(prompts)} prompts"
)
return prompts
except Exception as e:
verbose_logger.exception(f"Error in list_prompts endpoint: {str(e)}")
# Return empty list instead of failing completely
# This prevents the HTTP stream from failing and allows the client to get a response
return []
@server.get_prompt()
async def get_prompt(
name: str, arguments: dict[str, str] | None
) -> GetPromptResult:
"""
Get a specific prompt with the provided arguments
Args:
name (str): Name of the prompt to get
arguments (Dict[str, Any] | None): Arguments to pass to the prompt
Returns:
GetPromptResult: Getting prompt execution results
"""
# Validate arguments
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = get_auth_context()
verbose_logger.debug(
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
)
return await mcp_get_prompt(
name=name,
arguments=arguments,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
@server.list_resources()
async def list_resources() -> List[Resource]:
"""List all available resources."""
try:
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = get_auth_context()
verbose_logger.debug(
f"MCP list_resources - User API Key Auth from context: {user_api_key_auth}"
)
verbose_logger.debug(
f"MCP list_resources - MCP servers from context: {mcp_servers}"
)
verbose_logger.debug(
f"MCP list_resources - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
resources = await _list_mcp_resources(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
verbose_logger.info(
f"MCP list_resources - Successfully returned {len(resources)} resources"
)
return resources
except Exception as e:
verbose_logger.exception(f"Error in list_resources endpoint: {str(e)}")
return []
@server.list_resource_templates()
async def list_resource_templates() -> List[ResourceTemplate]:
"""List all available resource templates."""
try:
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = get_auth_context()
verbose_logger.debug(
f"MCP list_resource_templates - User API Key Auth from context: {user_api_key_auth}"
)
verbose_logger.debug(
f"MCP list_resource_templates - MCP servers from context: {mcp_servers}"
)
verbose_logger.debug(
f"MCP list_resource_templates - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
resource_templates = await _list_mcp_resource_templates(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
verbose_logger.info(
"MCP list_resource_templates - Successfully returned "
f"{len(resource_templates)} resource templates"
)
return resource_templates
except Exception as e:
verbose_logger.exception(
f"Error in list_resource_templates endpoint: {str(e)}"
)
return []
@server.read_resource()
async def read_resource(url: AnyUrl) -> list[ReadResourceContents]:
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = get_auth_context()
read_resource_result = await mcp_read_resource(
url=url,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
normalized_contents: List[ReadResourceContents] = []
for content in read_resource_result.contents:
if isinstance(content, TextResourceContents):
normalized_contents.append(
ReadResourceContents(
content=content.text,
mime_type=content.mimeType,
)
)
elif isinstance(content, BlobResourceContents):
normalized_contents.append(
ReadResourceContents(
content=content.blob,
mime_type=None,
)
)
return normalized_contents
########################################################
############ End of MCP Server Routes ##################
########################################################
@@ -379,7 +590,7 @@ if MCP_AVAILABLE:
True if the tool name (prefixed or unprefixed) is in the filter list
"""
from litellm.proxy._experimental.mcp_server.utils import (
get_server_name_prefix_tool_mcp,
split_server_prefix_from_name,
)
# Check if the full name is in the list
@@ -387,7 +598,7 @@ if MCP_AVAILABLE:
return True
# Check if the unprefixed name is in the list
unprefixed_name, _ = get_server_name_prefix_tool_mcp(tool_name)
unprefixed_name, _ = split_server_prefix_from_name(tool_name)
return unprefixed_name in filter_list
def filter_tools_by_allowed_tools(
@@ -428,6 +639,56 @@ if MCP_AVAILABLE:
return tools_to_return
async def _get_allowed_mcp_servers(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_servers: Optional[List[str]],
) -> List[MCPServer]:
"""Return allowed MCP servers for a request after applying filters."""
allowed_mcp_server_ids = (
await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth)
)
allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids(
allowed_mcp_server_ids
)
if mcp_servers is not None:
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
)
return allowed_mcp_servers
def _prepare_mcp_server_headers(
server: MCPServer,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
mcp_auth_header: Optional[str],
oauth2_headers: Optional[Dict[str, str]],
raw_headers: Optional[Dict[str, str]],
) -> Tuple[Optional[Union[Dict[str, str], str]], Optional[Dict[str, str]]]:
"""Build auth and extra headers for a server."""
server_auth_header: Optional[Union[Dict[str, str], str]] = None
if mcp_server_auth_headers and server.alias is not None:
server_auth_header = mcp_server_auth_headers.get(server.alias)
elif mcp_server_auth_headers and server.server_name is not None:
server_auth_header = mcp_server_auth_headers.get(server.server_name)
extra_headers: Optional[Dict[str, str]] = None
if server.auth_type == MCPAuth.oauth2:
extra_headers = oauth2_headers
if server.extra_headers and raw_headers:
if extra_headers is None:
extra_headers = {}
for header in server.extra_headers:
if header in raw_headers:
extra_headers[header] = raw_headers[header]
if server_auth_header is None:
server_auth_header = mcp_auth_header
return server_auth_header, extra_headers
async def _get_tools_from_mcp_servers(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str],
@@ -452,19 +713,10 @@ if MCP_AVAILABLE:
if not MCP_AVAILABLE:
return []
# Get allowed MCP servers based on user permissions
allowed_mcp_server_ids = (
await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth)
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
)
allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids(
allowed_mcp_server_ids
)
if mcp_servers is not None:
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
)
# Decide whether to add prefix based on number of allowed servers
add_prefix = not (len(allowed_mcp_servers) == 1)
@@ -475,27 +727,13 @@ if MCP_AVAILABLE:
if server is None:
continue
# Get server-specific auth header if available
server_auth_header: Optional[Union[Dict[str, str], str]] = None
if mcp_server_auth_headers and server.alias is not None:
server_auth_header = mcp_server_auth_headers.get(server.alias)
elif mcp_server_auth_headers and server.server_name is not None:
server_auth_header = mcp_server_auth_headers.get(server.server_name)
extra_headers: Optional[Dict[str, str]] = None
if server.auth_type == MCPAuth.oauth2:
extra_headers = oauth2_headers
if server.extra_headers and raw_headers:
if extra_headers is None:
extra_headers = {}
for header in server.extra_headers:
if header in raw_headers:
extra_headers[header] = raw_headers[header]
# Fall back to deprecated mcp_auth_header if no server-specific header found
if server_auth_header is None:
server_auth_header = mcp_auth_header
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
try:
tools = await global_mcp_server_manager._get_tools_from_server(
@@ -530,6 +768,195 @@ if MCP_AVAILABLE:
return all_tools
async def _get_prompts_from_mcp_servers(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str],
mcp_servers: Optional[List[str]],
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> List[Prompt]:
"""
Helper method to fetch prompt from MCP servers based on server filtering criteria.
Args:
user_api_key_auth: User authentication info for access control
mcp_auth_header: Optional auth header for MCP server (deprecated)
mcp_servers: Optional list of server names/aliases to filter by
mcp_server_auth_headers: Optional dict of server-specific auth headers
oauth2_headers: Optional dict of oauth2 headers
Returns:
List[Prompt]: Combined list of prompts from filtered servers
"""
if not MCP_AVAILABLE:
return []
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
)
# Decide whether to add prefix based on number of allowed servers
add_prefix = not (len(allowed_mcp_servers) == 1)
# Get prompts from each allowed server
all_prompts = []
for server in allowed_mcp_servers:
if server is None:
continue
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
try:
prompts = await global_mcp_server_manager.get_prompts_from_server(
server=server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=add_prefix,
)
all_prompts.extend(prompts)
verbose_logger.debug(
f"Successfully fetched {len(prompts)} prompts from server {server.name}"
)
except Exception as e:
verbose_logger.exception(
f"Error getting prompts from server {server.name}: {str(e)}"
)
# Continue with other servers instead of failing completely
verbose_logger.info(
f"Successfully fetched {len(all_prompts)} prompts total from all MCP servers"
)
return all_prompts
async def _get_resources_from_mcp_servers(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str],
mcp_servers: Optional[List[str]],
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> List[Resource]:
"""Fetch resources from allowed MCP servers."""
if not MCP_AVAILABLE:
return []
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
)
add_prefix = not (len(allowed_mcp_servers) == 1)
all_resources: List[Resource] = []
for server in allowed_mcp_servers:
if server is None:
continue
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
try:
resources = await global_mcp_server_manager.get_resources_from_server(
server=server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=add_prefix,
)
all_resources.extend(resources)
verbose_logger.debug(
f"Successfully fetched {len(resources)} resources from server {server.name}"
)
except Exception as e:
verbose_logger.exception(
f"Error getting resources from server {server.name}: {str(e)}"
)
verbose_logger.info(
f"Successfully fetched {len(all_resources)} resources total from all MCP servers"
)
return all_resources
async def _get_resource_templates_from_mcp_servers(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str],
mcp_servers: Optional[List[str]],
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> List[ResourceTemplate]:
"""Fetch resource templates from allowed MCP servers."""
if not MCP_AVAILABLE:
return []
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
)
add_prefix = not (len(allowed_mcp_servers) == 1)
all_resource_templates: List[ResourceTemplate] = []
for server in allowed_mcp_servers:
if server is None:
continue
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
try:
resource_templates = (
await global_mcp_server_manager.get_resource_templates_from_server(
server=server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=add_prefix,
)
)
all_resource_templates.extend(resource_templates)
verbose_logger.debug(
"Successfully fetched %s resource templates from server %s",
len(resource_templates),
server.name,
)
except Exception as e:
verbose_logger.exception(
"Error getting resource templates from server %s: %s",
server.name,
str(e),
)
verbose_logger.info(
"Successfully fetched %s resource templates total from all MCP servers",
len(all_resource_templates),
)
return all_resource_templates
async def filter_tools_by_key_team_permissions(
tools: List[MCPTool],
server_id: str,
@@ -553,7 +980,7 @@ if MCP_AVAILABLE:
filtered_tools = []
for t in tools:
# Get tool name without server prefix
unprefixed_tool_name, _ = get_server_name_prefix_tool_mcp(t.name)
unprefixed_tool_name, _ = split_server_prefix_from_name(t.name)
if unprefixed_tool_name in allowed_tool_names:
filtered_tools.append(t)
else:
@@ -606,6 +1033,118 @@ if MCP_AVAILABLE:
return managed_tools
async def _list_mcp_prompts(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> List[Prompt]:
"""
List all available MCP prompts.
Args:
user_api_key_auth: User authentication info for access control
mcp_auth_header: Optional auth header for MCP server (deprecated)
mcp_servers: Optional list of server names/aliases to filter by
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
Returns:
List[Prompt]: Combined list of tools from all accessible servers
"""
if not MCP_AVAILABLE:
return []
# Get tools from managed MCP servers with error handling
managed_prompts = []
try:
managed_prompts = await _get_prompts_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
verbose_logger.debug(
f"Successfully fetched {len(managed_prompts)} prompts from managed MCP servers"
)
except Exception as e:
verbose_logger.exception(
f"Error getting tools from managed MCP servers: {str(e)}"
)
# Continue with empty managed tools list instead of failing completely
return managed_prompts
async def _list_mcp_resources(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> List[Resource]:
"""List all available MCP resources."""
if not MCP_AVAILABLE:
return []
managed_resources: List[Resource] = []
try:
managed_resources = await _get_resources_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
verbose_logger.debug(
f"Successfully fetched {len(managed_resources)} resources from managed MCP servers"
)
except Exception as e:
verbose_logger.exception(
f"Error getting resources from managed MCP servers: {str(e)}"
)
return managed_resources
async def _list_mcp_resource_templates(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> List[ResourceTemplate]:
"""List all available MCP resource templates."""
if not MCP_AVAILABLE:
return []
managed_resource_templates: List[ResourceTemplate] = []
try:
managed_resource_templates = await _get_resource_templates_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
verbose_logger.debug(
"Successfully fetched %s resource templates from managed MCP servers",
len(managed_resource_templates),
)
except Exception as e:
verbose_logger.exception(
"Error getting resource templates from managed MCP servers: %s",
str(e),
)
return managed_resource_templates
@client
async def call_mcp_tool(
name: str,
@@ -647,7 +1186,7 @@ if MCP_AVAILABLE:
mcp_server: Optional[MCPServer] = None
# Remove prefix from tool name for logging and processing
original_tool_name, server_name = get_server_name_prefix_tool_mcp(name)
original_tool_name, server_name = split_server_prefix_from_name(name)
# If tool name is unprefixed, resolve its server so we can enforce permissions
if not server_name:
@@ -735,6 +1274,110 @@ if MCP_AVAILABLE:
)
return response
async def mcp_get_prompt(
name: str,
arguments: Optional[Dict[str, Any]] = None,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> GetPromptResult:
"""
Fetch a specific MCP prompt, handling both prefixed and unprefixed names.
"""
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
)
if not allowed_mcp_servers:
raise HTTPException(
status_code=403,
detail="User not allowed to get this prompt.",
)
# Decide whether to add prefix based on number of allowed servers
add_prefix = not (len(allowed_mcp_servers) == 1)
if add_prefix:
original_prompt_name, server_name = split_server_prefix_from_name(name)
else:
original_prompt_name = name
server_name = allowed_mcp_servers[0].name
server = next((s for s in allowed_mcp_servers if s.name == server_name), None)
if server is None:
raise HTTPException(
status_code=403,
detail="User not allowed to get this prompt.",
)
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
return await global_mcp_server_manager.get_prompt_from_server(
server=server,
prompt_name=original_prompt_name,
arguments=arguments,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
)
async def mcp_read_resource(
url: AnyUrl,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
) -> ReadResourceResult:
"""Read resource contents from upstream MCP servers."""
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
)
if not allowed_mcp_servers:
raise HTTPException(
status_code=403,
detail="User not allowed to read this resource.",
)
if len(allowed_mcp_servers) != 1:
raise HTTPException(
status_code=400,
detail=(
"Multiple MCP servers configured; read_resource currently "
"supports exactly one allowed server."
),
)
server = allowed_mcp_servers[0]
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
return await global_mcp_server_manager.read_resource_from_server(
server=server,
url=url,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
)
def _get_standard_logging_mcp_tool_call(
name: str,
arguments: Dict[str, Any],
+37 -44
View File
@@ -13,6 +13,7 @@ LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM"
MCP_TOOL_PREFIX_SEPARATOR = os.environ.get("MCP_TOOL_PREFIX_SEPARATOR", "-")
MCP_TOOL_PREFIX_FORMAT = "{server_name}{separator}{tool_name}"
def is_mcp_available() -> bool:
"""
Returns True if the MCP module is available, False otherwise
@@ -23,92 +24,81 @@ def is_mcp_available() -> bool:
except ImportError:
return False
def normalize_server_name(server_name: str) -> str:
"""
Normalize server name by replacing spaces with underscores
"""
return server_name.replace(" ", "_")
def validate_and_normalize_mcp_server_payload(payload: Any) -> None:
"""
Validate and normalize MCP server payload fields (server_name and alias).
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)
Args:
payload: The payload object containing server_name and alias fields
Raises:
HTTPException: If validation fails
"""
# Server name validation: disallow '-'
if hasattr(payload, 'server_name') and payload.server_name:
if hasattr(payload, "server_name") and payload.server_name:
validate_mcp_server_name(payload.server_name, raise_http_exception=True)
# Alias validation: disallow '-'
if hasattr(payload, 'alias') and payload.alias:
if hasattr(payload, "alias") and payload.alias:
validate_mcp_server_name(payload.alias, raise_http_exception=True)
# Alias normalization and defaulting
alias = getattr(payload, 'alias', None)
server_name = getattr(payload, 'server_name', None)
alias = getattr(payload, "alias", None)
server_name = getattr(payload, "server_name", None)
if not alias and server_name:
alias = normalize_server_name(server_name)
elif alias:
alias = normalize_server_name(alias)
# Update the payload with normalized alias
if hasattr(payload, 'alias'):
if hasattr(payload, "alias"):
payload.alias = alias
def add_server_prefix_to_tool_name(tool_name: str, server_name: str) -> str:
"""
Add server name prefix to tool name
Args:
tool_name: Original tool name
server_name: MCP server name
Returns:
Prefixed tool name in format: server_name::tool_name
"""
def add_server_prefix_to_name(name: str, server_name: str) -> str:
"""Add server name prefix to any MCP resource name."""
formatted_server_name = normalize_server_name(server_name)
return MCP_TOOL_PREFIX_FORMAT.format(
server_name=formatted_server_name,
separator=MCP_TOOL_PREFIX_SEPARATOR,
tool_name=tool_name
tool_name=name,
)
def get_server_prefix(server: Any) -> str:
"""Return the prefix for a server: alias if present, else server_name, else server_id"""
if hasattr(server, 'alias') and server.alias:
if hasattr(server, "alias") and server.alias:
return server.alias
if hasattr(server, 'server_name') and server.server_name:
if hasattr(server, "server_name") and server.server_name:
return server.server_name
if hasattr(server, 'server_id'):
if hasattr(server, "server_id"):
return server.server_id
return ""
def get_server_name_prefix_tool_mcp(prefixed_tool_name: str) -> Tuple[str, str]:
"""
Remove server name prefix from tool name
Args:
prefixed_tool_name: Tool name with server prefix
Returns:
Tuple of (original_tool_name, server_name)
"""
if MCP_TOOL_PREFIX_SEPARATOR in prefixed_tool_name:
parts = prefixed_tool_name.split(MCP_TOOL_PREFIX_SEPARATOR, 1)
def split_server_prefix_from_name(prefixed_name: str) -> Tuple[str, str]:
"""Return the unprefixed name plus the server name used as prefix."""
if MCP_TOOL_PREFIX_SEPARATOR in prefixed_name:
parts = prefixed_name.split(MCP_TOOL_PREFIX_SEPARATOR, 1)
if len(parts) == 2:
return parts[1], parts[0] # tool_name, server_name
return prefixed_tool_name, "" # No prefix found, return original name
return parts[1], parts[0]
return prefixed_name, ""
def is_tool_name_prefixed(tool_name: str) -> bool:
"""
@@ -122,14 +112,17 @@ def is_tool_name_prefixed(tool_name: str) -> bool:
"""
return MCP_TOOL_PREFIX_SEPARATOR in tool_name
def validate_mcp_server_name(server_name: str, raise_http_exception: bool = False) -> None:
def validate_mcp_server_name(
server_name: str, raise_http_exception: bool = False
) -> None:
"""
Validate that MCP server name does not contain 'MCP_TOOL_PREFIX_SEPARATOR'.
Args:
server_name: The server name to validate
raise_http_exception: If True, raises HTTPException instead of generic Exception
Raises:
Exception or HTTPException: If server name contains 'MCP_TOOL_PREFIX_SEPARATOR'
"""
@@ -138,9 +131,9 @@ def validate_mcp_server_name(server_name: str, raise_http_exception: bool = Fals
if raise_http_exception:
from fastapi import HTTPException
from starlette import status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": error_message}
status_code=status.HTTP_400_BAD_REQUEST, detail={"error": error_message}
)
else:
raise Exception(error_message)
@@ -1,7 +1,7 @@
from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Tuple, Union
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.utils import get_server_name_prefix_tool_mcp
from litellm.proxy._experimental.mcp_server.utils import split_server_prefix_from_name
from litellm.responses.main import aresponses
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
from litellm.types.llms.openai import ResponsesAPIResponse, ToolParam
@@ -123,9 +123,11 @@ class LiteLLM_Proxy_MCP_Handler:
for server in allowed_mcp_servers:
if server is None:
continue
server_name = getattr(server, "server_name", None) or getattr(
server, "alias", None
) or getattr(server, "name", None)
server_name = (
getattr(server, "server_name", None)
or getattr(server, "alias", None)
or getattr(server, "name", None)
)
if isinstance(server_name, str):
server_names.append(server_name)
@@ -161,7 +163,7 @@ class LiteLLM_Proxy_MCP_Handler:
if len(allowed_mcp_servers) == 1:
tool_server_map[tool_name] = allowed_mcp_servers[0]
else:
tool_server_map[tool_name], _ = get_server_name_prefix_tool_mcp(
tool_server_map[tool_name], _ = split_server_prefix_from_name(
tool_name
)
Generated
+2470 -1357
View File
File diff suppressed because it is too large Load Diff
+5 -5
View File
@@ -34,7 +34,7 @@ pydantic = "^2.5.0"
jsonschema = "^4.22.0"
numpydoc = {version = "*", optional = true} # used in utils.py
uvicorn = {version = "^0.29.0", optional = true}
uvicorn = {version = "^0.31.1", optional = true}
uvloop = {version = "^0.21.0", optional = true, markers="sys_platform != 'win32'"}
gunicorn = {version = "^23.0.0", optional = true}
fastapi = {version = ">=0.120.1", optional = true}
@@ -44,11 +44,11 @@ rq = {version = "*", optional = true}
orjson = {version = "^3.9.7", optional = true}
apscheduler = {version = "^3.10.4", optional = true}
fastapi-sso = { version = "^0.16.0", optional = true }
PyJWT = { version = "^2.8.0", optional = true }
PyJWT = { version = "^2.10.1", optional = true, python = ">=3.9" }
python-multipart = { version = "^0.0.18", optional = true}
cryptography = {version = "*", optional = true}
prisma = {version = "0.11.0", optional = true}
azure-identity = {version = "^1.15.0", optional = true}
azure-identity = {version = "^1.15.0", optional = true, python = ">=3.9"}
azure-keyvault-secrets = {version = "^4.8.0", optional = true}
azure-storage-blob = {version="^12.25.1", optional=true}
google-cloud-kms = {version = "^2.21.3", optional = true}
@@ -58,7 +58,7 @@ pynacl = {version = "^1.5.0", optional = true}
websockets = {version = "^13.1.0", optional = true}
boto3 = {version = "1.36.0", optional = true}
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
mcp = {version = "^1.10.0", optional = true, python = ">=3.10"}
mcp = {version = "^1.21.2", optional = true, python = ">=3.10"}
litellm-proxy-extras = {version = "0.4.6", optional = true}
rich = {version = "13.7.1", optional = true}
litellm-enterprise = {version = "0.1.22", optional = true}
@@ -152,7 +152,7 @@ prometheus-client = "0.20.0"
opentelemetry-api = "1.25.0"
opentelemetry-sdk = "1.25.0"
opentelemetry-exporter-otlp = "1.25.0"
azure-identity = "^1.15.0"
azure-identity = {version = "^1.15.0", python = ">=3.9"}
[build-system]
requires = ["poetry-core", "wheel"]
+4 -4
View File
@@ -6,7 +6,7 @@ fastapi==0.120.1 # server dep
starlette==0.49.1 # starlette fastapi dep
backoff==2.2.1 # server dep
pyyaml==6.0.2 # server dep
uvicorn==0.29.0 # server dep
uvicorn==0.31.1 # server dep
gunicorn==23.0.0 # server dep
fastuuid==0.13.5 # for uuid4
uvloop==0.21.0 # uvicorn dep, gives us much better performance under load
@@ -19,7 +19,7 @@ google-cloud-aiplatform==1.47.0 # for vertex ai calls
google-cloud-iam==2.19.1 # for GCP IAM Redis authentication
google-genai==1.22.0
anthropic[vertex]==0.54.0
mcp==1.10.1 # for MCP server
mcp==1.21.2 ; python_version >= "3.10" # for MCP server
google-generativeai==0.5.0 # for vertex ai calls
async_generator==1.10.0 # for async ollama calls
langfuse==2.59.7 # for langfuse self-hosted logging
@@ -29,11 +29,11 @@ orjson==3.11.2 # fast /embedding responses
polars==1.31.0 # for data processing
apscheduler==3.10.4 # for resetting budget in background
fastapi-sso==0.16.0 # admin UI, SSO
pyjwt[crypto]==2.9.0
pyjwt[crypto]==2.10.1 ; python_version >= "3.9"
python-multipart==0.0.18 # admin UI
Pillow==11.0.0
azure-ai-contentsafety==1.0.0 # for azure content safety
azure-identity==1.16.1 # for azure content safety
azure-identity==1.16.1 ; python_version >= "3.9" # for azure content safety
azure-keyvault==4.2.0 # for azure KMS integration
azure-storage-file-datalake==12.20.0 # for azure buck storage logging
opentelemetry-api==1.25.0
+8 -61
View File
@@ -1481,80 +1481,27 @@ def test_normalize_server_name():
assert normalize_server_name(" ") == "___"
def test_add_server_prefix_to_tool_name():
"""
Test that add_server_prefix_to_tool_name correctly formats tool names.
"""
from litellm.proxy._experimental.mcp_server.utils import (
add_server_prefix_to_tool_name,
)
def test_add_server_prefix_to_name():
"""Ensure add_server_prefix_to_name correctly formats resource names."""
from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name
# Test basic prefixing
result = add_server_prefix_to_tool_name("send_email", "My Server")
result = add_server_prefix_to_name("send_email", "My Server")
assert result == "My_Server-send_email"
# Test with server name that already has underscores
result = add_server_prefix_to_tool_name("create_event", "my_server")
result = add_server_prefix_to_name("create_event", "my_server")
assert result == "my_server-create_event"
# Test with empty tool name
result = add_server_prefix_to_tool_name("", "My Server")
# Test with empty name
result = add_server_prefix_to_name("", "My Server")
assert result == "My_Server-"
# Test with empty server name
result = add_server_prefix_to_tool_name("send_email", "")
result = add_server_prefix_to_name("send_email", "")
assert result == "-send_email"
@pytest.mark.asyncio
async def test_mcp_protocol_version_passed_to_client():
"""Test that MCP protocol version from request is correctly passed to MCPClient."""
# Create a test manager
test_manager = MCPServerManager()
# Mock MCPClient
mock_client = AsyncMock()
mock_client.list_tools = AsyncMock(return_value=[])
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=None)
def mock_client_constructor(*args, **kwargs):
# Verify that the protocol version from request is used
if "protocol_version" in kwargs:
assert kwargs["protocol_version"] == "2025-03-26"
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
mock_client_constructor,
):
# Load a test server
await test_manager.load_servers_from_config(
{
"test_server": {
"url": "https://test-server.com/mcp",
"transport": "http",
"description": "Test Server",
}
}
)
allowed_server_ids = list(test_manager.get_registry().keys())
assert allowed_server_ids, "Expected registry to contain configured server"
with patch.object(
test_manager,
"get_allowed_mcp_servers",
new=AsyncMock(return_value=allowed_server_ids),
):
# Call list_tools with a specific protocol version from request
await test_manager.list_tools()
# Verify the client was created with the correct protocol version
mock_client.list_tools.assert_called()
def test_get_server_auth_header_with_alias():
"""Test _get_server_auth_header function with server alias."""
from litellm.proxy._experimental.mcp_server.rest_endpoints import (
@@ -4,7 +4,11 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
from mcp import ReadResourceResult, Resource
from mcp.types import Prompt, ResourceTemplate, TextResourceContents
@pytest.mark.asyncio
@@ -70,6 +74,315 @@ async def test_mcp_server_tool_call_body_contains_request_data():
assert body["arguments"] == tool_arguments
@pytest.mark.asyncio
async def test_get_prompts_from_mcp_servers_success():
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_prompts_from_mcp_servers,
)
except ImportError:
pytest.skip("MCP server not available")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user")
server_a = MagicMock(name="server_a_obj")
server_a.name = "server_a"
server_a.alias = "server_a"
server_a.server_name = "server_a"
server_a.server_id = "a"
server_a.auth_type = None
server_a.extra_headers = None
server_b = MagicMock(name="server_b_obj")
server_b.name = "server_b"
server_b.alias = "server_b"
server_b.server_name = "server_b"
server_b.server_id = "b"
server_b.auth_type = None
server_b.extra_headers = None
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server_a, server_b]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
mock_manager.get_prompts_from_server = AsyncMock(
side_effect=[
[Prompt(name="hello", description="hi")],
[Prompt(name="howdy", description="hey")],
]
)
prompts = await _get_prompts_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
)
mock_allowed.assert_awaited_once()
assert mock_headers.call_count == 2
assert mock_manager.get_prompts_from_server.await_count == 2
assert {prompt.name for prompt in prompts} == {"hello", "howdy"}
@pytest.mark.asyncio
async def test_get_resources_from_mcp_servers_success():
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_resources_from_mcp_servers,
)
except ImportError:
pytest.skip("MCP server not available")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="user")
server_a = MagicMock(name="server_a_obj")
server_a.name = "server_a"
server_a.alias = "server_a"
server_a.server_name = "server_a"
server_a.server_id = "a"
server_a.auth_type = None
server_a.extra_headers = None
server_b = MagicMock(name="server_b_obj")
server_b.name = "server_b"
server_b.alias = "server_b"
server_b.server_name = "server_b"
server_b.server_id = "b"
server_b.auth_type = None
server_b.extra_headers = None
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server_a, server_b]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
mock_manager.get_resources_from_server = AsyncMock(
side_effect=[
[
Resource(
name="resource_a",
uri="https://example.com/a",
)
],
[
Resource(
name="resource_b",
uri="https://example.com/b",
)
],
]
)
resources = await _get_resources_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
)
mock_allowed.assert_awaited_once()
assert mock_headers.call_count == 2
assert mock_manager.get_resources_from_server.await_count == 2
assert {resource.name for resource in resources} == {
"resource_a",
"resource_b",
}
@pytest.mark.asyncio
async def test_get_resource_templates_from_mcp_servers_success():
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_resource_templates_from_mcp_servers,
)
except ImportError:
pytest.skip("MCP server not available")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="user")
server = MagicMock(name="server_obj")
server.name = "server"
server.alias = "server"
server.server_name = "server"
server.server_id = "server-id"
server.auth_type = None
server.extra_headers = None
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
mock_manager.get_resource_templates_from_server = AsyncMock(
return_value=[
ResourceTemplate(
name="template",
description="desc",
uriTemplate="https://example.com/resource/{id}",
)
]
)
templates = await _get_resource_templates_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
)
mock_allowed.assert_awaited_once()
mock_headers.assert_called_once()
mock_manager.get_resource_templates_from_server.assert_awaited_once()
assert [template.name for template in templates] == ["template"]
@pytest.mark.asyncio
async def test_mcp_get_prompt_success():
try:
from litellm.proxy._experimental.mcp_server.server import mcp_get_prompt
except ImportError:
pytest.skip("MCP server not available")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user")
server = MagicMock()
server.name = "server_a"
prompt_result = MagicMock(name="prompt_result")
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
mock_manager.get_prompt_from_server = AsyncMock(return_value=prompt_result)
result = await mcp_get_prompt(
name="hello",
arguments={"foo": "bar"},
user_api_key_auth=user_api_key_auth,
)
mock_allowed.assert_awaited_once()
mock_headers.assert_called_once_with(
server=server,
mcp_server_auth_headers=None,
mcp_auth_header=None,
oauth2_headers=None,
raw_headers=None,
)
mock_manager.get_prompt_from_server.assert_awaited_once_with(
server=server,
prompt_name="hello",
arguments={"foo": "bar"},
mcp_auth_header={"Authorization": "token"},
extra_headers={"X-Test": "1"},
)
assert result is prompt_result
@pytest.mark.asyncio
async def test_mcp_read_resource_success():
try:
from litellm.proxy._experimental.mcp_server.server import mcp_read_resource
except ImportError:
pytest.skip("MCP server not available")
user_api_key_auth = UserAPIKeyAuth(api_key="key", user_id="user")
server = MagicMock()
server.name = "server"
read_result = ReadResourceResult(
contents=[
TextResourceContents(
uri="https://example.com/resource",
text="hello world",
mimeType="text/plain",
)
]
)
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
mock_manager.read_resource_from_server = AsyncMock(return_value=read_result)
result = await mcp_read_resource(
url="https://example.com/resource",
user_api_key_auth=user_api_key_auth,
)
mock_allowed.assert_awaited_once()
mock_headers.assert_called_once_with(
server=server,
mcp_server_auth_headers=None,
mcp_auth_header=None,
oauth2_headers=None,
raw_headers=None,
)
mock_manager.read_resource_from_server.assert_awaited_once_with(
server=server,
url="https://example.com/resource",
mcp_auth_header={"Authorization": "token"},
extra_headers={"X-Test": "1"},
)
assert result is read_result
@pytest.mark.asyncio
async def test_mcp_read_resource_multiple_servers_error():
try:
from litellm.proxy._experimental.mcp_server.server import mcp_read_resource
except ImportError:
pytest.skip("MCP server not available")
user_api_key_auth = UserAPIKeyAuth(api_key="key", user_id="user")
server_a = MagicMock()
server_b = MagicMock()
server_a.name = "server_a"
server_b.name = "server_b"
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server_a, server_b]),
) as mock_allowed:
with pytest.raises(HTTPException) as exc_info:
await mcp_read_resource(
url="https://example.com/resource",
user_api_key_auth=user_api_key_auth,
)
mock_allowed.assert_awaited_once()
assert exc_info.value.status_code == 400
assert "Multiple MCP servers" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
"""Test that _get_tools_from_mcp_servers continues when one server fails"""
@@ -17,6 +17,8 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
from mcp import ReadResourceResult, Resource
from mcp.types import GetPromptResult, Prompt, ResourceTemplate, TextResourceContents
class TestMCPServerManager:
@@ -39,7 +41,7 @@ class TestMCPServerManager:
result = _deserialize_json_dict(invalid_json)
assert result is None
def test_add_update_server_stdio(self):
async def test_add_update_server_stdio(self):
"""Test adding stdio MCP server"""
manager = MCPServerManager()
@@ -221,6 +223,195 @@ class TestMCPServerManager:
assert len(result) == 1
assert result[0].name == "github_tool_1"
@pytest.mark.asyncio
async def test_get_prompts_from_server_success(self):
"""Ensure prompts are fetched and prefixed when requested."""
manager = MCPServerManager()
server = MCPServer(
server_id="server-1",
name="alias-server",
alias="alias-server",
server_name="alias-server",
url="https://example.com",
transport=MCPTransport.http,
)
mock_prompt = Prompt(name="hello", description="Say hi")
mock_client = AsyncMock()
mock_client.list_prompts = AsyncMock(return_value=[mock_prompt])
with patch.object(manager, "_create_mcp_client", return_value=mock_client):
prompts = await manager.get_prompts_from_server(server, add_prefix=True)
mock_client.list_prompts.assert_awaited_once()
assert len(prompts) == 1
assert prompts[0].name == "alias-server-hello"
@pytest.mark.asyncio
async def test_get_prompt_from_server_success(self):
"""Ensure a single prompt definition is requested via the MCP client."""
manager = MCPServerManager()
server = MCPServer(
server_id="server-1",
name="alias-server",
alias="alias-server",
server_name="alias-server",
url="https://example.com",
transport=MCPTransport.http,
)
mock_result = GetPromptResult(
description="Hello world prompt",
messages=[],
)
mock_client = AsyncMock()
mock_client.get_prompt = AsyncMock(return_value=mock_result)
with patch.object(manager, "_create_mcp_client", return_value=mock_client):
result = await manager.get_prompt_from_server(
server=server,
prompt_name="hello",
arguments={"tone": "casual"},
)
mock_client.get_prompt.assert_awaited_once()
awaited_call = mock_client.get_prompt.await_args
called_params = awaited_call.args[0]
assert called_params.name == "hello"
assert called_params.arguments == {"tone": "casual"}
assert result is mock_result
@pytest.mark.asyncio
async def test_get_resources_from_server_success(self):
manager = MCPServerManager()
server = MCPServer(
server_id="server-1",
name="alias-server",
alias="alias-server",
server_name="alias-server",
url="https://example.com",
transport=MCPTransport.http,
static_headers={"X-Static": "static"},
)
mock_client = AsyncMock()
mock_resources = [Resource(name="file", uri="https://example.com/file")]
mock_client.list_resources = AsyncMock(return_value=mock_resources)
prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")]
with patch.object(manager, "_create_mcp_client", return_value=mock_client) as mock_create_client, patch.object(
manager,
"_create_prefixed_resources",
return_value=prefixed_resources,
) as mock_prefix:
result = await manager.get_resources_from_server(
server=server,
mcp_auth_header="auth",
extra_headers={"X-Test": "1"},
add_prefix=True,
)
mock_create_client.assert_called_once()
called_kwargs = mock_create_client.call_args.kwargs
assert called_kwargs["server"] is server
assert called_kwargs["mcp_auth_header"] == "auth"
assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "static"}
mock_client.list_resources.assert_awaited_once()
mock_prefix.assert_called_once_with(mock_resources, server, add_prefix=True)
assert result == prefixed_resources
@pytest.mark.asyncio
async def test_get_resource_templates_from_server_success(self):
manager = MCPServerManager()
server = MCPServer(
server_id="server-1",
name="alias-server",
alias="alias-server",
server_name="alias-server",
url="https://example.com",
transport=MCPTransport.http,
)
mock_client = AsyncMock()
mock_templates = [
ResourceTemplate(
name="template",
uriTemplate="https://example.com/{id}",
)
]
mock_client.list_resource_templates = AsyncMock(return_value=mock_templates)
prefixed_templates = [
ResourceTemplate(
name="alias-server-template",
uriTemplate="https://example.com/{id}",
)
]
with patch.object(manager, "_create_mcp_client", return_value=mock_client) as mock_create_client, patch.object(
manager,
"_create_prefixed_resource_templates",
return_value=prefixed_templates,
) as mock_prefix:
result = await manager.get_resource_templates_from_server(
server=server,
mcp_auth_header="auth",
extra_headers=None,
add_prefix=False,
)
mock_create_client.assert_called_once_with(
server=server,
mcp_auth_header="auth",
extra_headers=None,
)
mock_client.list_resource_templates.assert_awaited_once()
mock_prefix.assert_called_once_with(mock_templates, server, add_prefix=False)
assert result == prefixed_templates
@pytest.mark.asyncio
async def test_read_resource_from_server_success(self):
manager = MCPServerManager()
server = MCPServer(
server_id="server-1",
name="alias-server",
alias="alias-server",
server_name="alias-server",
url="https://example.com",
transport=MCPTransport.http,
static_headers={"X-Static": "1"},
)
mock_client = AsyncMock()
read_result = ReadResourceResult(
contents=[
TextResourceContents(
uri="https://example.com/resource",
text="hello",
mimeType="text/plain",
)
]
)
mock_client.read_resource = AsyncMock(return_value=read_result)
with patch.object(manager, "_create_mcp_client", return_value=mock_client) as mock_create_client:
result = await manager.read_resource_from_server(
server=server,
url="https://example.com/resource",
mcp_auth_header="auth",
extra_headers={"X-Test": "1"},
)
mock_create_client.assert_called_once()
called_kwargs = mock_create_client.call_args.kwargs
assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"}
mock_client.read_resource.assert_awaited_once_with("https://example.com/resource")
assert result is read_result
@pytest.mark.asyncio
async def test_fetch_oauth_metadata_from_resource_returns_servers_and_scopes(self):
manager = MCPServerManager()
@@ -1044,7 +1235,7 @@ class TestMCPServerManager:
assert "tool_1" in tool_names
assert "tool_2" in tool_names
def test_add_db_mcp_server_to_registry(self):
async def test_add_db_mcp_server_to_registry(self):
"""Test that add_db_mcp_server_to_registry adds a MCP server to the registry"""
manager = MCPServerManager()
server = LiteLLM_MCPServerTable(