♻️ 大幅重构,拆分模块,增加一些MCP相关限制
Some checks failed
Pyright Lint / Pyright Lint (push) Has been cancelled
Ruff Lint / Ruff Lint (push) Has been cancelled

This commit is contained in:
FuQuan233 2026-07-29 17:00:22 +08:00
parent 41e6aeacb9
commit 0d6771eca6
17 changed files with 1142 additions and 912 deletions

View file

@ -24,6 +24,7 @@ class MCPClient:
cls,
server_config: dict[str, MCPServerConfig] | None = None,
default_command_cwd: str | None = None,
operation_timeout: int = 30,
):
if cls._instance is None:
cls._instance = super().__new__(cls)
@ -33,6 +34,7 @@ class MCPClient:
self,
server_config: dict[str, MCPServerConfig] | None = None,
default_command_cwd: str | None = None,
operation_timeout: int = 30,
):
if self._initialized:
return
@ -43,6 +45,7 @@ class MCPClient:
logger.info(f"正在初始化MCPClient单例共有{len(server_config)}个服务器配置")
self.server_config = server_config
self.default_command_cwd = default_command_cwd
self.operation_timeout = operation_timeout
self.sessions = {}
self.exit_stack = AsyncExitStack()
self._session_exit_stacks: dict[str, AsyncExitStack] = {}
@ -62,12 +65,13 @@ class MCPClient:
cls,
server_config: dict[str, MCPServerConfig] | None = None,
default_command_cwd: str | None = None,
operation_timeout: int = 30,
):
"""获取MCPClient实例"""
if cls._instance is None:
if server_config is None:
raise ValueError("server_config must be provided for first initialization")
cls._instance = cls(server_config, default_command_cwd)
cls._instance = cls(server_config, default_command_cwd, operation_timeout)
return cls._instance
@classmethod
@ -85,10 +89,24 @@ class MCPClient:
await self._get_or_create_session(server_name)
logger.info(f"已成功连接到MCP服务器[{server_name}]")
async def _create_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]:
async def _open_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]:
session_stack = AsyncExitStack()
try:
return await self._initialize_server_session(server_name, session_stack)
except BaseException:
try:
await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout)
except (asyncio.TimeoutError, RuntimeError):
logger.error(f"清理未完成的MCP会话[{server_name}]失败或超时")
raise
async def _initialize_server_session(
self,
server_name: str,
session_stack: AsyncExitStack,
) -> tuple[ClientSession, AsyncExitStack]:
"""创建并初始化一个新的服务器会话。"""
config = self.server_config[server_name]
session_stack = AsyncExitStack()
if config.url:
transport_type = config.transport
if transport_type == "streamable_http":
@ -137,6 +155,15 @@ class MCPClient:
await session.initialize()
return session, session_stack
async def _create_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]:
try:
return await asyncio.wait_for(
self._open_server_session(server_name),
timeout=self.operation_timeout,
)
except asyncio.TimeoutError as error:
raise TimeoutError(f"连接MCP服务器[{server_name}]超时") from error
async def _close_server_session(self, server_name: str):
"""关闭指定服务器会话。"""
session_stack = self._session_exit_stacks.pop(server_name, None)
@ -144,7 +171,10 @@ class MCPClient:
self._session_last_used.pop(server_name, None)
if session_stack is not None:
await session_stack.aclose()
try:
await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout)
except asyncio.TimeoutError:
logger.error(f"关闭MCP会话[{server_name}]超时")
async def _get_or_create_session(self, server_name: str) -> ClientSession:
"""获取可复用会话;若不存在或已过期则新建。"""
@ -204,7 +234,7 @@ class MCPClient:
for server_name in self.server_config.keys():
logger.debug(f"正在从服务器[{server_name}]获取工具列表")
session = await self._get_or_create_session(server_name)
response = await session.list_tools()
response = await asyncio.wait_for(session.list_tools(), timeout=self.operation_timeout)
tools = response.tools
logger.debug(f"在服务器[{server_name}]中找到{len(tools)}个工具")
@ -243,7 +273,13 @@ class MCPClient:
if group_id is None or bot_id is None:
return "QQ工具需要提供group_id和bot_id参数"
logger.info(f"调用OneBot工具[{tool_name}]")
return await self.onebot_tools.call_tool(tool_name, tool_args, group_id, bot_id)
try:
return await asyncio.wait_for(
self.onebot_tools.call_tool(tool_name, tool_args, group_id, bot_id),
timeout=self.operation_timeout,
)
except asyncio.TimeoutError:
return f"调用OneBot工具[{tool_name}]超时"
# 检查是否是MCP工具
if tool_name.startswith("mcp__"):
@ -254,12 +290,14 @@ class MCPClient:
server_name = parts[1]
real_tool_name = parts[2]
if server_name not in self.server_config:
return f"未知的MCP服务器: {server_name}"
logger.info(f"按需连接到服务器[{server_name}]调用工具[{real_tool_name}]")
try:
await self._ensure_cleanup_task()
session = await self._get_or_create_session(server_name)
response = await asyncio.wait_for(session.call_tool(real_tool_name, tool_args), timeout=30)
response = await asyncio.wait_for(session.call_tool(real_tool_name, tool_args), timeout=self.operation_timeout)
logger.debug(f"工具[{real_tool_name}]调用完成,响应: {response}")
return response.content
except asyncio.TimeoutError:
@ -289,7 +327,8 @@ class MCPClient:
server_name = parts[1]
real_tool_name = parts[2]
return (self.server_config[server_name].friendly_name or server_name) + " - " + real_tool_name
server = self.server_config.get(server_name)
return ((server.friendly_name if server else None) or server_name) + " - " + real_tool_name
# 未知工具类型,返回原名称
return tool_name