mirror of
https://github.com/FuQuan233/nonebot-plugin-llmchat.git
synced 2026-08-13 10:09:27 +00:00
Compare commits
No commits in common. "3320758cf9eab08fa39438fc79926c69e9d02a48" and "f0914976131705f68971f00a4af3e4f86a21f2e9" have entirely different histories.
3320758cf9
...
f091497613
2 changed files with 22 additions and 105 deletions
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
Loading…
Add table
Add a link
Reference in a new issue