From f2014723c5a4cb2fe1db096b17c5ed3a0073638c Mon Sep 17 00:00:00 2001 From: FuQuan233 Date: Mon, 10 Aug 2026 16:16:47 +0800 Subject: [PATCH] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20=E9=87=8D=E6=9E=84MCPClien?= =?UTF-8?q?t=E4=BC=9A=E8=AF=9D=E7=AE=A1=E7=90=86=EF=BC=8C=E5=A2=9E?= =?UTF-8?q?=E5=8A=A0=E4=BC=9A=E8=AF=9D=E6=89=80=E6=9C=89=E6=9D=83=E6=8E=A7?= =?UTF-8?q?=E5=88=B6=E7=9A=84=E6=B5=8B=E8=AF=95=E7=94=A8=E4=BE=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- nonebot_plugin_llmchat/mcpclient.py | 81 +++++++++++++++++++++-------- tests/test_mcpclient.py | 44 ++++++++++++++++ 2 files changed, 104 insertions(+), 21 deletions(-) create mode 100644 tests/test_mcpclient.py diff --git a/nonebot_plugin_llmchat/mcpclient.py b/nonebot_plugin_llmchat/mcpclient.py index ea7e10f..45f66aa 100644 --- a/nonebot_plugin_llmchat/mcpclient.py +++ b/nonebot_plugin_llmchat/mcpclient.py @@ -1,5 +1,6 @@ import asyncio from contextlib import AsyncExitStack +from dataclasses import dataclass from time import monotonic from typing import Any, cast @@ -14,6 +15,14 @@ from .config import MCPServerConfig from .onebottools import OneBotTools +@dataclass(slots=True) +class _SessionHandle: + """让创建会话的 Task 负责关闭该会话。""" + + stop_event: asyncio.Event + owner_task: asyncio.Task[None] + + class MCPClient: _instance = None _initialized = False @@ -48,7 +57,7 @@ class MCPClient: self.operation_timeout = operation_timeout self.sessions = {} self.exit_stack = AsyncExitStack() - self._session_exit_stacks: dict[str, AsyncExitStack] = {} + self._session_handles: dict[str, _SessionHandle] = {} self._session_last_used: dict[str, float] = {} self._session_lock = asyncio.Lock() self._session_cleanup_task: asyncio.Task | None = None @@ -89,16 +98,34 @@ class MCPClient: await self._get_or_create_session(server_name) logger.info(f"已成功连接到MCP服务器[{server_name}]") - async def _open_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]: + async def _run_server_session( + self, + server_name: str, + ready: asyncio.Future[ClientSession], + stop_event: asyncio.Event, + ) -> None: + """在同一个 Task 中创建并销毁会话。 + + MCP 的传输层使用 AnyIO cancel scope,异步上下文必须由进入它的 + Task 退出,因此不能把 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}]失败或超时") + session, _ = await self._initialize_server_session(server_name, session_stack) + ready.set_result(session) + await stop_event.wait() + except asyncio.CancelledError: + if not ready.done(): + ready.cancel() raise + except BaseException as error: + if not ready.done(): + ready.set_exception(error) + finally: + try: + await session_stack.aclose() + except Exception as error: + logger.opt(exception=error).error(f"关闭MCP会话[{server_name}]失败") async def _initialize_server_session( self, @@ -155,26 +182,38 @@ class MCPClient: await session.initialize() return session, session_stack - async def _create_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]: + async def _create_server_session(self, server_name: str) -> tuple[ClientSession, _SessionHandle]: + loop = asyncio.get_running_loop() + ready: asyncio.Future[ClientSession] = loop.create_future() + stop_event = asyncio.Event() + owner_task = asyncio.create_task( + self._run_server_session(server_name, ready, stop_event), + name=f"llmchat-mcp-{server_name}", + ) try: - return await asyncio.wait_for( - self._open_server_session(server_name), - timeout=self.operation_timeout, - ) - except asyncio.TimeoutError as error: + session = await asyncio.wait_for(asyncio.shield(ready), timeout=self.operation_timeout) + except TimeoutError as error: + owner_task.cancel() + await asyncio.gather(owner_task, return_exceptions=True) raise TimeoutError(f"连接MCP服务器[{server_name}]超时") from error + except BaseException: + owner_task.cancel() + await asyncio.gather(owner_task, return_exceptions=True) + raise + return session, _SessionHandle(stop_event=stop_event, owner_task=owner_task) async def _close_server_session(self, server_name: str): """关闭指定服务器会话。""" - session_stack = self._session_exit_stacks.pop(server_name, None) + handle = self._session_handles.pop(server_name, None) self.sessions.pop(server_name, None) self._session_last_used.pop(server_name, None) - if session_stack is not None: + if handle is not None: + handle.stop_event.set() try: - await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout) - except asyncio.TimeoutError: - logger.error(f"关闭MCP会话[{server_name}]超时") + await asyncio.wait_for(handle.owner_task, timeout=self.operation_timeout + 1) + except TimeoutError: + logger.error(f"等待MCP会话[{server_name}]关闭超时") async def _get_or_create_session(self, server_name: str) -> ClientSession: """获取可复用会话;若不存在或已过期则新建。""" @@ -191,9 +230,9 @@ class MCPClient: session = None if session is None: - session, session_stack = await self._create_server_session(server_name) + session, handle = await self._create_server_session(server_name) self.sessions[server_name] = session - self._session_exit_stacks[server_name] = session_stack + self._session_handles[server_name] = handle self._session_last_used[server_name] = now return self.sessions[server_name] diff --git a/tests/test_mcpclient.py b/tests/test_mcpclient.py new file mode 100644 index 0000000..4413e6d --- /dev/null +++ b/tests/test_mcpclient.py @@ -0,0 +1,44 @@ +import asyncio +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest import IsolatedAsyncioTestCase + +from anyio import CancelScope + +from nonebot_plugin_llmchat.mcpclient import MCPClient + + +class TestMCPClientSessionOwnership(IsolatedAsyncioTestCase): + async def test_session_context_is_closed_by_its_owner_task(self): + entered_by = None + exited_by = None + + @asynccontextmanager + async def task_bound_context(): + nonlocal entered_by, exited_by + entered_by = asyncio.current_task() + with CancelScope(): + try: + yield SimpleNamespace() + finally: + exited_by = asyncio.current_task() + + client = object.__new__(MCPClient) + client.operation_timeout = 1 + + async def initialize(server_name, session_stack): + session = await session_stack.enter_async_context(task_bound_context()) + return session, session_stack + + client._initialize_server_session = initialize + + session, handle = await client._create_server_session("test") + client.sessions = {"test": session} + client._session_handles = {"test": handle} + client._session_last_used = {"test": 0.0} + + await client._close_server_session("test") + + assert entered_by is handle.owner_task + assert exited_by is handle.owner_task + assert handle.owner_task.done()