Compare commits

..

No commits in common. "3320758cf9eab08fa39438fc79926c69e9d02a48" and "f0914976131705f68971f00a4af3e4f86a21f2e9" have entirely different histories.

2 changed files with 22 additions and 105 deletions

View file

@ -1,6 +1,5 @@
import asyncio import asyncio
from contextlib import AsyncExitStack from contextlib import AsyncExitStack
from dataclasses import dataclass
from time import monotonic from time import monotonic
from typing import Any, cast from typing import Any, cast
@ -15,14 +14,6 @@ from .config import MCPServerConfig
from .onebottools import OneBotTools from .onebottools import OneBotTools
@dataclass(slots=True)
class _SessionHandle:
"""让创建会话的 Task 负责关闭该会话。"""
stop_event: asyncio.Event
owner_task: asyncio.Task[None]
class MCPClient: class MCPClient:
_instance = None _instance = None
_initialized = False _initialized = False
@ -43,7 +34,7 @@ class MCPClient:
self, self,
server_config: dict[str, MCPServerConfig] | None = None, server_config: dict[str, MCPServerConfig] | None = None,
default_command_cwd: str | None = None, default_command_cwd: str | None = None,
operation_timeout: int = 100, operation_timeout: int = 30,
): ):
if self._initialized: if self._initialized:
return return
@ -57,7 +48,7 @@ class MCPClient:
self.operation_timeout = operation_timeout self.operation_timeout = operation_timeout
self.sessions = {} self.sessions = {}
self.exit_stack = AsyncExitStack() 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_last_used: dict[str, float] = {}
self._session_lock = asyncio.Lock() self._session_lock = asyncio.Lock()
self._session_cleanup_task: asyncio.Task | None = None self._session_cleanup_task: asyncio.Task | None = None
@ -98,34 +89,16 @@ class MCPClient:
await self._get_or_create_session(server_name) await self._get_or_create_session(server_name)
logger.info(f"已成功连接到MCP服务器[{server_name}]") logger.info(f"已成功连接到MCP服务器[{server_name}]")
async def _run_server_session( async def _open_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]:
self,
server_name: str,
ready: asyncio.Future[ClientSession],
stop_event: asyncio.Event,
) -> None:
"""在同一个 Task 中创建并销毁会话。
MCP 的传输层使用 AnyIO cancel scope异步上下文必须由进入它的
Task 退出因此不能把 AsyncExitStack 返回给调用方再关闭
"""
session_stack = AsyncExitStack() session_stack = AsyncExitStack()
try: try:
session, _ = await self._initialize_server_session(server_name, session_stack) return await self._initialize_server_session(server_name, session_stack)
ready.set_result(session) except BaseException:
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: try:
await session_stack.aclose() await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout)
except Exception as error: except (asyncio.TimeoutError, RuntimeError):
logger.opt(exception=error).error(f"关闭MCP会话[{server_name}]失败") logger.error(f"清理未完成的MCP会话[{server_name}]失败或超时")
raise
async def _initialize_server_session( async def _initialize_server_session(
self, self,
@ -182,38 +155,26 @@ class MCPClient:
await session.initialize() await session.initialize()
return session, session_stack return session, session_stack
async def _create_server_session(self, server_name: str) -> tuple[ClientSession, _SessionHandle]: async def _create_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]:
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: try:
session = await asyncio.wait_for(asyncio.shield(ready), timeout=self.operation_timeout) return await asyncio.wait_for(
except TimeoutError as error: self._open_server_session(server_name),
owner_task.cancel() timeout=self.operation_timeout,
await asyncio.gather(owner_task, return_exceptions=True) )
except asyncio.TimeoutError as error:
raise TimeoutError(f"连接MCP服务器[{server_name}]超时") from 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): 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.sessions.pop(server_name, None)
self._session_last_used.pop(server_name, None) self._session_last_used.pop(server_name, None)
if handle is not None: if session_stack is not None:
handle.stop_event.set()
try: try:
await asyncio.wait_for(handle.owner_task, timeout=self.operation_timeout + 1) await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout)
except TimeoutError: except asyncio.TimeoutError:
logger.error(f"等待MCP会话[{server_name}]关闭超时") logger.error(f"关闭MCP会话[{server_name}]超时")
async def _get_or_create_session(self, server_name: str) -> ClientSession: async def _get_or_create_session(self, server_name: str) -> ClientSession:
"""获取可复用会话;若不存在或已过期则新建。""" """获取可复用会话;若不存在或已过期则新建。"""
@ -230,9 +191,9 @@ class MCPClient:
session = None session = None
if session is 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.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 self._session_last_used[server_name] = now
return self.sessions[server_name] return self.sessions[server_name]

View file

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