Merge pull request #1105 from CaixyPromise/feature/mcp-exit-fix

fix(MCP): 重构MCPClient为后台协程 + AsyncExitStack管理,解决进程退出时“Attempted to exit cancel scope in a different task”错误
This commit is contained in:
Junsen Huang
2025-05-07 00:47:37 +08:00
committed by GitHub
2 changed files with 117 additions and 85 deletions
+108 -77
View File
@@ -1,94 +1,125 @@
from __future__ import annotations
from datetime import timedelta from datetime import timedelta
from typing import Optional import asyncio, os, shutil, concurrent.futures
from contextlib import AsyncExitStack from contextlib import AsyncExitStack
import os, shutil from typing import Optional, List, Dict, Any
from mcp import ClientSession, StdioServerParameters from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client from mcp.client.stdio import stdio_client
from mcp.client.sse import sse_client from mcp.client.sse import sse_client
from config.logger import setup_logging from config.logger import setup_logging
TAG = __name__ TAG = __name__
class MCPClient: class MCPClient:
def __init__(self, config): def __init__(self, config: Dict[str, Any]):
# Initialize session and client objects
self.session: Optional[ClientSession] = None
self.exit_stack = AsyncExitStack()
self.logger = setup_logging() self.logger = setup_logging()
self.config = config self.config = config
self.tolls = []
self._worker_task: Optional[asyncio.Task] = None
self._ready_evt = asyncio.Event()
self._shutdown_evt = asyncio.Event()
self.session: Optional[ClientSession] = None
self.tools: List = []
async def initialize(self): async def initialize(self):
args = self.config.get("args", []) if self._worker_task:
return
self._worker_task = asyncio.create_task(self._worker(), name="MCPClientWorker")
await self._ready_evt.wait()
if "command" in self.config: self.logger.bind(tag=TAG).info(
command = ( f"Connected, tools = {[t.name for t in self.tools]}"
shutil.which("npx") )
if self.config["command"] == "npx"
else self.config["command"]
)
env = {**os.environ}
if self.config.get("env"):
env.update(self.config["env"])
server_params = StdioServerParameters(
command=command,
args=args,
env=env
)
stdio_transport = await self.exit_stack.enter_async_context(stdio_client(server_params))
self.stdio, self.write = stdio_transport
time_out_delta = timedelta(seconds=15)
self.session = await self.exit_stack.enter_async_context(
ClientSession(read_stream=self.stdio, write_stream=self.write, read_timeout_seconds=time_out_delta)
)
elif "url" in self.config:
sse_transport = await self.exit_stack.enter_async_context(sse_client(self.config["url"]))
self.sse_read, self.sse_write = sse_transport
self.session = await self.exit_stack.enter_async_context(
ClientSession(read_stream=self.sse_read, write_stream=self.sse_write)
)
else:
raise ValueError("MCPClient config must have 'command' or 'url'.")
await self.session.initialize()
# List available tools
response = await self.session.list_tools()
tools = response.tools
self.tools = tools
self.logger.bind(tag=TAG).info(f"Connected to server with tools:{[tool.name for tool in tools]}")
def has_tool(self, tool_name):
return any(tool.name == tool_name for tool in self.tools)
def get_available_tools(self):
available_tools = [{"type": "function", "function":{
"name": tool.name,
"description": tool.description,
"parameters": tool.inputSchema
} } for tool in self.tools]
return available_tools
async def call_tool(self, tool_name: str, tool_args: dict):
self.logger.bind(tag=TAG).info(f"MCPClient Calling tool {tool_name} with args: {tool_args}")
try:
response = await self.session.call_tool(tool_name, tool_args)
except Exception as e:
self.logger.bind(tag=TAG).error(f"Error calling tool {tool_name}: {e}")
from types import SimpleNamespace
error_content = SimpleNamespace(
type='text',
text=f"Error calling tool {tool_name}: {e}"
)
error_response = SimpleNamespace(
content=[error_content],
isError=True
)
return error_response
self.logger.bind(tag=TAG).info(f"MCPClient Response from tool {tool_name}: {response}")
return response
async def cleanup(self): async def cleanup(self):
"""Clean up resources""" if not self._worker_task:
await self.exit_stack.aclose() return
self._shutdown_evt.set()
try:
await asyncio.wait_for(self._worker_task, timeout=20)
except (asyncio.TimeoutError, Exception) as e:
self.logger.bind(tag=TAG).error(f"worker shutdown err: {e}")
finally:
self._worker_task = None
def has_tool(self, name: str) -> bool:
return any(t.name == name for t in self.tools)
def get_available_tools(self):
return [
{
"type": "function",
"function": {
"name": t.name,
"description": t.description,
"parameters": t.inputSchema,
},
}
for t in self.tools
]
async def call_tool(self, name: str, args: dict):
if not self.session:
raise RuntimeError("MCPClient not initialized")
loop = self._worker_task.get_loop()
coro = self.session.call_tool(name, args)
if loop is asyncio.get_running_loop():
return await coro
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
return await asyncio.wrap_future(fut)
async def _worker(self):
async with AsyncExitStack() as stack:
try:
# 建立 StdioClient
if "command" in self.config:
cmd = (
shutil.which("npx")
if self.config["command"] == "npx"
else self.config["command"]
)
env = {**os.environ, **self.config.get("env", {})}
params = StdioServerParameters(
command=cmd,
args=self.config.get("args", []),
env=env,
)
stdio_r, stdio_w = await stack.enter_async_context(stdio_client(params))
read_stream, write_stream = stdio_r, stdio_w
# 建立SSEClient
elif "url" in self.config:
sse_r, sse_w = await stack.enter_async_context(sse_client(self.config["url"]))
read_stream, write_stream = sse_r, sse_w
else:
raise ValueError("MCPClient config must include 'command' or 'url'")
self.session = await stack.enter_async_context(
ClientSession(
read_stream=read_stream,
write_stream=write_stream,
read_timeout_seconds=timedelta(seconds=15),
)
)
await self.session.initialize()
# 获取工具
self.tools = (await self.session.list_tools()).tools
self._ready_evt.set()
# 挂起等待关闭
await self._shutdown_evt.wait()
except Exception as e:
self.logger.bind(tag=TAG).error(f"worker error: {e}")
self._ready_evt.set()
raise
+9 -8
View File
@@ -1,5 +1,5 @@
"""MCP服务管理器""" """MCP服务管理器"""
import asyncio
import os, json import os, json
from typing import Dict, Any, List from typing import Dict, Any, List
from .MCPClient import MCPClient from .MCPClient import MCPClient
@@ -94,8 +94,8 @@ class MCPManager:
""" """
for tool in self.tools: for tool in self.tools:
if ( if (
tool.get("function") != None tool.get("function") != None
and tool["function"].get("name") == tool_name and tool["function"].get("name") == tool_name
): ):
return True return True
return False return False
@@ -120,12 +120,13 @@ class MCPManager:
raise ValueError(f"Tool {tool_name} not found in any MCP server") raise ValueError(f"Tool {tool_name} not found in any MCP server")
async def cleanup_all(self) -> None: async def cleanup_all(self) -> None:
for name, client in self.client.items(): """依次关闭所有 MCPClient,不让异常阻断整体流程。"""
for name, client in list(self.client.items()):
try: try:
await client.cleanup() await asyncio.wait_for(client.cleanup(), timeout=20)
self.logger.bind(tag=TAG).info(f"Cleaned up MCP client: {name}") self.logger.bind(tag=TAG).info(f"MCP client closed: {name}")
except Exception as e: except (asyncio.TimeoutError, Exception) as e:
self.logger.bind(tag=TAG).error( self.logger.bind(tag=TAG).error(
f"Error cleaning up MCP client {name}: {e}" f"Error closing MCP client {name}: {e}"
) )
self.client.clear() self.client.clear()