diff --git a/nonebot_plugin_llmchat/mcpclient.py b/nonebot_plugin_llmchat/mcpclient.py index 193427a..ea7e10f 100644 --- a/nonebot_plugin_llmchat/mcpclient.py +++ b/nonebot_plugin_llmchat/mcpclient.py @@ -1,6 +1,5 @@ import asyncio from contextlib import AsyncExitStack -from dataclasses import dataclass from time import monotonic from typing import Any, cast @@ -15,14 +14,6 @@ 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 @@ -43,7 +34,7 @@ class MCPClient: self, server_config: dict[str, MCPServerConfig] | None = None, default_command_cwd: str | None = None, - operation_timeout: int = 100, + operation_timeout: int = 30, ): if self._initialized: return @@ -57,7 +48,7 @@ class MCPClient: self.operation_timeout = operation_timeout self.sessions = {} self.exit_stack = AsyncExitStack() - self._session_handles: dict[str, _SessionHandle] = {} + self._session_exit_stacks: dict[str, AsyncExitStack] = {} self._session_last_used: dict[str, float] = {} self._session_lock = asyncio.Lock() self._session_cleanup_task: asyncio.Task | None = None @@ -98,34 +89,16 @@ class MCPClient: await self._get_or_create_session(server_name) logger.info(f"已成功连接到MCP服务器[{server_name}]") - 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 返回给调用方再关闭。 - """ + async def _open_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]: session_stack = AsyncExitStack() try: - 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: + return await self._initialize_server_session(server_name, session_stack) + except BaseException: try: - await session_stack.aclose() - except Exception as error: - logger.opt(exception=error).error(f"关闭MCP会话[{server_name}]失败") + 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, @@ -182,38 +155,26 @@ class MCPClient: await session.initialize() return session, session_stack - 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}", - ) + async def _create_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]: try: - 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) + 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 - 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): """关闭指定服务器会话。""" - handle = self._session_handles.pop(server_name, None) + session_stack = self._session_exit_stacks.pop(server_name, None) self.sessions.pop(server_name, None) self._session_last_used.pop(server_name, None) - if handle is not None: - handle.stop_event.set() + if session_stack is not None: try: - await asyncio.wait_for(handle.owner_task, timeout=self.operation_timeout + 1) - except TimeoutError: - logger.error(f"等待MCP会话[{server_name}]关闭超时") + 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: """获取可复用会话;若不存在或已过期则新建。""" @@ -230,9 +191,9 @@ class MCPClient: session = None if session is None: - session, handle = await self._create_server_session(server_name) + session, session_stack = await self._create_server_session(server_name) self.sessions[server_name] = session - self._session_handles[server_name] = handle + self._session_exit_stacks[server_name] = session_stack self._session_last_used[server_name] = now return self.sessions[server_name] diff --git a/tests/test_mcpclient.py b/tests/test_mcpclient.py deleted file mode 100644 index 4413e6d..0000000 --- a/tests/test_mcpclient.py +++ /dev/null @@ -1,44 +0,0 @@ -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()