mirror of
https://github.com/FuQuan233/nonebot-plugin-llmchat.git
synced 2026-08-13 10:09:27 +00:00
♻️ 大幅重构,拆分模块,增加一些MCP相关限制
This commit is contained in:
parent
41e6aeacb9
commit
0d6771eca6
17 changed files with 1142 additions and 912 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue