update:优化服务端mcp

This commit is contained in:
hrz
2025-06-26 14:35:12 +08:00
parent da8435cfea
commit ec4694f859
5 changed files with 280 additions and 18 deletions
@@ -2,5 +2,6 @@
from .mcp_manager import ServerMCPManager
from .mcp_executor import ServerMCPExecutor
from .mcp_client import ServerMCPClient
__all__ = ["ServerMCPManager", "ServerMCPExecutor"]
__all__ = ["ServerMCPManager", "ServerMCPExecutor", "ServerMCPClient"]
@@ -0,0 +1,208 @@
"""服务端MCP客户端"""
from __future__ import annotations
from datetime import timedelta
import asyncio
import os
import shutil
import concurrent.futures
from contextlib import AsyncExitStack
from typing import Optional, List, Dict, Any
from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client
from mcp.client.sse import sse_client
from config.logger import setup_logging
from core.utils.util import sanitize_tool_name
TAG = __name__
class ServerMCPClient:
"""服务端MCP客户端,用于连接和管理MCP服务"""
def __init__(self, config: Dict[str, Any]):
"""初始化服务端MCP客户端
Args:
config: MCP服务配置字典
"""
self.logger = setup_logging()
self.config = config
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 = [] # 原始工具对象
self.tools_dict: Dict[str, Any] = {}
self.name_mapping: Dict[str, str] = {}
async def initialize(self):
"""初始化MCP客户端连接"""
if self._worker_task:
return
self._worker_task = asyncio.create_task(
self._worker(), name="ServerMCPClientWorker"
)
await self._ready_evt.wait()
self.logger.bind(tag=TAG).info(
f"服务端MCP客户端已连接,可用工具: {[name for name in self.name_mapping.values()]}"
)
async def cleanup(self):
"""清理MCP客户端资源"""
if not self._worker_task:
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"服务端MCP客户端关闭错误: {e}")
finally:
self._worker_task = None
def has_tool(self, name: str) -> bool:
"""检查是否包含指定工具
Args:
name: 工具名称
Returns:
bool: 是否包含该工具
"""
return name in self.tools_dict
def get_available_tools(self) -> List[Dict[str, Any]]:
"""获取所有可用工具的定义
Returns:
List[Dict[str, Any]]: 工具定义列表
"""
return [
{
"type": "function",
"function": {
"name": name,
"description": tool.description,
"parameters": tool.inputSchema,
},
}
for name, tool in self.tools_dict.items()
]
async def call_tool(self, name: str, args: dict) -> Any:
"""调用指定工具
Args:
name: 工具名称
args: 工具参数
Returns:
Any: 工具执行结果
Raises:
RuntimeError: 客户端未初始化时抛出
"""
if not self.session:
raise RuntimeError("服务端MCP客户端未初始化")
real_name = self.name_mapping.get(name, name)
loop = self._worker_task.get_loop()
coro = self.session.call_tool(real_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)
def is_connected(self) -> bool:
"""检查MCP客户端是否连接正常
Returns:
bool: 如果客户端已连接并正常工作,返回True,否则返回False
"""
# 检查工作任务是否存在
if self._worker_task is None:
return False
# 检查工作任务是否已经完成或取消
if self._worker_task.done():
return False
# 检查会话是否存在
if self.session is None:
return False
# 所有检查都通过,连接正常
return True
async def _worker(self):
"""MCP客户端工作协程"""
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:
if "API_ACCESS_TOKEN" in self.config:
headers = {
"Authorization": f"Bearer {self.config['API_ACCESS_TOKEN']}"
}
else:
headers = {}
sse_r, sse_w = await stack.enter_async_context(
sse_client(self.config["url"], headers=headers)
)
read_stream, write_stream = sse_r, sse_w
else:
raise ValueError("MCP客户端配置必须包含'command''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
for t in self.tools:
sanitized = sanitize_tool_name(t.name)
self.tools_dict[sanitized] = t
self.name_mapping[sanitized] = t.name
self._ready_evt.set()
# 挂起等待关闭
await self._shutdown_evt.wait()
except Exception as e:
self.logger.bind(tag=TAG).error(f"服务端MCP客户端工作协程错误: {e}")
self._ready_evt.set()
raise
@@ -62,8 +62,9 @@ class ServerMCPExecutor(ToolExecutor):
for tool in mcp_tools:
func_def = tool.get("function", {})
tool_name = f"mcp_{func_def.get('name', '')}"
tool_name = func_def.get("name", "")
if tool_name == "":
continue
tools[tool_name] = ToolDefinition(
name=tool_name, description=tool, tool_type=ToolType.SERVER_MCP
)
@@ -6,6 +6,7 @@ import json
from typing import Dict, Any, List
from config.config_loader import get_project_dir
from config.logger import setup_logging
from .mcp_client import ServerMCPClient
TAG = __name__
logger = setup_logging()
@@ -23,7 +24,7 @@ class ServerMCPManager:
logger.bind(tag=TAG).warning(
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
)
self.clients: Dict[str, Any] = {}
self.clients: Dict[str, ServerMCPClient] = {}
self.tools = []
def load_config(self) -> Dict[str, Any]:
@@ -52,14 +53,13 @@ class ServerMCPManager:
continue
try:
# 这里可以添加真正的MCP客户端初始化逻辑
# 暂时使用简化版本
# 初始化服务端MCP客户端
logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}")
# client = MCPClient(srv_config)
# await client.initialize()
# self.clients[name] = client
# client_tools = client.get_available_tools()
# self.tools.extend(client_tools)
client = ServerMCPClient(srv_config)
await client.initialize()
self.clients[name] = client
client_tools = client.get_available_tools()
self.tools.extend(client_tools)
except Exception as e:
logger.bind(tag=TAG).error(
@@ -81,12 +81,66 @@ class ServerMCPManager:
return False
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Any:
"""执行工具调用"""
"""执行工具调用,失败时会尝试重新连接"""
logger.bind(tag=TAG).info(f"执行服务端MCP工具 {tool_name},参数: {arguments}")
# 这里可以添加真正的工具执行逻辑
# 暂时返回模拟结果
return f"服务端MCP工具 {tool_name} 执行结果"
max_retries = 3 # 最大重试次数
retry_interval = 2 # 重试间隔(秒)
# 找到对应的客户端
client_name = None
target_client = None
for name, client in self.clients.items():
if client.has_tool(tool_name):
client_name = name
target_client = client
break
if not target_client:
raise ValueError(f"工具 {tool_name} 在任意MCP服务中未找到")
# 带重试机制的工具调用
for attempt in range(max_retries):
try:
return await target_client.call_tool(tool_name, arguments)
except Exception as e:
# 最后一次尝试失败时直接抛出异常
if attempt == max_retries - 1:
raise
logger.bind(tag=TAG).warning(
f"执行工具 {tool_name} 失败 (尝试 {attempt+1}/{max_retries}): {e}"
)
# 尝试重新连接
logger.bind(tag=TAG).info(
f"重试前尝试重新连接 MCP 客户端 {client_name}"
)
try:
# 关闭旧的连接
await target_client.cleanup()
# 重新初始化客户端
config = self.load_config()
if client_name in config:
client = ServerMCPClient(config[client_name])
await client.initialize()
self.clients[client_name] = client
target_client = client
logger.bind(tag=TAG).info(
f"成功重新连接 MCP 客户端: {client_name}"
)
else:
logger.bind(tag=TAG).error(
f"Cannot reconnect MCP client {client_name}: config not found"
)
except Exception as reconnect_error:
logger.bind(tag=TAG).error(
f"Failed to reconnect MCP client {client_name}: {reconnect_error}"
)
# 等待一段时间再重试
await asyncio.sleep(retry_interval)
async def cleanup_all(self) -> None:
"""关闭所有 MCP客户端"""
@@ -43,7 +43,6 @@ class TTSProviderBase(ABC):
self.tts_text_buff = []
self.punctuations = (
"",
".",
"",
"?",
"",
@@ -59,7 +58,6 @@ class TTSProviderBase(ABC):
"",
",",
"",
".",
"",
"?",
"",
@@ -171,7 +169,7 @@ class TTSProviderBase(ABC):
)
)
# 对于单句的文本,进行分段处理
segments = re.split(r'([。!?!?;\n])', content_detail)
segments = re.split(r"([。!?!?;\n])", content_detail)
for seg in segments:
self.tts_text_queue.put(
TTSMessageDTO(