mirror of
https://github.com/tiennm99/litellm.git
synced 2026-07-21 04:19:24 +00:00
[Feat] mcp resources support (#16800)
* feat: mcp prompts support * feat: mcp resources support
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
+5
-5
@@ -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
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user