feat: make lint & make format
This commit is contained in:
parent
c360252e56
commit
51b7824819
@ -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/*
|
||||
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -11,8 +11,7 @@ async def lifespan(app: FastAPI):
|
||||
"""FastAPI lifespan事件管理器"""
|
||||
# 初始化 MCP 服务器配置
|
||||
await init_mcp_servers()
|
||||
|
||||
|
||||
await tasker.start()
|
||||
yield
|
||||
await tasker.shutdown()
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user