From 51b782481982ad33c95e60c24de905a62972a901 Mon Sep 17 00:00:00 2001 From: hfdy2019 <5463852+hfdy2019@user.noreply.gitee.com> Date: Wed, 14 Jan 2026 00:47:26 +0800 Subject: [PATCH] feat: make lint & make format --- server/routers/__init__.py | 1 - server/routers/mcp_router.py | 60 ++++++++++++++++++------------------ server/utils/lifespan.py | 3 +- src/agents/common/mcp.py | 32 ++++++++++--------- 4 files changed, 49 insertions(+), 47 deletions(-) diff --git a/server/routers/__init__.py b/server/routers/__init__.py index aaae02fa..706a92eb 100644 --- a/server/routers/__init__.py +++ b/server/routers/__init__.py @@ -24,4 +24,3 @@ router.include_router(mindmap) # /api/mindmap/* router.include_router(graph) # /api/graph/* router.include_router(tasks) # /api/tasks/* router.include_router(mcp) # /api/system/mcp-servers/* - diff --git a/server/routers/mcp_router.py b/server/routers/mcp_router.py index b049318b..6214365b 100644 --- a/server/routers/mcp_router.py +++ b/server/routers/mcp_router.py @@ -54,14 +54,14 @@ async def create_mcp_server( # 校验传输类型 if transport not in ("sse", "streamable_http"): raise HTTPException(status_code=400, detail="传输类型必须是 sse 或 streamable_http") - + 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, @@ -79,10 +79,10 @@ async def create_mcp_server( 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 @@ -129,13 +129,13 @@ async def update_mcp_server( # 校验传输类型 if transport is not None and transport not in ("sse", "streamable_http"): raise HTTPException(status_code=400, detail="传输类型必须是 sse 或 streamable_http") - + 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 @@ -153,15 +153,15 @@ async def update_mcp_server( 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()) - + return {"success": True, "data": server.to_dict()} except HTTPException: raise @@ -182,13 +182,13 @@ async def delete_mcp_server( server = result.scalar_one_or_none() if not server: 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 @@ -214,10 +214,10 @@ async def test_mcp_server( server = result.scalar_one_or_none() if not server: raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在") - + # 获取配置用于测试 config = server.to_mcp_config() - + try: tools = await get_mcp_tools(name, {name: config}) return { @@ -249,19 +249,19 @@ async def toggle_mcp_server( 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) - + return { "success": True, "enabled": is_enabled, @@ -291,19 +291,19 @@ async def get_mcp_server_tools( server = result.scalar_one_or_none() if not server: raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在") - + # 获取配置 config = server.to_mcp_config() disabled_tools = server.disabled_tools or [] - + try: tools = await get_mcp_tools(name, {name: config}) tool_list = [] - + for tool in tools: original_name = tool.name unique_id = tool.metadata.get("id") if tool.metadata else original_name - + tool_info = { "name": original_name, "id": unique_id, @@ -319,7 +319,7 @@ async def get_mcp_server_tools( tool_info["parameters"] = {} tool_info["required"] = [] tool_list.append(tool_info) - + return { "success": True, "data": tool_list, @@ -352,13 +352,13 @@ async def refresh_mcp_server_tools( 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() - + try: tools = await get_mcp_tools(name, {name: config}) return { @@ -391,9 +391,9 @@ async def toggle_mcp_server_tool( 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) @@ -402,11 +402,11 @@ async def toggle_mcp_server_tool( # 当前启用,改为禁用 disabled_tools.append(tool_name) enabled = False - + server.disabled_tools = disabled_tools server.updated_by = current_user.username await db.commit() - + return { "success": True, "tool_name": tool_name, diff --git a/server/utils/lifespan.py b/server/utils/lifespan.py index 65e9e406..050d1e3b 100644 --- a/server/utils/lifespan.py +++ b/server/utils/lifespan.py @@ -11,8 +11,7 @@ async def lifespan(app: FastAPI): """FastAPI lifespan事件管理器""" # 初始化 MCP 服务器配置 await init_mcp_servers() - + await tasker.start() yield await tasker.shutdown() - diff --git a/src/agents/common/mcp.py b/src/agents/common/mcp.py index 00d0c471..390d5341 100644 --- a/src/agents/common/mcp.py +++ b/src/agents/common/mcp.py @@ -25,15 +25,17 @@ _DEFAULT_MCP_SERVERS = { }, } + 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)) @@ -48,41 +50,42 @@ async def load_mcp_servers_from_db() -> None: 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...") @@ -104,10 +107,10 @@ async def init_mcp_servers() -> None: 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()}") @@ -128,8 +131,9 @@ async def get_mcp_client( def to_camel_case(s: str) -> str: """将字符串转换为小驼峰格式""" import re + # 处理 - 和 _ - s = re.sub(r'[-_]+(.)', lambda m: m.group(1).upper(), s) + s = re.sub(r"[-_]+(.)", lambda m: m.group(1).upper(), s) # 首字母小写 if len(s) > 0: s = s[0].lower() + s[1:] @@ -159,18 +163,18 @@ async def get_mcp_tools(server_name: str, additional_servers: dict[str, dict] = # 渲染 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