feat: make lint & make format

This commit is contained in:
hfdy2019 2026-01-14 00:47:26 +08:00 committed by Wenjie Zhang
parent c360252e56
commit 51b7824819
4 changed files with 49 additions and 47 deletions

View File

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

View File

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

View File

@ -11,8 +11,7 @@ async def lifespan(app: FastAPI):
"""FastAPI lifespan事件管理器"""
# 初始化 MCP 服务器配置
await init_mcp_servers()
await tasker.start()
yield
await tasker.shutdown()

View File

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