ForcePilot/backend/server/routers/mcp_router.py
Wenjie Zhang c772e5ca3a refactor: Agent 运行时架构重构 - 移除 RuntimeConfigMiddleware,统一上下文准备与工具解析
- 删除 runtime_config_middleware.py,将功能合并到:
  - prepare_agent_runtime_context(): 统一上下文准备入口
  - resolve_configured_runtime_tools(): 运行时工具解析
- context.py: 规范化函数重命名 (_names → _keys),重构 config 加载流程
- chatbot/deep_agent graph: 集成新上下文准备流程,移除 RuntimeConfigMiddleware 引用
- skills_middleware: 抽取 normalize_string_list 为共享工具函数
- subagent_service: get_subagents_from_names → get_subagents_from_slugs,支持 slug 查询
- 对应更新 repositories/services/routers 适配新接口
2026-05-26 17:40:17 +08:00

414 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""MCP 服务器管理路由"""
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.services.mcp_service import (
create_mcp_server,
get_mcp_tools_stats,
delete_mcp_server,
get_all_mcp_servers,
get_all_mcp_tools,
get_mcp_server,
set_server_enabled,
toggle_tool_enabled,
update_mcp_server,
)
from yuxi.storage.postgres.models_business import User
from yuxi.utils import logger
from server.utils.auth_middleware import get_admin_user, get_db, get_required_user
mcp = APIRouter(prefix="/system/mcp-servers", tags=["mcp"])
# =============================================================================
# === DTOs ===
# =============================================================================
class CreateMcpServerRequest(BaseModel):
slug: str = Field(..., description="稳定标识")
name: str = Field(..., description="展示名称")
transport: str = Field(..., description="传输类型sse/streamable_http/stdio")
url: str | None = Field(None, description="服务器 URLsse/streamable_http")
command: str | None = Field(None, description="命令stdio")
args: list | None = Field(None, description="命令参数数组stdio")
env: dict | None = Field(None, description="环境变量stdio")
description: str | None = Field(None, description="描述")
headers: dict | None = Field(None, description="HTTP 请求头")
timeout: int | None = Field(None, description="HTTP 超时时间(秒)")
sse_read_timeout: int | None = Field(None, description="SSE 读取超时(秒)")
tags: list | None = Field(None, description="标签数组")
icon: str | None = Field(None, description="图标emoji")
class UpdateMcpServerRequest(BaseModel):
name: str | None = Field(None, description="展示名称")
transport: str | None = Field(None, description="传输类型")
url: str | None = Field(None, description="服务器 URL")
command: str | None = Field(None, description="命令stdio")
args: list | None = Field(None, description="命令参数数组stdio")
env: dict | None = Field(None, description="环境变量stdio")
description: str | None = Field(None, description="描述")
headers: dict | None = Field(None, description="HTTP 请求头")
timeout: int | None = Field(None, description="HTTP 超时时间(秒)")
sse_read_timeout: int | None = Field(None, description="SSE 读取超时(秒)")
tags: list | None = Field(None, description="标签数组")
icon: str | None = Field(None, description="图标emoji")
class UpdateMcpServerStatusRequest(BaseModel):
enabled: bool = Field(..., description="是否启用")
# =============================================================================
# === Helpers ===
# =============================================================================
async def get_server_or_404(db: AsyncSession, slug: str):
"""Helper to get server or raise 404."""
server = await get_mcp_server(db, slug)
if not server:
raise HTTPException(status_code=404, detail=f"服务器 '{slug}' 不存在")
return server
# =============================================================================
# === MCP 服务器 CRUD ===
# =============================================================================
@mcp.get("")
async def get_mcp_servers(
current_user: User = Depends(get_required_user),
db: AsyncSession = Depends(get_db),
):
"""获取所有 MCP 服务器配置(普通用户仅获取脱敏的基础信息)"""
try:
servers = await get_all_mcp_servers(db)
if current_user.role in ["admin", "superadmin"]:
return {"success": True, "data": [s.to_dict() for s in servers]}
else:
# NOTE: 针对普通用户采用高安全显式白名单字段准入投影,使用 getattr 兼容 Mock
# 仿真对象和历史数据,避免未来新增敏感字段或审计信息越权泄露
data = []
for s in servers:
data.append(
{
"name": getattr(s, "name", ""),
"description": getattr(s, "description", None),
"icon": getattr(s, "icon", None),
"enabled": bool(getattr(s, "enabled", True)),
"tags": getattr(s, "tags", None) or [],
}
)
return {"success": True, "data": data}
except Exception as e:
logger.error(f"Failed to get MCP servers: {e}")
raise HTTPException(status_code=500, detail=str(e))
@mcp.post("")
async def create_mcp_server_route(
request: CreateMcpServerRequest,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""创建新的 MCP 服务器"""
# 校验传输类型
valid_transports = ("sse", "streamable_http", "stdio")
if request.transport not in valid_transports:
raise HTTPException(status_code=400, detail=f"传输类型必须是 {', '.join(valid_transports)} 之一")
# 根据传输类型校验必填字段
if request.transport in ("sse", "streamable_http") and not request.url:
raise HTTPException(status_code=400, detail=f"传输类型为 {request.transport}url 必填")
if request.transport == "stdio" and not request.command:
raise HTTPException(status_code=400, detail="传输类型为 stdio 时command 必填")
try:
server = await create_mcp_server(
db,
slug=request.slug,
name=request.name,
transport=request.transport,
url=request.url,
command=request.command,
args=request.args,
env=request.env,
description=request.description,
headers=request.headers,
timeout=request.timeout,
sse_read_timeout=request.sse_read_timeout,
tags=request.tags,
icon=request.icon,
created_by=current_user.username,
)
return {"success": True, "data": server.to_dict()}
except ValueError as ve:
raise HTTPException(status_code=400, detail=str(ve))
except Exception as e:
logger.error(f"Failed to create MCP server: {e}")
raise HTTPException(status_code=500, detail=str(e))
@mcp.get("/{slug}")
async def get_mcp_server_route(
slug: str,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""获取单个 MCP 服务器配置"""
try:
server = await get_server_or_404(db, slug)
return {"success": True, "data": server.to_dict()}
except HTTPException:
raise
except Exception as e:
logger.error(f"Failed to get MCP server: {e}")
raise HTTPException(status_code=500, detail=str(e))
@mcp.put("/{slug}")
async def update_mcp_server_route(
slug: str,
request: UpdateMcpServerRequest,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""更新 MCP 服务器配置"""
# 校验传输类型
valid_transports = ("sse", "streamable_http", "stdio")
if request.transport is not None and request.transport not in valid_transports:
raise HTTPException(status_code=400, detail=f"传输类型必须是 {', '.join(valid_transports)} 之一")
try:
fields_set = getattr(request, "model_fields_set", getattr(request, "__fields_set__", set()))
update_kwargs = {}
if "env" in fields_set:
update_kwargs["env"] = request.env
server = await update_mcp_server(
db,
slug=slug,
name=request.name,
description=request.description,
transport=request.transport,
url=request.url,
command=request.command,
args=request.args,
headers=request.headers,
timeout=request.timeout,
sse_read_timeout=request.sse_read_timeout,
tags=request.tags,
icon=request.icon,
updated_by=current_user.username,
**update_kwargs,
)
return {"success": True, "data": server.to_dict()}
except ValueError as ve:
raise HTTPException(status_code=404, detail=str(ve))
except Exception as e:
logger.error(f"Failed to update MCP server: {e}")
raise HTTPException(status_code=500, detail=str(e))
@mcp.delete("/{slug}")
async def delete_mcp_server_route(
slug: str,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""删除 MCP 服务器"""
try:
# 检查是否为系统内置服务器
server = await get_mcp_server(db, slug)
if server and server.created_by == "system":
raise HTTPException(status_code=403, detail="系统内置的 MCP 服务器无法删除")
deleted = await delete_mcp_server(db, slug)
if not deleted:
raise HTTPException(status_code=404, detail=f"服务器 '{slug}' 不存在")
return {"success": True, "message": f"服务器 '{slug}' 已删除"}
except HTTPException:
raise
except Exception as e:
logger.error(f"Failed to delete MCP server: {e}")
raise HTTPException(status_code=500, detail=str(e))
# =============================================================================
# === MCP 服务器操作 ===
# =============================================================================
@mcp.post("/{slug}/test")
async def test_mcp_server(
slug: str,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""测试 MCP 服务器连接"""
try:
await get_server_or_404(db, slug)
try:
tools = await get_all_mcp_tools(slug)
return {
"success": True,
"message": f"连接成功,共发现 {len(tools)} 个工具",
"tool_count": len(tools),
}
except Exception as test_error:
raise HTTPException(status_code=500, detail=f"连接失败: {str(test_error)}")
except HTTPException:
raise
except Exception as e:
logger.error(f"Failed to test MCP server: {e}")
raise HTTPException(status_code=500, detail=str(e))
@mcp.put("/{slug}/status")
async def update_mcp_server_status_route(
slug: str,
request: UpdateMcpServerStatusRequest,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""更新 MCP 服务器启用状态"""
try:
is_enabled, server = await set_server_enabled(db, slug, request.enabled, current_user.username)
return {
"success": True,
"enabled": is_enabled,
"data": server.to_dict(),
"message": f"MCP '{slug}'{'添加' if is_enabled else '移除'}",
}
except ValueError as ve:
raise HTTPException(status_code=404, detail=str(ve))
except Exception as e:
logger.error(f"Failed to toggle MCP server: {e}")
raise HTTPException(status_code=500, detail=str(e))
# =============================================================================
# === MCP 工具管理 ===
# =============================================================================
@mcp.get("/{slug}/tools")
async def get_mcp_server_tools(
slug: str,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""获取 MCP 服务器的工具列表"""
try:
server = await get_server_or_404(db, slug)
disabled_tools = server.disabled_tools or []
try:
# 获取所有工具(不过滤 disabled_tools
tools = await get_all_mcp_tools(slug)
tool_list = []
for tool in tools:
original_name = tool.name
unique_id = tool.metadata.get("id") if tool.metadata else original_name
tool_info = {
"name": original_name,
"id": unique_id,
"description": getattr(tool, "description", ""),
"enabled": original_name not in disabled_tools,
}
# 提取参数信息
if hasattr(tool, "args_schema") and tool.args_schema:
schema = tool.args_schema.schema() if hasattr(tool.args_schema, "schema") else {}
tool_info["parameters"] = schema.get("properties", {})
tool_info["required"] = schema.get("required", [])
else:
tool_info["parameters"] = {}
tool_info["required"] = []
tool_list.append(tool_info)
return {
"success": True,
"data": tool_list,
"total": len(tool_list),
}
except Exception as tool_error:
logger.error(f"Failed to get tools from MCP server '{slug}': {tool_error}")
raise HTTPException(status_code=500, detail=f"获取工具失败: {str(tool_error)}")
except HTTPException:
raise
except Exception as e:
logger.error(f"Failed to get MCP server tools: {e}")
raise HTTPException(status_code=500, detail=str(e))
@mcp.post("/{slug}/tools/refresh")
async def refresh_mcp_server_tools(
slug: str,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""刷新 MCP 服务器的工具列表(清除缓存重新获取)"""
try:
await get_server_or_404(db, slug)
try:
# 获取所有工具(不过滤 disabled_tools
tools = await get_all_mcp_tools(slug)
# 获取统计信息
stats = get_mcp_tools_stats(slug)
enabled_count = stats.get("enabled", len(tools)) if stats else len(tools)
disabled_count = stats.get("disabled", 0) if stats else 0
message = "工具列表已刷新"
if disabled_count > 0:
message += f"{enabled_count} 个已启用,{disabled_count} 个已禁用"
else:
message += f",共发现 {enabled_count} 个工具"
return {
"success": True,
"message": message,
"tool_count": enabled_count,
"enabled_count": enabled_count,
"disabled_count": disabled_count,
}
except Exception as tool_error:
raise HTTPException(status_code=500, detail=f"刷新失败: {str(tool_error)}")
except HTTPException:
raise
except Exception as e:
logger.error(f"Failed to refresh MCP server tools: {e}")
raise HTTPException(status_code=500, detail=str(e))
@mcp.put("/{slug}/tools/{tool_name}/toggle")
async def toggle_mcp_server_tool_route(
slug: str,
tool_name: str,
current_user: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""切换单个工具的启用状态"""
try:
enabled, server = await toggle_tool_enabled(db, slug, tool_name, current_user.username)
return {
"success": True,
"tool_name": tool_name,
"enabled": enabled,
"message": f"工具 '{tool_name}'{'启用' if enabled else '禁用'}",
}
except ValueError as ve:
raise HTTPException(status_code=404, detail=str(ve))
except Exception as e:
logger.error(f"Failed to toggle MCP server tool: {e}")
raise HTTPException(status_code=500, detail=str(e))