mirror of
https://github.com/FuQuan233/nonebot-plugin-llmchat.git
synced 2026-08-13 10:09:27 +00:00
Compare commits
2 commits
f091497613
...
3320758cf9
| Author | SHA1 | Date | |
|---|---|---|---|
| 3320758cf9 | |||
| f2014723c5 |
2 changed files with 105 additions and 22 deletions
|
|
@ -1,5 +1,6 @@
|
||||||
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
|
||||||
|
|
||||||
|
|
@ -14,6 +15,14 @@ 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
|
||||||
|
|
@ -34,7 +43,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 = 30,
|
operation_timeout: int = 100,
|
||||||
):
|
):
|
||||||
if self._initialized:
|
if self._initialized:
|
||||||
return
|
return
|
||||||
|
|
@ -48,7 +57,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_exit_stacks: dict[str, AsyncExitStack] = {}
|
self._session_handles: dict[str, _SessionHandle] = {}
|
||||||
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
|
||||||
|
|
@ -89,16 +98,34 @@ 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 _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()
|
session_stack = AsyncExitStack()
|
||||||
try:
|
try:
|
||||||
return await self._initialize_server_session(server_name, session_stack)
|
session, _ = await self._initialize_server_session(server_name, session_stack)
|
||||||
except BaseException:
|
ready.set_result(session)
|
||||||
try:
|
await stop_event.wait()
|
||||||
await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout)
|
except asyncio.CancelledError:
|
||||||
except (asyncio.TimeoutError, RuntimeError):
|
if not ready.done():
|
||||||
logger.error(f"清理未完成的MCP会话[{server_name}]失败或超时")
|
ready.cancel()
|
||||||
raise
|
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(
|
async def _initialize_server_session(
|
||||||
self,
|
self,
|
||||||
|
|
@ -155,26 +182,38 @@ 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, 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:
|
try:
|
||||||
return await asyncio.wait_for(
|
session = await asyncio.wait_for(asyncio.shield(ready), timeout=self.operation_timeout)
|
||||||
self._open_server_session(server_name),
|
except TimeoutError as error:
|
||||||
timeout=self.operation_timeout,
|
owner_task.cancel()
|
||||||
)
|
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):
|
||||||
"""关闭指定服务器会话。"""
|
"""关闭指定服务器会话。"""
|
||||||
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.sessions.pop(server_name, None)
|
||||||
self._session_last_used.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:
|
try:
|
||||||
await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout)
|
await asyncio.wait_for(handle.owner_task, timeout=self.operation_timeout + 1)
|
||||||
except asyncio.TimeoutError:
|
except 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:
|
||||||
"""获取可复用会话;若不存在或已过期则新建。"""
|
"""获取可复用会话;若不存在或已过期则新建。"""
|
||||||
|
|
@ -191,9 +230,9 @@ class MCPClient:
|
||||||
session = None
|
session = None
|
||||||
|
|
||||||
if session is 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.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
|
self._session_last_used[server_name] = now
|
||||||
return self.sessions[server_name]
|
return self.sessions[server_name]
|
||||||
|
|
|
||||||
44
tests/test_mcpclient.py
Normal file
44
tests/test_mcpclient.py
Normal 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()
|
||||||
Loading…
Add table
Add a link
Reference in a new issue