feat: 实现MCP服务以统一业务逻辑和状态管理
- 新增 mcp_service 文件,移除原本的 src/agents/common/mcp.py,所有功能合并到mcp_service,MCP服务来处理服务器配置的增删改查操作、数据库与缓存之间的同步以及MCP客户端和工具的管理。 - 将重构的 MCP 和现有智能体做适配,引入全局缓存和状态管理机制用于MCP服务器和工具。 - 数据库表更新,在MCPServer模型中为StdIO传输类型添加了command和args字段。 - 其他代码优化、增强样式以提升用户体验和可读性。
This commit is contained in:
parent
51b7824819
commit
ebbb831fc0
@ -106,7 +106,7 @@ async def get_graph(self):
|
||||
|
||||
系统会根据配置自动组装工具集合,涵盖知识图谱查询、向量检索生成的动态工具、MySQL 只读查询能力、Tavily 搜索以及所有注册的 MCP 工具。
|
||||
|
||||
工具的启用状态和描述由配置文件或环境变量决定,当依赖缺失时会被中间件自动忽略,从而避免在图中加载不可用能力。MCP Server 的接入方式保持不变,只需在 `src/agents/common/mcp.py` 的 `MCP_SERVERS` 中填入服务地址与 `transport` 类型,如需更多范式可参阅 LangChain 官方文档。
|
||||
工具的启用状态和描述由配置文件或环境变量决定,当依赖缺失时会被中间件自动忽略,从而避免在图中加载不可用能力。MCP Server 的接入方式保持不变,只需在 `src/services/mcp_service.py` 的 `MCP_SERVERS` 中填入服务地址与 `transport` 类型,如需更多范式可参阅 LangChain 官方文档。
|
||||
|
||||
### MCP 服务器配置方式
|
||||
|
||||
@ -198,23 +198,25 @@ MCP_SERVERS = {
|
||||
|
||||
### 动态工具加载
|
||||
|
||||
系统支持动态加载 MCP 工具:
|
||||
系统提供统一的 MCP 服务层 (`src/services/mcp_service.py`) 封装所有 MCP 相关操作。
|
||||
|
||||
#### 智能体获取工具(自动过滤禁用工具)
|
||||
|
||||
智能体应使用 `get_enabled_mcp_tools()` 获取工具,该函数会自动过滤 `disabled_tools` 中的工具:
|
||||
|
||||
```python
|
||||
from src.agents.common.mcp import get_mcp_tools, add_mcp_server
|
||||
from src.services.mcp_service import get_enabled_mcp_tools, add_mcp_server
|
||||
|
||||
# 获取特定服务器的工具
|
||||
tools = await get_mcp_tools("sequentialthinking")
|
||||
# 获取指定服务器的工具(自动过滤 disabled_tools)
|
||||
tools = await get_enabled_mcp_tools("sequentialthinking")
|
||||
|
||||
# 动态添加新的 MCP 服务器
|
||||
add_mcp_server("custom-server", {
|
||||
"url": "https://your-mcp-server.com/mcp",
|
||||
"transport": "streamable_http"
|
||||
})
|
||||
|
||||
# 获取所有 MCP 工具
|
||||
all_tools = await get_all_mcp_tools()
|
||||
```
|
||||
|
||||
### MySQL 数据库
|
||||
|
||||
在 数据库报表助手(SqlReporterAgent) 中,可以通过配置下面环境变量,让 Agent 能够连接到 MySQL 数据库。并通过执行 SQL 查询,获取数据库中的数据。
|
||||
|
||||
@ -170,6 +170,10 @@ $env:HTTPS_PROXY="http://IP:PORT"
|
||||
|
||||
如果已配置代理但构建失败,尝试移除代理后重试。
|
||||
|
||||
如果出现,FetchError: request to https://registry.npmjs.org/npm failed, reason: connect ECONNREFUSED 127.0.0.1:7890
|
||||
|
||||
新建一个终端重新执行,并确保没有代理干扰。
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
|
||||
@ -1,21 +1,72 @@
|
||||
"""MCP 服务器管理路由"""
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
from sqlalchemy import select
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.agents.common.mcp import (
|
||||
clear_mcp_server_tools_cache,
|
||||
get_mcp_tools,
|
||||
sync_mcp_server_to_cache,
|
||||
from src.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,
|
||||
toggle_server_enabled,
|
||||
toggle_tool_enabled,
|
||||
update_mcp_server,
|
||||
)
|
||||
from src.storage.db.models import MCPServer, User
|
||||
from src.storage.db.models import User
|
||||
from src.utils import logger
|
||||
from server.utils.auth_middleware import get_admin_user, get_db
|
||||
|
||||
mcp = APIRouter(prefix="/system/mcp-servers", tags=["mcp"])
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# === DTOs ===
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class CreateMcpServerRequest(BaseModel):
|
||||
name: str = Field(..., description="服务器名称")
|
||||
transport: str = Field(..., description="传输类型:sse/streamable_http/stdio")
|
||||
url: str | None = Field(None, description="服务器 URL(sse/streamable_http)")
|
||||
command: str | None = Field(None, description="命令(stdio)")
|
||||
args: list | 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):
|
||||
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)")
|
||||
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)")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# === Helpers ===
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def get_server_or_404(db: AsyncSession, name: str):
|
||||
"""Helper to get server or raise 404."""
|
||||
server = await get_mcp_server(db, name)
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
return server
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# === MCP 服务器 CRUD ===
|
||||
# =============================================================================
|
||||
@ -28,8 +79,7 @@ async def get_mcp_servers(
|
||||
):
|
||||
"""获取所有 MCP 服务器配置"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer))
|
||||
servers = result.scalars().all()
|
||||
servers = await get_all_mcp_servers(db)
|
||||
return {"success": True, "data": [s.to_dict() for s in servers]}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get MCP servers: {e}")
|
||||
@ -37,72 +87,56 @@ async def get_mcp_servers(
|
||||
|
||||
|
||||
@mcp.post("")
|
||||
async def create_mcp_server(
|
||||
name: str = Body(..., description="服务器名称"),
|
||||
transport: str = Body(..., description="传输类型:sse/streamable_http"),
|
||||
url: str = Body(..., description="服务器 URL"),
|
||||
description: str = Body(None, description="描述"),
|
||||
headers: dict = Body(None, description="HTTP 请求头"),
|
||||
timeout: int = Body(None, description="HTTP 超时时间(秒)"),
|
||||
sse_read_timeout: int = Body(None, description="SSE 读取超时(秒)"),
|
||||
tags: list = Body(None, description="标签数组"),
|
||||
icon: str = Body(None, description="图标(emoji)"),
|
||||
async def create_mcp_server_route(
|
||||
request: CreateMcpServerRequest,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""创建新的 MCP 服务器"""
|
||||
# 校验传输类型
|
||||
if transport not in ("sse", "streamable_http"):
|
||||
raise HTTPException(status_code=400, detail="传输类型必须是 sse 或 streamable_http")
|
||||
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:
|
||||
# 检查名称是否已存在
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
existing = result.scalar_one_or_none()
|
||||
if existing:
|
||||
raise HTTPException(status_code=400, detail=f"服务器名称 '{name}' 已存在")
|
||||
|
||||
server = MCPServer(
|
||||
name=name,
|
||||
description=description,
|
||||
transport=transport,
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
sse_read_timeout=sse_read_timeout,
|
||||
tags=tags,
|
||||
icon=icon,
|
||||
enabled=1,
|
||||
server = await create_mcp_server(
|
||||
db,
|
||||
name=request.name,
|
||||
transport=request.transport,
|
||||
url=request.url,
|
||||
command=request.command,
|
||||
args=request.args,
|
||||
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,
|
||||
updated_by=current_user.username,
|
||||
)
|
||||
db.add(server)
|
||||
await db.commit()
|
||||
await db.refresh(server)
|
||||
|
||||
# 同步到缓存
|
||||
sync_mcp_server_to_cache(name, server.to_mcp_config())
|
||||
|
||||
return {"success": True, "data": server.to_dict()}
|
||||
except HTTPException:
|
||||
raise
|
||||
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("/{name}")
|
||||
async def get_mcp_server(
|
||||
async def get_mcp_server_route(
|
||||
name: str,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取单个 MCP 服务器配置"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
server = await get_server_or_404(db, name)
|
||||
return {"success": True, "data": server.to_dict()}
|
||||
except HTTPException:
|
||||
raise
|
||||
@ -112,83 +146,58 @@ async def get_mcp_server(
|
||||
|
||||
|
||||
@mcp.put("/{name}")
|
||||
async def update_mcp_server(
|
||||
async def update_mcp_server_route(
|
||||
name: str,
|
||||
description: str = Body(None, description="描述"),
|
||||
transport: str = Body(None, description="传输类型"),
|
||||
url: str = Body(None, description="服务器 URL"),
|
||||
headers: dict = Body(None, description="HTTP 请求头"),
|
||||
timeout: int = Body(None, description="HTTP 超时时间(秒)"),
|
||||
sse_read_timeout: int = Body(None, description="SSE 读取超时(秒)"),
|
||||
tags: list = Body(None, description="标签数组"),
|
||||
icon: str = Body(None, description="图标(emoji)"),
|
||||
request: UpdateMcpServerRequest,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""更新 MCP 服务器配置"""
|
||||
# 校验传输类型
|
||||
if transport is not None and transport not in ("sse", "streamable_http"):
|
||||
raise HTTPException(status_code=400, detail="传输类型必须是 sse 或 streamable_http")
|
||||
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:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
|
||||
# 更新字段
|
||||
if description is not None:
|
||||
server.description = description
|
||||
if transport is not None:
|
||||
server.transport = transport
|
||||
if url is not None:
|
||||
server.url = url
|
||||
if headers is not None:
|
||||
server.headers = headers
|
||||
if timeout is not None:
|
||||
server.timeout = timeout
|
||||
if sse_read_timeout is not None:
|
||||
server.sse_read_timeout = sse_read_timeout
|
||||
if tags is not None:
|
||||
server.tags = tags
|
||||
if icon is not None:
|
||||
server.icon = icon
|
||||
|
||||
server.updated_by = current_user.username
|
||||
await db.commit()
|
||||
await db.refresh(server)
|
||||
|
||||
# 同步到缓存(如果启用)
|
||||
if server.enabled:
|
||||
sync_mcp_server_to_cache(name, server.to_mcp_config())
|
||||
|
||||
server = await update_mcp_server(
|
||||
db,
|
||||
name=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,
|
||||
)
|
||||
return {"success": True, "data": server.to_dict()}
|
||||
except HTTPException:
|
||||
raise
|
||||
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("/{name}")
|
||||
async def delete_mcp_server(
|
||||
async def delete_mcp_server_route(
|
||||
name: str,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""删除 MCP 服务器"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
# 检查是否为系统内置服务器
|
||||
server = await get_mcp_server(db, name)
|
||||
if server and server.created_by == "system":
|
||||
raise HTTPException(status_code=403, detail="系统内置的 MCP 服务器无法删除")
|
||||
|
||||
deleted = await delete_mcp_server(db, name)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
|
||||
await db.delete(server)
|
||||
await db.commit()
|
||||
|
||||
# 从缓存中删除
|
||||
sync_mcp_server_to_cache(name, None)
|
||||
|
||||
return {"success": True, "message": f"服务器 '{name}' 已删除"}
|
||||
except HTTPException:
|
||||
raise
|
||||
@ -210,26 +219,17 @@ async def test_mcp_server(
|
||||
):
|
||||
"""测试 MCP 服务器连接"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
|
||||
# 获取配置用于测试
|
||||
config = server.to_mcp_config()
|
||||
await get_server_or_404(db, name)
|
||||
|
||||
try:
|
||||
tools = await get_mcp_tools(name, {name: config})
|
||||
tools = await get_all_mcp_tools(name)
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"连接成功,共发现 {len(tools)} 个工具",
|
||||
"tool_count": len(tools),
|
||||
}
|
||||
except Exception as test_error:
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"连接失败: {str(test_error)}",
|
||||
}
|
||||
raise HTTPException(status_code=500, detail=f"连接失败: {str(test_error)}")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
@ -238,37 +238,21 @@ async def test_mcp_server(
|
||||
|
||||
|
||||
@mcp.put("/{name}/toggle")
|
||||
async def toggle_mcp_server(
|
||||
async def toggle_mcp_server_route(
|
||||
name: str,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""切换 MCP 服务器启用状态"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
|
||||
# 切换状态
|
||||
server.enabled = 0 if server.enabled else 1
|
||||
server.updated_by = current_user.username
|
||||
await db.commit()
|
||||
|
||||
# 获取更新后的状态
|
||||
is_enabled = bool(server.enabled)
|
||||
server_config = server.to_mcp_config() if is_enabled else None
|
||||
|
||||
# 同步到缓存
|
||||
sync_mcp_server_to_cache(name, server_config)
|
||||
|
||||
is_enabled, server = await toggle_server_enabled(db, name, current_user.username)
|
||||
return {
|
||||
"success": True,
|
||||
"enabled": is_enabled,
|
||||
"message": f"服务器 '{name}' 已{'启用' if is_enabled else '禁用'}",
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
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))
|
||||
@ -287,17 +271,12 @@ async def get_mcp_server_tools(
|
||||
):
|
||||
"""获取 MCP 服务器的工具列表"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
|
||||
# 获取配置
|
||||
config = server.to_mcp_config()
|
||||
server = await get_server_or_404(db, name)
|
||||
disabled_tools = server.disabled_tools or []
|
||||
|
||||
try:
|
||||
tools = await get_mcp_tools(name, {name: config})
|
||||
# 获取所有工具(不过滤 disabled_tools)
|
||||
tools = await get_all_mcp_tools(name)
|
||||
tool_list = []
|
||||
|
||||
for tool in tools:
|
||||
@ -327,12 +306,7 @@ async def get_mcp_server_tools(
|
||||
}
|
||||
except Exception as tool_error:
|
||||
logger.error(f"Failed to get tools from MCP server '{name}': {tool_error}")
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"获取工具失败: {str(tool_error)}",
|
||||
"data": [],
|
||||
"total": 0,
|
||||
}
|
||||
raise HTTPException(status_code=500, detail=f"获取工具失败: {str(tool_error)}")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
@ -348,29 +322,32 @@ async def refresh_mcp_server_tools(
|
||||
):
|
||||
"""刷新 MCP 服务器的工具列表(清除缓存重新获取)"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
|
||||
# 清除该服务器的工具缓存
|
||||
clear_mcp_server_tools_cache(name)
|
||||
|
||||
# 获取配置
|
||||
config = server.to_mcp_config()
|
||||
await get_server_or_404(db, name)
|
||||
|
||||
try:
|
||||
tools = await get_mcp_tools(name, {name: config})
|
||||
# 获取所有工具(不过滤 disabled_tools)
|
||||
tools = await get_all_mcp_tools(name)
|
||||
|
||||
# 获取统计信息
|
||||
stats = get_mcp_tools_stats(name)
|
||||
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": f"工具列表已刷新,共发现 {len(tools)} 个工具",
|
||||
"tool_count": len(tools),
|
||||
"message": message,
|
||||
"tool_count": enabled_count,
|
||||
"enabled_count": enabled_count,
|
||||
"disabled_count": disabled_count,
|
||||
}
|
||||
except Exception as tool_error:
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"刷新失败: {str(tool_error)}",
|
||||
}
|
||||
raise HTTPException(status_code=500, detail=f"刷新失败: {str(tool_error)}")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
@ -379,7 +356,7 @@ async def refresh_mcp_server_tools(
|
||||
|
||||
|
||||
@mcp.put("/{name}/tools/{tool_name}/toggle")
|
||||
async def toggle_mcp_server_tool(
|
||||
async def toggle_mcp_server_tool_route(
|
||||
name: str,
|
||||
tool_name: str,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
@ -387,34 +364,15 @@ async def toggle_mcp_server_tool(
|
||||
):
|
||||
"""切换单个工具的启用状态"""
|
||||
try:
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if not server:
|
||||
raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在")
|
||||
|
||||
disabled_tools = list(server.disabled_tools or [])
|
||||
|
||||
if tool_name in disabled_tools:
|
||||
# 当前禁用,改为启用
|
||||
disabled_tools.remove(tool_name)
|
||||
enabled = True
|
||||
else:
|
||||
# 当前启用,改为禁用
|
||||
disabled_tools.append(tool_name)
|
||||
enabled = False
|
||||
|
||||
server.disabled_tools = disabled_tools
|
||||
server.updated_by = current_user.username
|
||||
await db.commit()
|
||||
|
||||
enabled, server = await toggle_tool_enabled(db, name, tool_name, current_user.username)
|
||||
return {
|
||||
"success": True,
|
||||
"tool_name": tool_name,
|
||||
"enabled": enabled,
|
||||
"message": f"工具 '{tool_name}' 已{'启用' if enabled else '禁用'}",
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
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))
|
||||
|
||||
@ -3,14 +3,18 @@ from contextlib import asynccontextmanager
|
||||
from fastapi import FastAPI
|
||||
|
||||
from server.services import tasker
|
||||
from src.agents.common.mcp import init_mcp_servers
|
||||
from src.services.mcp_service import init_mcp_servers
|
||||
from src.utils import logger
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
"""FastAPI lifespan事件管理器"""
|
||||
# 初始化 MCP 服务器配置
|
||||
await init_mcp_servers()
|
||||
try:
|
||||
await init_mcp_servers()
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize MCP servers during startup: {e}")
|
||||
|
||||
await tasker.start()
|
||||
yield
|
||||
|
||||
@ -2,8 +2,8 @@ from dataclasses import dataclass, field
|
||||
from typing import Annotated
|
||||
|
||||
from src.agents.common import BaseContext, gen_tool_info
|
||||
from src.agents.common.mcp import MCP_SERVERS
|
||||
from src.knowledge import knowledge_base
|
||||
from src.services.mcp_service import get_mcp_server_names
|
||||
|
||||
from .tools import get_tools
|
||||
|
||||
@ -33,7 +33,7 @@ class Context(BaseContext):
|
||||
default_factory=list,
|
||||
metadata={
|
||||
"name": "MCP服务器",
|
||||
"options": lambda: list(MCP_SERVERS.keys()),
|
||||
"options": lambda: get_mcp_server_names(),
|
||||
"description": (
|
||||
"MCP服务器列表,建议使用支持 SSE 的 MCP 服务器,"
|
||||
"如果需要使用 uvx 或 npx 运行的服务器,也请在项目外部启动 MCP 服务器,并在项目中配置 MCP 服务器。"
|
||||
|
||||
@ -2,11 +2,11 @@ from langchain.agents import create_agent
|
||||
from langchain.agents.middleware import ModelRetryMiddleware
|
||||
|
||||
from src.agents.common import BaseAgent, load_chat_model
|
||||
from src.agents.common.mcp import get_mcp_tools
|
||||
from src.agents.common.middlewares import (
|
||||
inject_attachment_context,
|
||||
)
|
||||
from src.agents.common.tools import get_kb_based_tools
|
||||
from src.services.mcp_service import get_enabled_mcp_tools
|
||||
|
||||
from .context import Context
|
||||
from .tools import get_tools
|
||||
@ -38,10 +38,10 @@ class ChatbotAgent(BaseAgent):
|
||||
kb_tools = get_kb_based_tools(db_names=knowledges)
|
||||
selected_tools.extend(kb_tools)
|
||||
|
||||
# 3. MCP 工具
|
||||
# 3. MCP 工具(使用统一入口,自动过滤 disabled_tools)
|
||||
if mcps:
|
||||
for server_name in mcps:
|
||||
mcp_tools = await get_mcp_tools(server_name)
|
||||
mcp_tools = await get_enabled_mcp_tools(server_name)
|
||||
selected_tools.extend(mcp_tools)
|
||||
|
||||
return selected_tools
|
||||
|
||||
@ -7,16 +7,13 @@ allowing simplified imports like:
|
||||
|
||||
For other specific functions, use the original import style:
|
||||
from src.agents.common.tools import query_knowledge_graph
|
||||
from src.agents.common.mcp import MCP_SERVERS
|
||||
from src.services.mcp_service import MCP_SERVERS
|
||||
"""
|
||||
|
||||
# Base classes - 核心基类
|
||||
from src.agents.common.base import BaseAgent
|
||||
from src.agents.common.context import BaseContext
|
||||
|
||||
# MCP - 核心 MCP 函数
|
||||
from src.agents.common.mcp import get_mcp_tools
|
||||
|
||||
# Model utilities - 模型加载
|
||||
from src.agents.common.models import load_chat_model
|
||||
from src.agents.common.state import BaseState
|
||||
@ -24,6 +21,9 @@ from src.agents.common.state import BaseState
|
||||
# Tools - 核心工具函数
|
||||
from src.agents.common.tools import gen_tool_info, get_buildin_tools
|
||||
|
||||
# MCP - Agent 层统一入口(自动过滤 disabled_tools)
|
||||
from src.services.mcp_service import get_enabled_mcp_tools
|
||||
|
||||
__all__ = [
|
||||
# Base classes
|
||||
"BaseAgent",
|
||||
@ -35,5 +35,5 @@ __all__ = [
|
||||
"get_buildin_tools",
|
||||
"gen_tool_info",
|
||||
# Core MCP
|
||||
"get_mcp_tools",
|
||||
"get_enabled_mcp_tools",
|
||||
]
|
||||
|
||||
@ -1,217 +0,0 @@
|
||||
"""MCP Client setup and management for LangGraph ReAct Agent."""
|
||||
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from typing import Any, cast
|
||||
|
||||
from langchain_mcp_adapters.client import MultiServerMCPClient
|
||||
|
||||
from src.utils import logger
|
||||
|
||||
# Global MCP tools cache
|
||||
_mcp_tools_cache: dict[str, list[Callable[..., Any]]] = {}
|
||||
|
||||
# MCP Server configurations(运行时缓存,从数据库加载)
|
||||
MCP_SERVERS: dict[str, dict[str, Any]] = {}
|
||||
|
||||
# 默认 MCP 服务器配置(首次启动时导入数据库)
|
||||
_DEFAULT_MCP_SERVERS = {
|
||||
"sequentialthinking": {
|
||||
"url": "https://remote.mcpservers.org/sequentialthinking/mcp",
|
||||
"transport": "streamable_http",
|
||||
"description": "顺序思考工具,帮助 AI 将复杂问题分解为多个步骤",
|
||||
"icon": "🧠",
|
||||
"tags": ["工具", "AI"],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async def load_mcp_servers_from_db() -> None:
|
||||
"""从数据库加载所有启用的 MCP 服务器配置到 MCP_SERVERS 缓存"""
|
||||
global MCP_SERVERS
|
||||
|
||||
# 延迟导入以避免循环引用
|
||||
from sqlalchemy import select
|
||||
|
||||
from src.storage.db.manager import db_manager
|
||||
from src.storage.db.models import MCPServer
|
||||
|
||||
try:
|
||||
async with db_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).filter(MCPServer.enabled == 1))
|
||||
servers = result.scalars().all()
|
||||
MCP_SERVERS.clear()
|
||||
for server in servers:
|
||||
MCP_SERVERS[server.name] = server.to_mcp_config()
|
||||
logger.info(f"Loaded {len(MCP_SERVERS)} MCP servers from database: {list(MCP_SERVERS.keys())}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load MCP servers from database: {e}")
|
||||
|
||||
|
||||
def sync_mcp_server_to_cache(name: str, config: dict[str, Any] | None) -> None:
|
||||
"""同步单个 MCP 服务器配置到缓存
|
||||
|
||||
Args:
|
||||
name: 服务器名称
|
||||
config: 服务器配置,如果为 None 则从缓存中删除
|
||||
"""
|
||||
global MCP_SERVERS
|
||||
|
||||
if config is None:
|
||||
MCP_SERVERS.pop(name, None)
|
||||
logger.info(f"Removed MCP server '{name}' from cache")
|
||||
else:
|
||||
MCP_SERVERS[name] = config
|
||||
logger.info(f"Synced MCP server '{name}' to cache")
|
||||
|
||||
# 清除该服务器的工具缓存
|
||||
_mcp_tools_cache.pop(name, None)
|
||||
|
||||
|
||||
async def init_mcp_servers() -> None:
|
||||
"""初始化 MCP 服务器配置
|
||||
|
||||
首次启动时,如果数据库为空,将默认配置导入数据库
|
||||
然后从数据库加载配置到 MCP_SERVERS 缓存
|
||||
"""
|
||||
# 延迟导入以避免循环引用
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from src.storage.db.manager import db_manager
|
||||
from src.storage.db.models import MCPServer
|
||||
|
||||
try:
|
||||
async with db_manager.get_async_session_context() as session:
|
||||
# 检查数据库是否有 MCP 配置
|
||||
result = await session.execute(select(func.count(MCPServer.name)))
|
||||
count = result.scalar()
|
||||
|
||||
if count == 0:
|
||||
# 数据库为空,导入默认配置
|
||||
logger.info("No MCP servers in database, importing default configurations...")
|
||||
for name, config in _DEFAULT_MCP_SERVERS.items():
|
||||
server = MCPServer(
|
||||
name=name,
|
||||
description=config.get("description"),
|
||||
transport=config["transport"],
|
||||
url=config["url"],
|
||||
headers=config.get("headers"),
|
||||
timeout=config.get("timeout"),
|
||||
sse_read_timeout=config.get("sse_read_timeout"),
|
||||
tags=config.get("tags"),
|
||||
icon=config.get("icon"),
|
||||
enabled=1,
|
||||
created_by="system",
|
||||
updated_by="system",
|
||||
)
|
||||
session.add(server)
|
||||
await session.commit()
|
||||
logger.info(f"Imported {len(_DEFAULT_MCP_SERVERS)} default MCP servers to database")
|
||||
|
||||
# 从数据库加载配置到缓存
|
||||
await load_mcp_servers_from_db()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize MCP servers: {e}, traceback: {traceback.format_exc()}")
|
||||
|
||||
|
||||
async def get_mcp_client(
|
||||
server_configs: dict[str, Any] | None = None,
|
||||
) -> MultiServerMCPClient | None:
|
||||
"""Initializes an MCP client with the given server configurations."""
|
||||
try:
|
||||
client = MultiServerMCPClient(server_configs) # pyright: ignore[reportArgumentType]
|
||||
logger.info(f"Initialized MCP client with servers: {list(server_configs.keys())}")
|
||||
return client
|
||||
except Exception as e:
|
||||
logger.error("Failed to initialize MCP client: {}", e)
|
||||
return None
|
||||
|
||||
|
||||
def to_camel_case(s: str) -> str:
|
||||
"""将字符串转换为小驼峰格式"""
|
||||
import re
|
||||
|
||||
# 处理 - 和 _
|
||||
s = re.sub(r"[-_]+(.)", lambda m: m.group(1).upper(), s)
|
||||
# 首字母小写
|
||||
if len(s) > 0:
|
||||
s = s[0].lower() + s[1:]
|
||||
return s
|
||||
|
||||
|
||||
async def get_mcp_tools(server_name: str, additional_servers: dict[str, dict] = None) -> list[Callable[..., Any]]:
|
||||
"""Get MCP tools for a specific server, initializing client if needed and rendering unique IDs."""
|
||||
global _mcp_tools_cache
|
||||
|
||||
# Return cached tools if available
|
||||
if server_name in _mcp_tools_cache:
|
||||
return _mcp_tools_cache[server_name]
|
||||
|
||||
mcp_servers = MCP_SERVERS | (additional_servers or {})
|
||||
|
||||
try:
|
||||
assert server_name in mcp_servers, f"Server {server_name} not found in ({list(mcp_servers.keys())})"
|
||||
client = await get_mcp_client({server_name: mcp_servers[server_name]})
|
||||
if client is None:
|
||||
return []
|
||||
|
||||
# Get all tools
|
||||
all_tools = await client.get_tools()
|
||||
raw_tools = cast(list[Any], all_tools)
|
||||
|
||||
# 渲染 ID 规则: mcp__[camelCaseServer]__[camelCaseTool]
|
||||
server_cc = to_camel_case(server_name)
|
||||
processed_tools = []
|
||||
|
||||
for tool in raw_tools:
|
||||
# 渲染唯一 ID 规则: mcp__[camelCaseServer]__[camelCaseTool]
|
||||
original_name = tool.name
|
||||
tool_cc = to_camel_case(original_name)
|
||||
unique_id = f"mcp__{server_cc}__{tool_cc}"
|
||||
|
||||
# 使用 metadata 存储,这是 LangChain 工具扩展属性的标准做法
|
||||
if tool.metadata is None:
|
||||
tool.metadata = {}
|
||||
tool.metadata["id"] = unique_id
|
||||
|
||||
processed_tools.append(tool)
|
||||
|
||||
_mcp_tools_cache[server_name] = processed_tools
|
||||
logger.info(f"Loaded {len(processed_tools)} tools from MCP server '{server_name}' with extra tool IDs")
|
||||
return processed_tools
|
||||
except AssertionError as e:
|
||||
logger.warning(f"[assert] Failed to load tools from MCP server '{server_name}': {e}")
|
||||
return []
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load tools from MCP server '{server_name}': {e}, traceback: {traceback.format_exc()}")
|
||||
return []
|
||||
|
||||
|
||||
async def get_all_mcp_tools() -> list[Callable[..., Any]]:
|
||||
"""Get all tools from all configured MCP servers."""
|
||||
all_tools = []
|
||||
for server_name in MCP_SERVERS.keys():
|
||||
tools = await get_mcp_tools(server_name)
|
||||
all_tools.extend(tools)
|
||||
return all_tools
|
||||
|
||||
|
||||
def add_mcp_server(name: str, config: dict[str, Any]) -> None:
|
||||
"""Add a new MCP server configuration."""
|
||||
MCP_SERVERS[name] = config
|
||||
# Clear client to force reinitialization with new config
|
||||
clear_mcp_cache()
|
||||
|
||||
|
||||
def clear_mcp_cache() -> None:
|
||||
"""Clear the MCP tools cache (useful for testing)."""
|
||||
global _mcp_tools_cache
|
||||
_mcp_tools_cache = {}
|
||||
|
||||
|
||||
def clear_mcp_server_tools_cache(server_name: str) -> None:
|
||||
"""Clear the tools cache for a specific MCP server."""
|
||||
global _mcp_tools_cache
|
||||
_mcp_tools_cache.pop(server_name, None)
|
||||
logger.info(f"Cleared tools cache for MCP server '{server_name}'")
|
||||
@ -3,7 +3,7 @@ from typing import Any
|
||||
|
||||
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse
|
||||
|
||||
from src.agents.common import get_mcp_tools
|
||||
from src.services.mcp_service import get_mcp_tools
|
||||
from src.utils import logger
|
||||
|
||||
|
||||
|
||||
@ -1,11 +1,10 @@
|
||||
from langchain.agents import create_agent
|
||||
|
||||
from src.agents.common import BaseAgent, get_mcp_tools, load_chat_model
|
||||
from src.agents.common import BaseAgent, load_chat_model
|
||||
from src.agents.common.toolkits.mysql import get_mysql_tools
|
||||
from src.services.mcp_service import get_mcp_tools
|
||||
from src.utils import logger
|
||||
|
||||
_mcp_servers = {"mcp-server-chart": {"command": "npx", "args": ["-y", "@antv/mcp-server-chart"], "transport": "stdio"}}
|
||||
|
||||
|
||||
class SqlReporterAgent(BaseAgent):
|
||||
name = "数据库报表助手"
|
||||
@ -15,9 +14,9 @@ class SqlReporterAgent(BaseAgent):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
async def get_tools(self):
|
||||
chart_tools = await get_mcp_tools("mcp-server-chart", additional_servers=_mcp_servers)
|
||||
mysql_tools = get_mysql_tools()
|
||||
return chart_tools + mysql_tools
|
||||
chart_tools = await get_mcp_tools("mcp-server-chart")
|
||||
return mysql_tools + chart_tools
|
||||
|
||||
async def get_graph(self, **kwargs):
|
||||
if self.graph:
|
||||
|
||||
615
src/services/mcp_service.py
Normal file
615
src/services/mcp_service.py
Normal file
@ -0,0 +1,615 @@
|
||||
"""MCP Service - Unified business logic and state management for MCP.
|
||||
|
||||
Responsibilities:
|
||||
- Server configuration CRUD operations
|
||||
- Configuration synchronization (Database <-> Cache)
|
||||
- Unified entry point for Agent tool retrieval (auto-filtering disabled_tools)
|
||||
- MCP Client and Tools management (formerly in agents/common/mcp.py)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from typing import Any, cast
|
||||
|
||||
from langchain_mcp_adapters.client import MultiServerMCPClient
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.storage.db.models import MCPServer
|
||||
from src.utils import logger
|
||||
|
||||
# =============================================================================
|
||||
# === Global Cache & State ===
|
||||
# =============================================================================
|
||||
|
||||
# Global Lock for MCP state
|
||||
_mcp_lock = asyncio.Lock()
|
||||
|
||||
# Global MCP tools cache
|
||||
_mcp_tools_cache: dict[str, list[Callable[..., Any]]] = {}
|
||||
|
||||
# MCP tools statistics (for reporting enabled/disabled counts)
|
||||
_mcp_tools_stats: dict[str, dict[str, int]] = {}
|
||||
|
||||
# MCP Server configurations (Runtime cache, loaded from DB)
|
||||
MCP_SERVERS: dict[str, dict[str, Any]] = {}
|
||||
|
||||
# Default MCP Server configurations (Imported to DB on first run)
|
||||
_DEFAULT_MCP_SERVERS = {
|
||||
"sequentialthinking": {
|
||||
"url": "https://remote.mcpservers.org/sequentialthinking/mcp",
|
||||
"transport": "streamable_http",
|
||||
"description": "顺序思考工具,帮助 AI 将复杂问题分解为多个步骤",
|
||||
"icon": "🧠",
|
||||
"tags": ["内置", "AI"],
|
||||
},
|
||||
"mcp-server-chart": {
|
||||
"command": "npx",
|
||||
"args": ["-y", "@antv/mcp-server-chart"],
|
||||
"transport": "stdio",
|
||||
"description": "图表生成工具,支持生成各类图表(柱状图、折线图、饼图等)",
|
||||
"icon": "📊",
|
||||
"tags": ["内置", "图表"],
|
||||
},
|
||||
}
|
||||
|
||||
# =============================================================================
|
||||
# === Core Logic (Moved from agents/common/mcp.py) ===
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def load_mcp_servers_from_db() -> None:
|
||||
"""Load all enabled MCP server configurations from database to MCP_SERVERS cache."""
|
||||
global MCP_SERVERS
|
||||
|
||||
# Delayed import to avoid circular references
|
||||
from src.storage.db.manager import db_manager
|
||||
|
||||
try:
|
||||
async with db_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).filter(MCPServer.enabled == 1))
|
||||
servers = result.scalars().all()
|
||||
|
||||
async with _mcp_lock:
|
||||
MCP_SERVERS.clear()
|
||||
for server in servers:
|
||||
MCP_SERVERS[server.name] = server.to_mcp_config()
|
||||
|
||||
logger.info(f"Loaded {len(MCP_SERVERS)} MCP servers from database: {list(MCP_SERVERS.keys())}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load MCP servers from database: {e}")
|
||||
|
||||
|
||||
async def sync_mcp_server_to_cache(name: str, config: dict[str, Any] | None) -> None:
|
||||
"""Sync a single MCP server configuration to cache.
|
||||
|
||||
Args:
|
||||
name: Server name
|
||||
config: Server configuration, or None to remove from cache
|
||||
"""
|
||||
global MCP_SERVERS
|
||||
|
||||
async with _mcp_lock:
|
||||
if config is None:
|
||||
MCP_SERVERS.pop(name, None)
|
||||
logger.info(f"Removed MCP server '{name}' from cache")
|
||||
else:
|
||||
MCP_SERVERS[name] = config
|
||||
logger.info(f"Synced MCP server '{name}' to cache")
|
||||
|
||||
# Clear tools cache for this server
|
||||
_mcp_tools_cache.pop(name, None)
|
||||
|
||||
|
||||
async def init_mcp_servers() -> None:
|
||||
"""Initialize MCP server configurations.
|
||||
|
||||
On first run, if database is empty, import default configurations.
|
||||
Then load configurations from database to MCP_SERVERS cache.
|
||||
Also ensures all built-in MCP servers are present in the database.
|
||||
"""
|
||||
# Delayed import to avoid circular references
|
||||
from src.storage.db.manager import db_manager
|
||||
|
||||
try:
|
||||
async with db_manager.get_async_session_context() as session:
|
||||
# Check if database has MCP configurations
|
||||
result = await session.execute(select(func.count(MCPServer.name)))
|
||||
count = result.scalar()
|
||||
|
||||
if count == 0:
|
||||
# Database is empty, import default configurations
|
||||
logger.info("No MCP servers in database, importing default configurations...")
|
||||
for name, config in _DEFAULT_MCP_SERVERS.items():
|
||||
server = MCPServer(
|
||||
name=name,
|
||||
description=config.get("description"),
|
||||
transport=config["transport"],
|
||||
url=config.get("url"),
|
||||
command=config.get("command"),
|
||||
args=config.get("args"),
|
||||
headers=config.get("headers"),
|
||||
timeout=config.get("timeout"),
|
||||
sse_read_timeout=config.get("sse_read_timeout"),
|
||||
tags=config.get("tags"),
|
||||
icon=config.get("icon"),
|
||||
enabled=1,
|
||||
created_by="system",
|
||||
updated_by="system",
|
||||
)
|
||||
session.add(server)
|
||||
await session.commit()
|
||||
logger.info(f"Imported {len(_DEFAULT_MCP_SERVERS)} default MCP servers to database")
|
||||
else:
|
||||
# Ensure all built-in MCP servers exist in database
|
||||
for name, config in _DEFAULT_MCP_SERVERS.items():
|
||||
result = await session.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
existing = result.scalar_one_or_none()
|
||||
if not existing:
|
||||
server = MCPServer(
|
||||
name=name,
|
||||
description=config.get("description"),
|
||||
transport=config["transport"],
|
||||
url=config.get("url"),
|
||||
command=config.get("command"),
|
||||
args=config.get("args"),
|
||||
headers=config.get("headers"),
|
||||
timeout=config.get("timeout"),
|
||||
sse_read_timeout=config.get("sse_read_timeout"),
|
||||
tags=config.get("tags"),
|
||||
icon=config.get("icon"),
|
||||
enabled=1,
|
||||
created_by="system",
|
||||
updated_by="system",
|
||||
)
|
||||
session.add(server)
|
||||
logger.info(f"Added built-in MCP server '{name}' to database")
|
||||
# Commit if any new servers were added (check session state)
|
||||
if session.new:
|
||||
await session.commit()
|
||||
|
||||
# Load configurations from database to cache
|
||||
await load_mcp_servers_from_db()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize MCP servers: {e}, traceback: {traceback.format_exc()}")
|
||||
|
||||
|
||||
async def get_mcp_client(
|
||||
server_configs: dict[str, Any] | None = None,
|
||||
) -> MultiServerMCPClient | None:
|
||||
"""Initializes an MCP client with the given server configurations."""
|
||||
try:
|
||||
client = MultiServerMCPClient(server_configs) # pyright: ignore[reportArgumentType]
|
||||
logger.info(f"Initialized MCP client with servers: {list(server_configs.keys())}")
|
||||
return client
|
||||
except Exception as e:
|
||||
logger.error("Failed to initialize MCP client: {}", e)
|
||||
return None
|
||||
|
||||
|
||||
def to_camel_case(s: str) -> str:
|
||||
"""Convert string to lowerCamelCase."""
|
||||
|
||||
# Handle - and _
|
||||
s = re.sub(r"[-_]+(.)", lambda m: m.group(1).upper(), s)
|
||||
# Lowercase first letter
|
||||
if len(s) > 0:
|
||||
s = s[0].lower() + s[1:]
|
||||
return s
|
||||
|
||||
|
||||
async def get_mcp_tools(
|
||||
server_name: str,
|
||||
additional_servers: dict[str, dict] = None,
|
||||
disabled_tools: list[str] = None,
|
||||
cache: bool = True,
|
||||
force_refresh: bool = False,
|
||||
) -> list[Callable[..., Any]]:
|
||||
"""Get MCP tools for a specific server.
|
||||
|
||||
Architecture:
|
||||
1. Fetching: Connects to MCP server to get ALL tools.
|
||||
2. Caching: Stores the FULL, UNFILTERED list of tools in `_mcp_tools_cache`.
|
||||
3. Filtering: Filters the return value based on `disabled_tools` argument.
|
||||
|
||||
Args:
|
||||
server_name: Server name
|
||||
additional_servers: Additional server configurations
|
||||
disabled_tools: List of tool names to filter out from the RETURN value (does not affect cache)
|
||||
cache: Whether to use/update the cache (default: True)
|
||||
force_refresh: Whether to force a refresh from the server (default: False)
|
||||
"""
|
||||
global _mcp_tools_cache
|
||||
|
||||
# 1. Prepare Server Config
|
||||
async with _mcp_lock:
|
||||
mcp_servers = MCP_SERVERS | (additional_servers or {})
|
||||
|
||||
all_processed_tools = []
|
||||
|
||||
# 2. Check Cache / Fetch Strategy
|
||||
# If we have it in cache and don't need to force refresh, use cache.
|
||||
if not force_refresh and cache and server_name in _mcp_tools_cache:
|
||||
all_processed_tools = _mcp_tools_cache[server_name]
|
||||
else:
|
||||
# Need to fetch from server
|
||||
try:
|
||||
assert server_name in mcp_servers, f"Server {server_name} not found in ({list(mcp_servers.keys())})"
|
||||
|
||||
# Extract connection config
|
||||
server_config = mcp_servers[server_name]
|
||||
client_config = {k: v for k, v in server_config.items() if k not in ("disabled_tools",)}
|
||||
|
||||
client = await get_mcp_client({server_name: client_config})
|
||||
if client is None:
|
||||
return []
|
||||
|
||||
# Get ALL tools (Raw)
|
||||
raw_tools = cast(list[Any], await client.get_tools())
|
||||
|
||||
# Render IDs for ALL tools
|
||||
server_cc = to_camel_case(server_name)
|
||||
for tool in raw_tools:
|
||||
# Render unique ID rule: mcp__[camelCaseServer]__[camelCaseTool]
|
||||
original_name = tool.name
|
||||
tool_cc = to_camel_case(original_name)
|
||||
unique_id = f"mcp__{server_cc}__{tool_cc}"
|
||||
|
||||
# Use metadata to store
|
||||
if tool.metadata is None:
|
||||
tool.metadata = {}
|
||||
tool.metadata["id"] = unique_id
|
||||
|
||||
all_processed_tools.append(tool)
|
||||
|
||||
# Update Cache (Store the FULL list)
|
||||
if cache:
|
||||
_mcp_tools_cache[server_name] = all_processed_tools
|
||||
|
||||
# Update Stats
|
||||
# Stats should reflect the GLOBAL configuration state
|
||||
# (How many are disabled in the stored config, not the transient arg)
|
||||
global_config_disabled = mcp_servers.get(server_name, {}).get("disabled_tools") or []
|
||||
enabled_count = len([t for t in all_processed_tools if t.name not in global_config_disabled])
|
||||
|
||||
_mcp_tools_stats[server_name] = {
|
||||
"total": len(all_processed_tools),
|
||||
"enabled": enabled_count,
|
||||
"disabled": len(all_processed_tools) - enabled_count,
|
||||
}
|
||||
|
||||
logger.info(f"Refreshed MCP tools cache for '{server_name}': {len(all_processed_tools)} tools loaded.")
|
||||
|
||||
except AssertionError as e:
|
||||
logger.warning(f"[assert] Failed to load tools from MCP server '{server_name}': {e}")
|
||||
return []
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to load tools from MCP server '{server_name}': {e}, traceback: {traceback.format_exc()}"
|
||||
)
|
||||
return []
|
||||
|
||||
# 3. Filtering (Apply to Return Value Only)
|
||||
if disabled_tools:
|
||||
filtered_tools = [t for t in all_processed_tools if t.name not in disabled_tools]
|
||||
logger.debug(
|
||||
f"Returning {len(filtered_tools)}/{len(all_processed_tools)} tools for '{server_name}' "
|
||||
f"(filtered {len(disabled_tools)} by argument)"
|
||||
)
|
||||
return filtered_tools
|
||||
|
||||
return all_processed_tools
|
||||
|
||||
|
||||
async def get_tools_from_all_servers() -> list[Callable[..., Any]]:
|
||||
"""Get all tools from all configured MCP servers."""
|
||||
all_tools = []
|
||||
for server_name in MCP_SERVERS.keys():
|
||||
tools = await get_mcp_tools(server_name)
|
||||
all_tools.extend(tools)
|
||||
return all_tools
|
||||
|
||||
|
||||
def add_mcp_server(name: str, config: dict[str, Any]) -> None:
|
||||
"""Add a new MCP server configuration."""
|
||||
MCP_SERVERS[name] = config
|
||||
# Clear client to force reinitialization with new config
|
||||
clear_mcp_cache()
|
||||
|
||||
|
||||
def clear_mcp_cache() -> None:
|
||||
"""Clear the MCP tools cache (useful for testing)."""
|
||||
global _mcp_tools_cache, _mcp_tools_stats
|
||||
_mcp_tools_cache = {}
|
||||
_mcp_tools_stats = {}
|
||||
|
||||
|
||||
def clear_mcp_server_tools_cache(server_name: str) -> None:
|
||||
"""Clear the tools cache for a specific MCP server."""
|
||||
global _mcp_tools_cache, _mcp_tools_stats
|
||||
_mcp_tools_cache.pop(server_name, None)
|
||||
_mcp_tools_stats.pop(server_name, None)
|
||||
logger.info(f"Cleared tools cache for MCP server '{server_name}'")
|
||||
|
||||
|
||||
def get_mcp_tools_stats(server_name: str) -> dict[str, int] | None:
|
||||
"""Get tools statistics for a MCP server.
|
||||
|
||||
Returns:
|
||||
dict with 'total', 'enabled', 'disabled' counts, or None if not available
|
||||
"""
|
||||
return _mcp_tools_stats.get(server_name)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# === Server Config CRUD (Existing in mcp_service.py) ===
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def get_mcp_server(db: AsyncSession, name: str) -> MCPServer | None:
|
||||
"""Get single server configuration."""
|
||||
result = await db.execute(select(MCPServer).filter(MCPServer.name == name))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def get_all_mcp_servers(db: AsyncSession) -> list[MCPServer]:
|
||||
"""Get all server configurations."""
|
||||
result = await db.execute(select(MCPServer))
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def create_mcp_server(
|
||||
db: AsyncSession,
|
||||
name: str,
|
||||
transport: str,
|
||||
url: str = None,
|
||||
command: str = None,
|
||||
args: list = None,
|
||||
description: str = None,
|
||||
headers: dict = None,
|
||||
timeout: int = None,
|
||||
sse_read_timeout: int = None,
|
||||
tags: list = None,
|
||||
icon: str = None,
|
||||
created_by: str = None,
|
||||
) -> MCPServer:
|
||||
"""Create server."""
|
||||
# Check if name exists
|
||||
existing = await get_mcp_server(db, name)
|
||||
if existing:
|
||||
raise ValueError(f"Server name '{name}' already exists")
|
||||
|
||||
server = MCPServer(
|
||||
name=name,
|
||||
description=description,
|
||||
transport=transport,
|
||||
url=url,
|
||||
command=command,
|
||||
args=args,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
sse_read_timeout=sse_read_timeout,
|
||||
tags=tags,
|
||||
icon=icon,
|
||||
enabled=1,
|
||||
created_by=created_by,
|
||||
updated_by=created_by,
|
||||
)
|
||||
db.add(server)
|
||||
await db.commit()
|
||||
await db.refresh(server)
|
||||
|
||||
# Sync to cache
|
||||
await sync_mcp_server_to_cache(name, server.to_mcp_config())
|
||||
|
||||
logger.info(f"Created MCP server '{name}'")
|
||||
return server
|
||||
|
||||
|
||||
async def update_mcp_server(
|
||||
db: AsyncSession,
|
||||
name: str,
|
||||
description: str = None,
|
||||
transport: str = None,
|
||||
url: str = None,
|
||||
command: str = None,
|
||||
args: list = None,
|
||||
headers: dict = None,
|
||||
timeout: int = None,
|
||||
sse_read_timeout: int = None,
|
||||
tags: list = None,
|
||||
icon: str = None,
|
||||
updated_by: str = None,
|
||||
) -> MCPServer:
|
||||
"""Update server configuration."""
|
||||
server = await get_mcp_server(db, name)
|
||||
if not server:
|
||||
raise ValueError(f"Server '{name}' does not exist")
|
||||
|
||||
if description is not None:
|
||||
server.description = description
|
||||
if transport is not None:
|
||||
server.transport = transport
|
||||
if url is not None:
|
||||
server.url = url
|
||||
if command is not None:
|
||||
server.command = command
|
||||
if args is not None:
|
||||
server.args = args
|
||||
if headers is not None:
|
||||
server.headers = headers
|
||||
if timeout is not None:
|
||||
server.timeout = timeout
|
||||
if sse_read_timeout is not None:
|
||||
server.sse_read_timeout = sse_read_timeout
|
||||
if tags is not None:
|
||||
server.tags = tags
|
||||
if icon is not None:
|
||||
server.icon = icon
|
||||
if updated_by is not None:
|
||||
server.updated_by = updated_by
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(server)
|
||||
|
||||
# Sync to cache (if enabled)
|
||||
if server.enabled:
|
||||
await sync_mcp_server_to_cache(name, server.to_mcp_config())
|
||||
|
||||
logger.info(f"Updated MCP server '{name}'")
|
||||
return server
|
||||
|
||||
|
||||
async def delete_mcp_server(db: AsyncSession, name: str) -> bool:
|
||||
"""Delete server."""
|
||||
server = await get_mcp_server(db, name)
|
||||
if not server:
|
||||
return False
|
||||
|
||||
await db.delete(server)
|
||||
await db.commit()
|
||||
|
||||
# Remove from cache
|
||||
await sync_mcp_server_to_cache(name, None)
|
||||
|
||||
logger.info(f"Deleted MCP server '{name}'")
|
||||
return True
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# === Tool Management ===
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def toggle_server_enabled(db: AsyncSession, name: str, updated_by: str = None) -> tuple[bool, MCPServer]:
|
||||
"""Toggle server enabled status."""
|
||||
server = await get_mcp_server(db, name)
|
||||
if not server:
|
||||
raise ValueError(f"Server '{name}' does not exist")
|
||||
|
||||
server.enabled = 0 if server.enabled else 1
|
||||
if updated_by is not None:
|
||||
server.updated_by = updated_by
|
||||
await db.commit()
|
||||
|
||||
# Sync to cache
|
||||
is_enabled = bool(server.enabled)
|
||||
server_config = server.to_mcp_config() if is_enabled else None
|
||||
await sync_mcp_server_to_cache(name, server_config)
|
||||
|
||||
logger.info(f"Toggled MCP server '{name}' enabled={is_enabled}")
|
||||
return is_enabled, server
|
||||
|
||||
|
||||
async def toggle_tool_enabled(
|
||||
db: AsyncSession,
|
||||
server_name: str,
|
||||
tool_name: str,
|
||||
updated_by: str = None,
|
||||
) -> tuple[bool, MCPServer]:
|
||||
"""Toggle single tool enabled status.
|
||||
|
||||
Args:
|
||||
db: Database session
|
||||
server_name: Server name
|
||||
tool_name: Tool name
|
||||
updated_by: Updater
|
||||
|
||||
Returns:
|
||||
(enabled, server): Tool enabled status and updated server object
|
||||
"""
|
||||
server = await get_mcp_server(db, server_name)
|
||||
if not server:
|
||||
raise ValueError(f"Server '{server_name}' does not exist")
|
||||
|
||||
disabled_tools = list(server.disabled_tools or [])
|
||||
|
||||
if tool_name in disabled_tools:
|
||||
disabled_tools.remove(tool_name)
|
||||
enabled = True
|
||||
else:
|
||||
disabled_tools.append(tool_name)
|
||||
enabled = False
|
||||
|
||||
server.disabled_tools = disabled_tools
|
||||
if updated_by is not None:
|
||||
server.updated_by = updated_by
|
||||
await db.commit()
|
||||
|
||||
# Clear tool cache (re-filtered on next fetch)
|
||||
clear_mcp_server_tools_cache(server_name)
|
||||
|
||||
logger.info(f"Toggled tool '{tool_name}' for server '{server_name}' enabled={enabled}")
|
||||
return enabled, server
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# === Unified Entry Points (Wrappers) ===
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def get_mcp_server_names() -> list[str]:
|
||||
"""Get list of loaded MCP server names.
|
||||
|
||||
Returns a copy of keys to avoid runtime modification issues during iteration.
|
||||
"""
|
||||
return list(MCP_SERVERS.keys())
|
||||
|
||||
|
||||
async def get_enabled_mcp_tools(server_name: str) -> list:
|
||||
"""Get MCP server tools (auto-filtering disabled_tools).
|
||||
|
||||
Unified entry point for Agents, automatically:
|
||||
1. Gets server config from cache
|
||||
2. Gets all tools
|
||||
3. Filters out disabled_tools
|
||||
|
||||
Args:
|
||||
server_name: Server name
|
||||
|
||||
Returns:
|
||||
List of enabled tools
|
||||
"""
|
||||
config = MCP_SERVERS.get(server_name)
|
||||
if not config:
|
||||
logger.warning(f"MCP server '{server_name}' not found in cache")
|
||||
return []
|
||||
|
||||
disabled_tools = config.get("disabled_tools") or []
|
||||
return await get_mcp_tools(server_name, disabled_tools=disabled_tools)
|
||||
|
||||
|
||||
async def get_servers_config(names: list[str]) -> dict[str, dict[str, Any]]:
|
||||
"""Batch get server configurations.
|
||||
|
||||
Args:
|
||||
names: List of server names
|
||||
|
||||
Returns:
|
||||
{name: config} dictionary, containing only found servers
|
||||
"""
|
||||
return {name: MCP_SERVERS[name] for name in names if name in MCP_SERVERS}
|
||||
|
||||
|
||||
async def get_all_mcp_tools(server_name: str) -> list:
|
||||
"""Get all tools of an MCP server (no filtering).
|
||||
|
||||
For management UI to display tool list, supports viewing all tools and their enabled status.
|
||||
Does NOT update the global tools cache to avoid polluting agent's filtered view.
|
||||
|
||||
Args:
|
||||
server_name: Server name
|
||||
|
||||
Returns:
|
||||
List of all tools (unfiltered)
|
||||
"""
|
||||
config = MCP_SERVERS.get(server_name)
|
||||
if not config:
|
||||
logger.warning(f"MCP server '{server_name}' not found in cache")
|
||||
return []
|
||||
|
||||
# Get all tools (no filtering, force refresh, no cache update)
|
||||
return await get_mcp_tools(server_name, disabled_tools=[], cache=False, force_refresh=True)
|
||||
@ -9,6 +9,15 @@ from src.utils.datetime_utils import coerce_datetime, utc_isoformat, utc_now
|
||||
Base = declarative_base()
|
||||
|
||||
|
||||
def _format_utc_datetime(dt_value):
|
||||
"""Helper to format datetime to UTC ISO string, assuming naive datetimes are UTC."""
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
|
||||
## Removed legacy RDBMS knowledge models (KnowledgeDatabase/KnowledgeFile/KnowledgeNode)
|
||||
|
||||
|
||||
@ -34,13 +43,6 @@ class Conversation(Base):
|
||||
)
|
||||
|
||||
def to_dict(self):
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
return {
|
||||
"id": self.id,
|
||||
"thread_id": self.thread_id,
|
||||
@ -48,8 +50,8 @@ class Conversation(Base):
|
||||
"agent_id": self.agent_id,
|
||||
"title": self.title,
|
||||
"status": self.status,
|
||||
"created_at": format_utc_datetime(self.created_at),
|
||||
"updated_at": format_utc_datetime(self.updated_at),
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"updated_at": _format_utc_datetime(self.updated_at),
|
||||
"metadata": self.extra_metadata or {},
|
||||
}
|
||||
|
||||
@ -76,20 +78,13 @@ class Message(Base):
|
||||
tool_calls = relationship("ToolCall", back_populates="message", cascade="all, delete-orphan")
|
||||
|
||||
def to_dict(self):
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
return {
|
||||
"id": self.id,
|
||||
"conversation_id": self.conversation_id,
|
||||
"role": self.role,
|
||||
"content": self.content,
|
||||
"message_type": self.message_type,
|
||||
"created_at": format_utc_datetime(self.created_at),
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"token_count": self.token_count,
|
||||
"metadata": self.extra_metadata or {},
|
||||
"image_content": self.image_content,
|
||||
@ -124,13 +119,6 @@ class ToolCall(Base):
|
||||
message = relationship("Message", back_populates="tool_calls")
|
||||
|
||||
def to_dict(self):
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
return {
|
||||
"id": self.id,
|
||||
"message_id": self.message_id,
|
||||
@ -140,7 +128,7 @@ class ToolCall(Base):
|
||||
"tool_output": self.tool_output,
|
||||
"status": self.status,
|
||||
"error_message": self.error_message,
|
||||
"created_at": format_utc_datetime(self.created_at),
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
}
|
||||
|
||||
|
||||
@ -164,13 +152,6 @@ class ConversationStats(Base):
|
||||
conversation = relationship("Conversation", back_populates="stats")
|
||||
|
||||
def to_dict(self):
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
return {
|
||||
"id": self.id,
|
||||
"conversation_id": self.conversation_id,
|
||||
@ -178,8 +159,8 @@ class ConversationStats(Base):
|
||||
"total_tokens": self.total_tokens,
|
||||
"model_used": self.model_used,
|
||||
"user_feedback": self.user_feedback or {},
|
||||
"created_at": format_utc_datetime(self.created_at),
|
||||
"updated_at": format_utc_datetime(self.updated_at),
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"updated_at": _format_utc_datetime(self.updated_at),
|
||||
}
|
||||
|
||||
|
||||
@ -212,14 +193,6 @@ class User(Base):
|
||||
|
||||
def to_dict(self, include_password=False):
|
||||
# SQLite 存储 naive datetime,需要标记为 UTC 后再转换
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
# 如果是 naive datetime,假设它是 UTC(因为代码中使用 utc_now() 存储)
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
result = {
|
||||
"id": self.id,
|
||||
"username": self.username,
|
||||
@ -227,13 +200,13 @@ class User(Base):
|
||||
"phone_number": self.phone_number,
|
||||
"avatar": self.avatar,
|
||||
"role": self.role,
|
||||
"created_at": format_utc_datetime(self.created_at),
|
||||
"last_login": format_utc_datetime(self.last_login),
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"last_login": _format_utc_datetime(self.last_login),
|
||||
"login_failed_count": self.login_failed_count,
|
||||
"last_failed_login": format_utc_datetime(self.last_failed_login),
|
||||
"login_locked_until": format_utc_datetime(self.login_locked_until),
|
||||
"last_failed_login": _format_utc_datetime(self.last_failed_login),
|
||||
"login_locked_until": _format_utc_datetime(self.login_locked_until),
|
||||
"is_deleted": self.is_deleted,
|
||||
"deleted_at": format_utc_datetime(self.deleted_at),
|
||||
"deleted_at": _format_utc_datetime(self.deleted_at),
|
||||
}
|
||||
if include_password:
|
||||
result["password_hash"] = self.password_hash
|
||||
@ -298,20 +271,13 @@ class OperationLog(Base):
|
||||
user = relationship("User", back_populates="operation_logs")
|
||||
|
||||
def to_dict(self):
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
return {
|
||||
"id": self.id,
|
||||
"user_id": self.user_id,
|
||||
"operation": self.operation,
|
||||
"details": self.details,
|
||||
"ip_address": self.ip_address,
|
||||
"timestamp": format_utc_datetime(self.timestamp),
|
||||
"timestamp": _format_utc_datetime(self.timestamp),
|
||||
}
|
||||
|
||||
|
||||
@ -333,20 +299,13 @@ class MessageFeedback(Base):
|
||||
message = relationship("Message", backref="feedbacks")
|
||||
|
||||
def to_dict(self):
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
return {
|
||||
"id": self.id,
|
||||
"message_id": self.message_id,
|
||||
"user_id": self.user_id,
|
||||
"rating": self.rating,
|
||||
"reason": self.reason,
|
||||
"created_at": format_utc_datetime(self.created_at),
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
}
|
||||
|
||||
|
||||
@ -360,8 +319,10 @@ class MCPServer(Base):
|
||||
description = Column(String(500), nullable=True, comment="描述")
|
||||
|
||||
# 连接配置
|
||||
transport = Column(String(20), nullable=False, comment="传输类型:sse/streamable_http")
|
||||
url = Column(String(500), nullable=False, comment="服务器 URL")
|
||||
transport = Column(String(20), nullable=False, comment="传输类型:sse/streamable_http/stdio")
|
||||
url = Column(String(500), nullable=True, comment="服务器 URL(sse/streamable_http)")
|
||||
command = Column(String(500), nullable=True, comment="命令(stdio)")
|
||||
args = Column(JSON, nullable=True, comment="命令参数数组(stdio)")
|
||||
headers = Column(JSON, nullable=True, comment="HTTP 请求头")
|
||||
timeout = Column(Integer, nullable=True, comment="HTTP 超时时间(秒)")
|
||||
sse_read_timeout = Column(Integer, nullable=True, comment="SSE 读取超时(秒)")
|
||||
@ -383,18 +344,13 @@ class MCPServer(Base):
|
||||
updated_at = Column(DateTime, default=utc_now, onupdate=utc_now, comment="更新时间")
|
||||
|
||||
def to_dict(self):
|
||||
def format_utc_datetime(dt_value):
|
||||
if dt_value is None:
|
||||
return None
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=dt.UTC)
|
||||
return utc_isoformat(dt_value)
|
||||
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"transport": self.transport,
|
||||
"url": self.url,
|
||||
"command": self.command,
|
||||
"args": self.args or [],
|
||||
"headers": self.headers or {},
|
||||
"timeout": self.timeout,
|
||||
"sse_read_timeout": self.sse_read_timeout,
|
||||
@ -404,20 +360,27 @@ class MCPServer(Base):
|
||||
"disabled_tools": self.disabled_tools or [],
|
||||
"created_by": self.created_by,
|
||||
"updated_by": self.updated_by,
|
||||
"created_at": format_utc_datetime(self.created_at),
|
||||
"updated_at": format_utc_datetime(self.updated_at),
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"updated_at": _format_utc_datetime(self.updated_at),
|
||||
}
|
||||
|
||||
def to_mcp_config(self) -> dict:
|
||||
"""转换为 MCP 配置格式(用于加载到 MCP_SERVERS 缓存)"""
|
||||
config = {
|
||||
"transport": self.transport,
|
||||
"url": self.url,
|
||||
}
|
||||
if self.url:
|
||||
config["url"] = self.url
|
||||
if self.command:
|
||||
config["command"] = self.command
|
||||
if self.args:
|
||||
config["args"] = self.args
|
||||
if self.headers:
|
||||
config["headers"] = self.headers
|
||||
if self.timeout is not None:
|
||||
config["timeout"] = self.timeout
|
||||
if self.sse_read_timeout is not None:
|
||||
config["sse_read_timeout"] = self.sse_read_timeout
|
||||
if self.disabled_tools:
|
||||
config["disabled_tools"] = self.disabled_tools
|
||||
return config
|
||||
|
||||
@ -35,31 +35,52 @@
|
||||
<div class="info-item">
|
||||
<label>传输类型</label>
|
||||
<span>
|
||||
<a-tag :color="server.transport === 'sse' ? 'orange' : 'blue'">
|
||||
<a-tag :color="getTransportColor(server.transport)">
|
||||
{{ server.transport }}
|
||||
</a-tag>
|
||||
</span>
|
||||
</div>
|
||||
<div class="info-item">
|
||||
<label>服务器 URL</label>
|
||||
<span class="url-text">{{ server.url }}</span>
|
||||
</div>
|
||||
|
||||
<!-- HTTP 类型显示 URL -->
|
||||
<template v-if="server.transport === 'streamable_http' || server.transport === 'sse'">
|
||||
<div class="info-item">
|
||||
<label>服务器 URL</label>
|
||||
<span class="url-text">{{ server.url || '-' }}</span>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.headers && Object.keys(server.headers).length > 0">
|
||||
<label>请求头</label>
|
||||
<pre class="headers-pre">{{ JSON.stringify(server.headers, null, 2) }}</pre>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.timeout">
|
||||
<label>HTTP 超时</label>
|
||||
<span>{{ server.timeout }} 秒</span>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.sse_read_timeout">
|
||||
<label>SSE 读取超时</label>
|
||||
<span>{{ server.sse_read_timeout }} 秒</span>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<!-- StdIO 类型显示 command/args -->
|
||||
<template v-if="server.transport === 'stdio'">
|
||||
<div class="info-item">
|
||||
<label>命令</label>
|
||||
<span class="command-text">{{ server.command || '-' }}</span>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.args && server.args.length > 0">
|
||||
<label>参数</label>
|
||||
<span>
|
||||
<a-tag v-for="(arg, index) in server.args" :key="index" size="small">
|
||||
{{ arg }}
|
||||
</a-tag>
|
||||
</span>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<div class="info-item" v-if="server.description">
|
||||
<label>描述</label>
|
||||
<span>{{ server.description }}</span>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.timeout">
|
||||
<label>HTTP 超时</label>
|
||||
<span>{{ server.timeout }} 秒</span>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.sse_read_timeout">
|
||||
<label>SSE 读取超时</label>
|
||||
<span>{{ server.sse_read_timeout }} 秒</span>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.headers && Object.keys(server.headers).length > 0">
|
||||
<label>请求头</label>
|
||||
<pre class="headers-pre">{{ JSON.stringify(server.headers, null, 2) }}</pre>
|
||||
</div>
|
||||
<div class="info-item" v-if="server.tags && server.tags.length > 0">
|
||||
<label>标签</label>
|
||||
<span>
|
||||
@ -344,6 +365,16 @@ const copyToolName = async (name) => {
|
||||
// 格式化时间
|
||||
const formatTime = (timeStr) => formatDateTime(timeStr)
|
||||
|
||||
// 获取传输类型颜色
|
||||
const getTransportColor = (transport) => {
|
||||
const colors = {
|
||||
sse: 'orange',
|
||||
stdio: 'green',
|
||||
streamable_http: 'blue',
|
||||
}
|
||||
return colors[transport] || 'blue'
|
||||
}
|
||||
|
||||
// 关闭弹框
|
||||
const handleClose = () => {
|
||||
emit('update:visible', false)
|
||||
@ -442,6 +473,14 @@ const handleClose = () => {
|
||||
border-radius: 4px;
|
||||
}
|
||||
|
||||
.command-text {
|
||||
font-family: 'Monaco', 'Consolas', monospace;
|
||||
font-size: 13px;
|
||||
background: var(--gray-50);
|
||||
padding: 8px 12px;
|
||||
border-radius: 4px;
|
||||
}
|
||||
|
||||
.headers-pre {
|
||||
font-family: 'Monaco', 'Consolas', monospace;
|
||||
font-size: 12px;
|
||||
|
||||
@ -18,7 +18,7 @@
|
||||
<div class="stats-section" v-if="servers.length > 0">
|
||||
<span class="stats-text">
|
||||
已配置 {{ servers.length }} 个 MCP 服务器:
|
||||
HTTP: {{ httpCount }} · SSE: {{ sseCount }}
|
||||
HTTP: {{ httpCount }} · SSE: {{ sseCount }} · StdIO: {{ stdioCount }}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
@ -36,9 +36,9 @@
|
||||
</a-empty>
|
||||
</div>
|
||||
<div v-else class="server-cards-grid">
|
||||
<div
|
||||
v-for="server in servers"
|
||||
:key="server.name"
|
||||
<div
|
||||
v-for="server in servers"
|
||||
:key="server.name"
|
||||
class="server-card"
|
||||
:class="{ disabled: !server.enabled }"
|
||||
>
|
||||
@ -48,29 +48,22 @@
|
||||
<div class="server-basic-info">
|
||||
<h4 class="server-name">{{ server.name }}</h4>
|
||||
<div class="server-transport">
|
||||
<a-tag :color="getTransportColor(server.transport)" size="small">
|
||||
<a-tag size="small" class="transport-tag">
|
||||
{{ server.transport }}
|
||||
</a-tag>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<a-switch
|
||||
:checked="server.enabled"
|
||||
<a-switch
|
||||
:checked="server.enabled"
|
||||
@change="handleToggleServer(server)"
|
||||
:loading="toggleLoading === server.name"
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="card-content">
|
||||
<div class="server-description" v-if="server.description">
|
||||
{{ server.description }}
|
||||
</div>
|
||||
<div class="server-url">
|
||||
<span class="url-label">URL:</span>
|
||||
<span class="url-value">{{ truncateUrl(server.url) }}</span>
|
||||
</div>
|
||||
<div class="server-tags" v-if="server.tags && server.tags.length > 0">
|
||||
<a-tag v-for="tag in server.tags" :key="tag" size="small">{{ tag }}</a-tag>
|
||||
<div class="server-description">
|
||||
{{ server.description || '暂无描述' }}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@ -82,10 +75,10 @@
|
||||
</a-button>
|
||||
</a-tooltip>
|
||||
<a-tooltip title="测试连接">
|
||||
<a-button
|
||||
type="text"
|
||||
size="small"
|
||||
@click="handleTestServer(server)"
|
||||
<a-button
|
||||
type="text"
|
||||
size="small"
|
||||
@click="handleTestServer(server)"
|
||||
class="action-btn"
|
||||
:loading="testLoading === server.name"
|
||||
>
|
||||
@ -99,11 +92,12 @@
|
||||
<span>编辑</span>
|
||||
</a-button>
|
||||
</a-tooltip>
|
||||
<a-tooltip title="删除服务器">
|
||||
<a-tooltip :title="server.created_by === 'system' ? '内置 MCP 无法删除' : '删除服务器'">
|
||||
<a-button
|
||||
type="text"
|
||||
size="small"
|
||||
danger
|
||||
:disabled="server.created_by === 'system'"
|
||||
@click="confirmDeleteServer(server)"
|
||||
class="action-btn"
|
||||
>
|
||||
@ -160,6 +154,7 @@
|
||||
<a-select v-model:value="form.transport">
|
||||
<a-select-option value="streamable_http">streamable_http</a-select-option>
|
||||
<a-select-option value="sse">sse</a-select-option>
|
||||
<a-select-option value="stdio">stdio</a-select-option>
|
||||
</a-select>
|
||||
</a-form-item>
|
||||
</a-col>
|
||||
@ -170,33 +165,55 @@
|
||||
</a-col>
|
||||
</a-row>
|
||||
|
||||
<a-form-item label="服务器 URL" required class="form-item">
|
||||
<a-input
|
||||
v-model:value="form.url"
|
||||
placeholder="https://example.com/mcp"
|
||||
/>
|
||||
</a-form-item>
|
||||
<!-- HTTP 类型 -->
|
||||
<template v-if="form.transport === 'streamable_http' || form.transport === 'sse'">
|
||||
<a-form-item label="服务器 URL" required class="form-item">
|
||||
<a-input
|
||||
v-model:value="form.url"
|
||||
placeholder="https://example.com/mcp"
|
||||
/>
|
||||
</a-form-item>
|
||||
|
||||
<a-form-item label="HTTP 请求头" class="form-item">
|
||||
<a-textarea
|
||||
v-model:value="form.headersText"
|
||||
placeholder='JSON 格式,如:{"Authorization": "Bearer xxx"}'
|
||||
:rows="3"
|
||||
/>
|
||||
</a-form-item>
|
||||
<a-form-item label="HTTP 请求头" class="form-item">
|
||||
<a-textarea
|
||||
v-model:value="form.headersText"
|
||||
placeholder='JSON 格式,如:{"Authorization": "Bearer xxx"}'
|
||||
:rows="3"
|
||||
/>
|
||||
</a-form-item>
|
||||
|
||||
<a-row :gutter="16">
|
||||
<a-col :span="12">
|
||||
<a-form-item label="HTTP 超时(秒)" class="form-item">
|
||||
<a-input-number v-model:value="form.timeout" :min="1" :max="300" style="width: 100%" />
|
||||
</a-form-item>
|
||||
</a-col>
|
||||
<a-col :span="12">
|
||||
<a-form-item label="SSE 读取超时(秒)" class="form-item">
|
||||
<a-input-number v-model:value="form.sse_read_timeout" :min="1" :max="300" style="width: 100%" />
|
||||
</a-form-item>
|
||||
</a-col>
|
||||
</a-row>
|
||||
<a-row :gutter="16">
|
||||
<a-col :span="12">
|
||||
<a-form-item label="HTTP 超时(秒)" class="form-item">
|
||||
<a-input-number v-model:value="form.timeout" :min="1" :max="300" style="width: 100%" />
|
||||
</a-form-item>
|
||||
</a-col>
|
||||
<a-col :span="12">
|
||||
<a-form-item label="SSE 读取超时(秒)" class="form-item">
|
||||
<a-input-number v-model:value="form.sse_read_timeout" :min="1" :max="300" style="width: 100%" />
|
||||
</a-form-item>
|
||||
</a-col>
|
||||
</a-row>
|
||||
</template>
|
||||
|
||||
<!-- StdIO 类型 -->
|
||||
<template v-if="form.transport === 'stdio'">
|
||||
<a-form-item label="命令" required class="form-item">
|
||||
<a-input
|
||||
v-model:value="form.command"
|
||||
placeholder="例如:npx 或 /path/to/server"
|
||||
/>
|
||||
</a-form-item>
|
||||
|
||||
<a-form-item label="参数" class="form-item">
|
||||
<a-select
|
||||
v-model:value="form.args"
|
||||
mode="tags"
|
||||
placeholder="输入参数后回车添加,如:-m"
|
||||
style="width: 100%"
|
||||
/>
|
||||
</a-form-item>
|
||||
</template>
|
||||
|
||||
<a-form-item label="标签" class="form-item">
|
||||
<a-select
|
||||
@ -271,6 +288,8 @@ const form = reactive({
|
||||
description: '',
|
||||
transport: 'streamable_http',
|
||||
url: '',
|
||||
command: '',
|
||||
args: [],
|
||||
headersText: '',
|
||||
timeout: null,
|
||||
sse_read_timeout: null,
|
||||
@ -285,6 +304,7 @@ const selectedServer = ref(null)
|
||||
// 计算属性
|
||||
const httpCount = computed(() => servers.value.filter(s => s.transport === 'streamable_http').length)
|
||||
const sseCount = computed(() => servers.value.filter(s => s.transport === 'sse').length)
|
||||
const stdioCount = computed(() => servers.value.filter(s => s.transport === 'stdio').length)
|
||||
|
||||
// 获取服务器列表
|
||||
const fetchServers = async () => {
|
||||
@ -314,6 +334,8 @@ const showAddModal = () => {
|
||||
description: '',
|
||||
transport: 'streamable_http',
|
||||
url: '',
|
||||
command: '',
|
||||
args: [],
|
||||
headersText: '',
|
||||
timeout: null,
|
||||
sse_read_timeout: null,
|
||||
@ -332,7 +354,9 @@ const showEditModal = (server) => {
|
||||
name: server.name,
|
||||
description: server.description || '',
|
||||
transport: server.transport,
|
||||
url: server.url,
|
||||
url: server.url || '',
|
||||
command: server.command || '',
|
||||
args: server.args || [],
|
||||
headersText: server.headers ? JSON.stringify(server.headers, null, 2) : '',
|
||||
timeout: server.timeout,
|
||||
sse_read_timeout: server.sse_read_timeout,
|
||||
@ -352,7 +376,7 @@ const showDetailModal = (server) => {
|
||||
const handleFormSubmit = async () => {
|
||||
try {
|
||||
formLoading.value = true
|
||||
|
||||
|
||||
let data
|
||||
if (formMode.value === 'json') {
|
||||
try {
|
||||
@ -372,12 +396,14 @@ const handleFormSubmit = async () => {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
data = {
|
||||
name: form.name,
|
||||
description: form.description || null,
|
||||
transport: form.transport,
|
||||
url: form.url,
|
||||
url: form.url || null,
|
||||
command: form.command || null,
|
||||
args: form.args.length > 0 ? form.args : null,
|
||||
headers,
|
||||
timeout: form.timeout || null,
|
||||
sse_read_timeout: form.sse_read_timeout || null,
|
||||
@ -385,21 +411,31 @@ const handleFormSubmit = async () => {
|
||||
icon: form.icon || null,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// 校验必填字段
|
||||
if (!data.name?.trim()) {
|
||||
notification.error({ message: '服务器名称不能为空' })
|
||||
return
|
||||
}
|
||||
if (!data.url?.trim()) {
|
||||
notification.error({ message: '服务器 URL 不能为空' })
|
||||
return
|
||||
}
|
||||
if (!data.transport) {
|
||||
notification.error({ message: '请选择传输类型' })
|
||||
return
|
||||
}
|
||||
|
||||
// HTTP 类型校验 URL
|
||||
if (['sse', 'streamable_http'].includes(data.transport)) {
|
||||
if (!data.url?.trim()) {
|
||||
notification.error({ message: 'HTTP 类型必须填写服务器 URL' })
|
||||
return
|
||||
}
|
||||
}
|
||||
// StdIO 类型校验 command
|
||||
if (data.transport === 'stdio') {
|
||||
if (!data.command?.trim()) {
|
||||
notification.error({ message: 'StdIO 类型必须填写命令' })
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if (editMode.value) {
|
||||
const result = await mcpApi.updateMcpServer(data.name, data)
|
||||
if (result.success) {
|
||||
@ -417,7 +453,7 @@ const handleFormSubmit = async () => {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
formModalVisible.value = false
|
||||
await fetchServers()
|
||||
} catch (err) {
|
||||
@ -467,6 +503,15 @@ const handleTestServer = async (server) => {
|
||||
|
||||
// 确认删除服务器
|
||||
const confirmDeleteServer = (server) => {
|
||||
// system 创建的服务器不允许删除
|
||||
if (server.created_by === 'system') {
|
||||
notification.warning({
|
||||
message: '无法删除系统服务器',
|
||||
description: '系统内置的 MCP 服务器无法删除,如需停用可切换禁用开关。',
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
Modal.confirm({
|
||||
title: '确认删除服务器',
|
||||
content: `确定要删除服务器 "${server.name}" 吗?此操作不可撤销。`,
|
||||
@ -514,6 +559,8 @@ const parseJsonToForm = () => {
|
||||
description: obj.description || '',
|
||||
transport: obj.transport || 'streamable_http',
|
||||
url: obj.url || '',
|
||||
command: obj.command || '',
|
||||
args: obj.args || [],
|
||||
headersText: obj.headers ? JSON.stringify(obj.headers, null, 2) : '',
|
||||
timeout: obj.timeout || null,
|
||||
sse_read_timeout: obj.sse_read_timeout || null,
|
||||
@ -527,16 +574,6 @@ const parseJsonToForm = () => {
|
||||
}
|
||||
}
|
||||
|
||||
// 辅助函数
|
||||
const getTransportColor = (transport) => {
|
||||
return transport === 'sse' ? 'orange' : 'blue'
|
||||
}
|
||||
|
||||
const truncateUrl = (url) => {
|
||||
if (!url) return '-'
|
||||
return url.length > 40 ? url.substring(0, 40) + '...' : url
|
||||
}
|
||||
|
||||
// 初始化
|
||||
onMounted(() => {
|
||||
fetchServers()
|
||||
@ -568,7 +605,7 @@ onMounted(() => {
|
||||
|
||||
.stats-section {
|
||||
margin-bottom: 16px;
|
||||
|
||||
|
||||
.stats-text {
|
||||
font-size: 13px;
|
||||
color: var(--gray-600);
|
||||
@ -632,40 +669,32 @@ onMounted(() => {
|
||||
font-weight: 600;
|
||||
color: var(--gray-900);
|
||||
}
|
||||
|
||||
.server-transport {
|
||||
.transport-tag {
|
||||
background: var(--gray-100);
|
||||
border: none;
|
||||
color: var(--gray-600);
|
||||
border-radius: 4px;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
.card-content {
|
||||
margin-bottom: 12px;
|
||||
min-height: 44px;
|
||||
|
||||
.server-description {
|
||||
font-size: 13px;
|
||||
color: var(--gray-600);
|
||||
margin-bottom: 8px;
|
||||
line-height: 1.4;
|
||||
}
|
||||
|
||||
.server-url {
|
||||
font-size: 12px;
|
||||
margin-bottom: 8px;
|
||||
|
||||
.url-label {
|
||||
color: var(--gray-500);
|
||||
margin-right: 4px;
|
||||
}
|
||||
|
||||
.url-value {
|
||||
color: var(--gray-700);
|
||||
font-family: 'Monaco', 'Consolas', monospace;
|
||||
word-break: break-all;
|
||||
}
|
||||
}
|
||||
|
||||
.server-tags {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 4px;
|
||||
display: -webkit-box;
|
||||
-webkit-line-clamp: 2;
|
||||
line-clamp: 2;
|
||||
-webkit-box-orient: vertical;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user