mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 09:33:55 +08:00
update:优化服务端mcp
This commit is contained in:
@@ -2,5 +2,6 @@
|
|||||||
|
|
||||||
from .mcp_manager import ServerMCPManager
|
from .mcp_manager import ServerMCPManager
|
||||||
from .mcp_executor import ServerMCPExecutor
|
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:
|
for tool in mcp_tools:
|
||||||
func_def = tool.get("function", {})
|
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(
|
tools[tool_name] = ToolDefinition(
|
||||||
name=tool_name, description=tool, tool_type=ToolType.SERVER_MCP
|
name=tool_name, description=tool, tool_type=ToolType.SERVER_MCP
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import json
|
|||||||
from typing import Dict, Any, List
|
from typing import Dict, Any, List
|
||||||
from config.config_loader import get_project_dir
|
from config.config_loader import get_project_dir
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
from .mcp_client import ServerMCPClient
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -23,7 +24,7 @@ class ServerMCPManager:
|
|||||||
logger.bind(tag=TAG).warning(
|
logger.bind(tag=TAG).warning(
|
||||||
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
|
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
|
||||||
)
|
)
|
||||||
self.clients: Dict[str, Any] = {}
|
self.clients: Dict[str, ServerMCPClient] = {}
|
||||||
self.tools = []
|
self.tools = []
|
||||||
|
|
||||||
def load_config(self) -> Dict[str, Any]:
|
def load_config(self) -> Dict[str, Any]:
|
||||||
@@ -52,14 +53,13 @@ class ServerMCPManager:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# 这里可以添加真正的MCP客户端初始化逻辑
|
# 初始化服务端MCP客户端
|
||||||
# 暂时使用简化版本
|
|
||||||
logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}")
|
logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}")
|
||||||
# client = MCPClient(srv_config)
|
client = ServerMCPClient(srv_config)
|
||||||
# await client.initialize()
|
await client.initialize()
|
||||||
# self.clients[name] = client
|
self.clients[name] = client
|
||||||
# client_tools = client.get_available_tools()
|
client_tools = client.get_available_tools()
|
||||||
# self.tools.extend(client_tools)
|
self.tools.extend(client_tools)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(
|
||||||
@@ -81,12 +81,66 @@ class ServerMCPManager:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Any:
|
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Any:
|
||||||
"""执行工具调用"""
|
"""执行工具调用,失败时会尝试重新连接"""
|
||||||
logger.bind(tag=TAG).info(f"执行服务端MCP工具 {tool_name},参数: {arguments}")
|
logger.bind(tag=TAG).info(f"执行服务端MCP工具 {tool_name},参数: {arguments}")
|
||||||
|
|
||||||
# 这里可以添加真正的工具执行逻辑
|
max_retries = 3 # 最大重试次数
|
||||||
# 暂时返回模拟结果
|
retry_interval = 2 # 重试间隔(秒)
|
||||||
return f"服务端MCP工具 {tool_name} 执行结果"
|
|
||||||
|
# 找到对应的客户端
|
||||||
|
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:
|
async def cleanup_all(self) -> None:
|
||||||
"""关闭所有 MCP客户端"""
|
"""关闭所有 MCP客户端"""
|
||||||
|
|||||||
@@ -43,7 +43,6 @@ class TTSProviderBase(ABC):
|
|||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.punctuations = (
|
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:
|
for seg in segments:
|
||||||
self.tts_text_queue.put(
|
self.tts_text_queue.put(
|
||||||
TTSMessageDTO(
|
TTSMessageDTO(
|
||||||
|
|||||||
Reference in New Issue
Block a user