diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/__init__.py b/main/xiaozhi-server/core/providers/tools/server_mcp/__init__.py index ba73b522..aecd80b8 100644 --- a/main/xiaozhi-server/core/providers/tools/server_mcp/__init__.py +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/__init__.py @@ -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"] diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_client.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_client.py new file mode 100644 index 00000000..8b60ac1b --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_client.py @@ -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 diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py index a41b29dc..9ae15d4b 100644 --- a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py @@ -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 ) diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py index 92e9d8cd..6589c302 100644 --- a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py @@ -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客户端""" diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 1a9d069c..223cb807 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -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(