♻️ 重构MCPClient会话管理,增加会话所有权控制的测试用例

This commit is contained in:
FuQuan233 2026-08-10 16:16:47 +08:00
parent f091497613
commit f2014723c5
2 changed files with 104 additions and 21 deletions

View file

@ -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]

44
tests/test_mcpclient.py Normal file
View file

@ -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()