From ebbb831fc08ea2c165a52fbccb74510cecf987ac Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 14 Jan 2026 17:32:19 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=9E=E7=8E=B0MCP=E6=9C=8D=E5=8A=A1?= =?UTF-8?q?=E4=BB=A5=E7=BB=9F=E4=B8=80=E4=B8=9A=E5=8A=A1=E9=80=BB=E8=BE=91?= =?UTF-8?q?=E5=92=8C=E7=8A=B6=E6=80=81=E7=AE=A1=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 mcp_service 文件,移除原本的 src/agents/common/mcp.py,所有功能合并到mcp_service,MCP服务来处理服务器配置的增删改查操作、数据库与缓存之间的同步以及MCP客户端和工具的管理。 - 将重构的 MCP 和现有智能体做适配,引入全局缓存和状态管理机制用于MCP服务器和工具。 - 数据库表更新,在MCPServer模型中为StdIO传输类型添加了command和args字段。 - 其他代码优化、增强样式以提升用户体验和可读性。 --- docs/latest/advanced/agents-config.md | 18 +- docs/latest/intro/quick-start.md | 4 + server/routers/mcp_router.py | 346 +++++----- server/utils/lifespan.py | 8 +- src/agents/chatbot/context.py | 4 +- src/agents/chatbot/graph.py | 6 +- src/agents/common/__init__.py | 10 +- src/agents/common/mcp.py | 217 ------ .../middlewares/dynamic_tool_middleware.py | 2 +- src/agents/reporter/graph.py | 9 +- src/services/mcp_service.py | 615 ++++++++++++++++++ src/storage/db/models.py | 113 ++-- web/src/components/McpServerDetailModal.vue | 73 ++- web/src/components/McpServersComponent.vue | 213 +++--- 14 files changed, 1017 insertions(+), 621 deletions(-) delete mode 100644 src/agents/common/mcp.py create mode 100644 src/services/mcp_service.py diff --git a/docs/latest/advanced/agents-config.md b/docs/latest/advanced/agents-config.md index ccc44177..932470e9 100644 --- a/docs/latest/advanced/agents-config.md +++ b/docs/latest/advanced/agents-config.md @@ -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 查询,获取数据库中的数据。 diff --git a/docs/latest/intro/quick-start.md b/docs/latest/intro/quick-start.md index 8f32647c..eb9a532a 100644 --- a/docs/latest/intro/quick-start.md +++ b/docs/latest/intro/quick-start.md @@ -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 + +新建一个终端重新执行,并确保没有代理干扰。 +
diff --git a/server/routers/mcp_router.py b/server/routers/mcp_router.py index 6214365b..d886ec01 100644 --- a/server/routers/mcp_router.py +++ b/server/routers/mcp_router.py @@ -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)) diff --git a/server/utils/lifespan.py b/server/utils/lifespan.py index 050d1e3b..feb642f8 100644 --- a/server/utils/lifespan.py +++ b/server/utils/lifespan.py @@ -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 diff --git a/src/agents/chatbot/context.py b/src/agents/chatbot/context.py index 5d6efda2..2f772e36 100644 --- a/src/agents/chatbot/context.py +++ b/src/agents/chatbot/context.py @@ -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 服务器。" diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index cc4ebc92..ce6ce771 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -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 diff --git a/src/agents/common/__init__.py b/src/agents/common/__init__.py index e9cf7fef..5a21c7f6 100644 --- a/src/agents/common/__init__.py +++ b/src/agents/common/__init__.py @@ -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", ] diff --git a/src/agents/common/mcp.py b/src/agents/common/mcp.py deleted file mode 100644 index 390d5341..00000000 --- a/src/agents/common/mcp.py +++ /dev/null @@ -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}'") diff --git a/src/agents/common/middlewares/dynamic_tool_middleware.py b/src/agents/common/middlewares/dynamic_tool_middleware.py index a8dc196b..6f11e4c9 100644 --- a/src/agents/common/middlewares/dynamic_tool_middleware.py +++ b/src/agents/common/middlewares/dynamic_tool_middleware.py @@ -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 diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index c8b991a9..e688d378 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -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: diff --git a/src/services/mcp_service.py b/src/services/mcp_service.py new file mode 100644 index 00000000..90651a53 --- /dev/null +++ b/src/services/mcp_service.py @@ -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) diff --git a/src/storage/db/models.py b/src/storage/db/models.py index 77cfb099..6606410a 100644 --- a/src/storage/db/models.py +++ b/src/storage/db/models.py @@ -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 diff --git a/web/src/components/McpServerDetailModal.vue b/web/src/components/McpServerDetailModal.vue index febb4d82..6ad4a1ec 100644 --- a/web/src/components/McpServerDetailModal.vue +++ b/web/src/components/McpServerDetailModal.vue @@ -35,31 +35,52 @@
- + {{ server.transport }}
-
- - {{ server.url }} -
+ + + + + + +
{{ server.description }}
-
- - {{ server.timeout }} 秒 -
-
- - {{ server.sse_read_timeout }} 秒 -
-
- -
{{ JSON.stringify(server.headers, null, 2) }}
-
@@ -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; diff --git a/web/src/components/McpServersComponent.vue b/web/src/components/McpServersComponent.vue index 2f08b5a4..d990838d 100644 --- a/web/src/components/McpServersComponent.vue +++ b/web/src/components/McpServersComponent.vue @@ -18,7 +18,7 @@
已配置 {{ servers.length }} 个 MCP 服务器: - HTTP: {{ httpCount }} · SSE: {{ sseCount }} + HTTP: {{ httpCount }} · SSE: {{ sseCount }} · StdIO: {{ stdioCount }}
@@ -36,9 +36,9 @@
-
@@ -48,29 +48,22 @@

{{ server.name }}

- + {{ server.transport }}
-
-
- {{ server.description }} -
-
- URL: - {{ truncateUrl(server.url) }} -
-
- {{ tag }} +
+ {{ server.description || '暂无描述' }}
@@ -82,10 +75,10 @@ - @@ -99,11 +92,12 @@ 编辑 - + @@ -160,6 +154,7 @@ streamable_http sse + stdio @@ -170,33 +165,55 @@ - - - + + + + + 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; } }