feat: 实现MCP服务以统一业务逻辑和状态管理

- 新增 mcp_service 文件,移除原本的 src/agents/common/mcp.py,所有功能合并到mcp_service,MCP服务来处理服务器配置的增删改查操作、数据库与缓存之间的同步以及MCP客户端和工具的管理。
- 将重构的 MCP 和现有智能体做适配,引入全局缓存和状态管理机制用于MCP服务器和工具。
- 数据库表更新,在MCPServer模型中为StdIO传输类型添加了command和args字段。
- 其他代码优化、增强样式以提升用户体验和可读性。
This commit is contained in:
Wenjie Zhang 2026-01-14 17:32:19 +08:00
parent 51b7824819
commit ebbb831fc0
14 changed files with 1017 additions and 621 deletions

View File

@ -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 查询,获取数据库中的数据。

View File

@ -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>

View File

@ -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="服务器 URLsse/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))

View File

@ -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

View File

@ -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 服务器。"

View File

@ -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

View File

@ -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",
]

View File

@ -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}'")

View File

@ -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

View File

@ -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
View 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)

View File

@ -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="服务器 URLsse/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

View File

@ -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;

View File

@ -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;
}
}