mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
add MCP prompts REST endpoints (list + get), MCP-compliant
This commit is contained in:
parent
0f01802dde
commit
36784f7f41
1 changed files with 274 additions and 0 deletions
|
|
@ -491,6 +491,280 @@ if MCP_AVAILABLE:
|
|||
},
|
||||
)
|
||||
|
||||
@router.get("/prompts/list", dependencies=[Depends(user_api_key_auth)])
|
||||
async def list_prompts_rest_api(
|
||||
request: Request,
|
||||
server_id: Optional[str] = Query(
|
||||
None, description="The server id to list prompts for"
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> dict:
|
||||
"""
|
||||
List all available MCP prompts from configured upstream servers.
|
||||
|
||||
MCP-compliant: response shape matches prompts/list (prompts, nextCursor)
|
||||
per https://modelcontextprotocol.io/specification/2025-06-18/server/prompts
|
||||
|
||||
Example response:
|
||||
{
|
||||
"prompts": [
|
||||
{
|
||||
"name": "code_review",
|
||||
"title": "Request Code Review",
|
||||
"description": "Asks the LLM to analyze code quality",
|
||||
"arguments": [
|
||||
{"name": "code", "description": "The code to review", "required": true}
|
||||
]
|
||||
}
|
||||
],
|
||||
"nextCursor": null,
|
||||
"error": null,
|
||||
"message": "Successfully retrieved prompts"
|
||||
}
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
|
||||
try:
|
||||
mcp_auth_header, mcp_server_auth_headers, raw_headers_from_request = (
|
||||
_extract_mcp_headers_from_request(request, MCPRequestHandler)
|
||||
)
|
||||
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
_rest_client_ip = IPAddressUtils.get_mcp_client_ip(request)
|
||||
|
||||
allowed_server_ids_set: set = set()
|
||||
for auth_context in auth_contexts:
|
||||
servers = await global_mcp_server_manager.get_allowed_mcp_servers(
|
||||
user_api_key_auth=auth_context,
|
||||
)
|
||||
allowed_server_ids_set.update(servers)
|
||||
|
||||
allowed_server_ids = global_mcp_server_manager.filter_server_ids_by_ip(
|
||||
list(allowed_server_ids_set), _rest_client_ip
|
||||
)
|
||||
|
||||
prompts_result: list = []
|
||||
error_message: Optional[str] = None
|
||||
|
||||
if server_id:
|
||||
if server_id not in allowed_server_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "access_denied",
|
||||
"message": f"The key is not allowed to access server {server_id}",
|
||||
},
|
||||
)
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
|
||||
if server is None:
|
||||
return {
|
||||
"prompts": [],
|
||||
"nextCursor": None,
|
||||
"error": "server_not_found",
|
||||
"message": f"Server with id {server_id} not found",
|
||||
}
|
||||
|
||||
server_auth_header = _get_server_auth_header(
|
||||
server, mcp_server_auth_headers, mcp_auth_header
|
||||
)
|
||||
|
||||
try:
|
||||
prompts = await global_mcp_server_manager.get_prompts_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
add_prefix=False,
|
||||
raw_headers=raw_headers_from_request,
|
||||
)
|
||||
prompts_result = [p.model_dump() for p in prompts]
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"Error getting prompts from %s: %s", server.name, e
|
||||
)
|
||||
return {
|
||||
"prompts": [],
|
||||
"nextCursor": None,
|
||||
"error": "server_error",
|
||||
"message": f"Failed to get prompts from server {server.name}: {str(e)}",
|
||||
}
|
||||
else:
|
||||
if not allowed_server_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "access_denied",
|
||||
"message": "The key is not allowed to access any MCP servers.",
|
||||
},
|
||||
)
|
||||
|
||||
errors: list = []
|
||||
for allowed_server_id in allowed_server_ids:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(
|
||||
allowed_server_id
|
||||
)
|
||||
if server is None:
|
||||
continue
|
||||
|
||||
server_auth_header = _get_server_auth_header(
|
||||
server, mcp_server_auth_headers, mcp_auth_header
|
||||
)
|
||||
|
||||
try:
|
||||
prompts = await global_mcp_server_manager.get_prompts_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
add_prefix=False,
|
||||
raw_headers=raw_headers_from_request,
|
||||
)
|
||||
prompts_result.extend(p.model_dump() for p in prompts)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"Error getting prompts from %s: %s", server.name, e
|
||||
)
|
||||
errors.append(f"{server.name}: {str(e)}")
|
||||
continue
|
||||
|
||||
if errors and not prompts_result:
|
||||
error_message = (
|
||||
"Failed to get prompts from servers: " + "; ".join(errors)
|
||||
)
|
||||
|
||||
return {
|
||||
"prompts": prompts_result,
|
||||
"nextCursor": None,
|
||||
"error": "partial_failure" if error_message else None,
|
||||
"message": (
|
||||
error_message
|
||||
if error_message
|
||||
else "Successfully retrieved prompts"
|
||||
),
|
||||
}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"Unexpected error in list_prompts_rest_api: %s", str(e)
|
||||
)
|
||||
return {
|
||||
"prompts": [],
|
||||
"nextCursor": None,
|
||||
"error": "unexpected_error",
|
||||
"message": f"An unexpected error occurred: {str(e)}",
|
||||
}
|
||||
|
||||
@router.post("/prompts/get", dependencies=[Depends(user_api_key_auth)])
|
||||
async def get_prompt_rest_api(
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> dict:
|
||||
"""
|
||||
Get a specific MCP prompt by name with optional arguments.
|
||||
|
||||
MCP-compliant: request (name, arguments) and response (description, messages)
|
||||
match prompts/get per https://modelcontextprotocol.io/specification/2025-06-18/server/prompts
|
||||
|
||||
Request body:
|
||||
{
|
||||
"server_id": "server-uuid",
|
||||
"name": "code_review",
|
||||
"arguments": {"code": "def hello(): ..."}
|
||||
}
|
||||
|
||||
Example response:
|
||||
{
|
||||
"description": "Code review prompt",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": {"type": "text", "text": "Please review ..."}
|
||||
}
|
||||
],
|
||||
"error": null,
|
||||
"message": "Successfully retrieved prompt"
|
||||
}
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
|
||||
try:
|
||||
data = await request.json()
|
||||
|
||||
server_id = data.get("server_id")
|
||||
if not server_id:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "missing_parameter",
|
||||
"message": "server_id is required in request body",
|
||||
},
|
||||
)
|
||||
|
||||
prompt_name = data.get("name")
|
||||
if not prompt_name:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "missing_parameter",
|
||||
"message": "name is required in request body",
|
||||
},
|
||||
)
|
||||
|
||||
arguments = data.get("arguments")
|
||||
|
||||
mcp_auth_header, mcp_server_auth_headers, raw_headers_from_request = (
|
||||
_extract_mcp_headers_from_request(request, MCPRequestHandler)
|
||||
)
|
||||
|
||||
allowed_mcp_servers = await _resolve_allowed_mcp_servers_with_ip_filter(
|
||||
request, user_api_key_dict, server_id
|
||||
)
|
||||
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
|
||||
if server is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={
|
||||
"error": "not_found",
|
||||
"message": f"Server with id {server_id} not found",
|
||||
},
|
||||
)
|
||||
|
||||
server_auth_header = _get_server_auth_header(
|
||||
server, mcp_server_auth_headers, mcp_auth_header
|
||||
)
|
||||
|
||||
result = await global_mcp_server_manager.get_prompt_from_server(
|
||||
server=server,
|
||||
prompt_name=prompt_name,
|
||||
arguments=arguments,
|
||||
mcp_auth_header=server_auth_header,
|
||||
raw_headers=raw_headers_from_request,
|
||||
)
|
||||
|
||||
return {
|
||||
"description": result.description,
|
||||
"messages": [m.model_dump() for m in result.messages],
|
||||
"error": None,
|
||||
"message": "Successfully retrieved prompt",
|
||||
}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"Unexpected error in get_prompt_rest_api: %s", str(e)
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "internal_server_error",
|
||||
"message": f"An unexpected error occurred: {str(e)}",
|
||||
},
|
||||
)
|
||||
|
||||
########################################################
|
||||
# MCP Connection testing routes
|
||||
# /health -> Test if we can connect to the MCP server
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue