mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-22 23:23:55 +08:00
375 lines
14 KiB
Python
375 lines
14 KiB
Python
import json
|
|
import asyncio
|
|
from concurrent.futures import Future
|
|
from core.utils.util import get_vision_url
|
|
from core.utils.auth import AuthToken
|
|
|
|
TAG = __name__
|
|
|
|
|
|
class MCPClient:
|
|
"""MCPClient,用于管理MCP状态和工具"""
|
|
|
|
def __init__(self):
|
|
self.tools = {} # Dictionary for O(1) lookup
|
|
self.ready = False
|
|
self.call_results = {} # To store Futures for tool call responses
|
|
self.next_id = 1
|
|
self.lock = asyncio.Lock()
|
|
self._cached_available_tools = None # Cache for get_available_tools
|
|
|
|
def has_tool(self, name: str) -> bool:
|
|
return name in self.tools
|
|
|
|
def get_available_tools(self) -> list:
|
|
# Check if the cache is valid
|
|
if self._cached_available_tools is not None:
|
|
return self._cached_available_tools
|
|
|
|
# If cache is not valid, regenerate the list
|
|
result = []
|
|
for tool_name, tool_data in self.tools.items():
|
|
function_def = {
|
|
"name": tool_data["name"],
|
|
"description": tool_data["description"],
|
|
"parameters": {
|
|
"type": tool_data["inputSchema"].get("type", "object"),
|
|
"properties": tool_data["inputSchema"].get("properties", {}),
|
|
"required": tool_data["inputSchema"].get("required", []),
|
|
},
|
|
}
|
|
result.append({"type": "function", "function": function_def})
|
|
|
|
self._cached_available_tools = result # Store the generated list in cache
|
|
return result
|
|
|
|
async def is_ready(self) -> bool:
|
|
async with self.lock:
|
|
return self.ready
|
|
|
|
async def set_ready(self, status: bool):
|
|
async with self.lock:
|
|
self.ready = status
|
|
|
|
async def add_tool(self, tool_data: dict):
|
|
async with self.lock:
|
|
self.tools[tool_data["name"]] = tool_data
|
|
self._cached_available_tools = (
|
|
None # Invalidate the cache when a tool is added
|
|
)
|
|
|
|
async def get_next_id(self) -> int:
|
|
async with self.lock:
|
|
current_id = self.next_id
|
|
self.next_id += 1
|
|
return current_id
|
|
|
|
async def register_call_result_future(self, id: int, future: Future):
|
|
async with self.lock:
|
|
self.call_results[id] = future
|
|
|
|
async def resolve_call_result(self, id: int, result: any):
|
|
async with self.lock:
|
|
if id in self.call_results:
|
|
future = self.call_results.pop(id)
|
|
if not future.done():
|
|
future.set_result(result)
|
|
|
|
async def reject_call_result(self, id: int, exception: Exception):
|
|
async with self.lock:
|
|
if id in self.call_results:
|
|
future = self.call_results.pop(id)
|
|
if not future.done():
|
|
future.set_exception(exception)
|
|
|
|
async def cleanup_call_result(self, id: int):
|
|
async with self.lock:
|
|
if id in self.call_results:
|
|
self.call_results.pop(id)
|
|
|
|
|
|
async def send_mcp_message(conn, payload: dict):
|
|
"""Helper to send MCP messages, encapsulating common logic."""
|
|
if not conn.features.get("mcp"):
|
|
conn.logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
|
|
return
|
|
|
|
message = json.dumps({"type": "mcp", "payload": payload})
|
|
|
|
try:
|
|
await conn.websocket.send(message)
|
|
conn.logger.bind(tag=TAG).info(f"成功发送MCP消息: {message}")
|
|
except Exception as e:
|
|
conn.logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}")
|
|
|
|
|
|
async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
|
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
|
|
conn.logger.bind(tag=TAG).info(f"处理MCP消息: {payload}")
|
|
|
|
if not isinstance(payload, dict):
|
|
conn.logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误")
|
|
return
|
|
|
|
# Handle result
|
|
if "result" in payload:
|
|
result = payload["result"]
|
|
msg_id = int(payload.get("id", 0))
|
|
|
|
# Check for tool call response first
|
|
if msg_id in mcp_client.call_results:
|
|
conn.logger.bind(tag=TAG).debug(
|
|
f"收到工具调用响应,ID: {msg_id}, 结果: {result}"
|
|
)
|
|
await mcp_client.resolve_call_result(msg_id, result)
|
|
return
|
|
|
|
if msg_id == 1: # mcpInitializeID
|
|
conn.logger.bind(tag=TAG).debug("收到MCP初始化响应")
|
|
server_info = result.get("serverInfo")
|
|
if isinstance(server_info, dict):
|
|
name = server_info.get("name")
|
|
version = server_info.get("version")
|
|
conn.logger.bind(tag=TAG).info(
|
|
f"客户端MCP服务器信息: name={name}, version={version}"
|
|
)
|
|
await send_mcp_tools_list_request(
|
|
conn
|
|
) # After initialization, request tool list
|
|
return
|
|
|
|
elif msg_id == 2: # mcpToolsListID
|
|
conn.logger.bind(tag=TAG).debug("收到MCP工具列表响应")
|
|
if isinstance(result, dict) and "tools" in result:
|
|
tools_data = result["tools"]
|
|
if not isinstance(tools_data, list):
|
|
conn.logger.bind(tag=TAG).error("工具列表格式错误")
|
|
return
|
|
|
|
conn.logger.bind(tag=TAG).info(
|
|
f"客户端设备支持的工具数量: {len(tools_data)}"
|
|
)
|
|
|
|
for i, tool in enumerate(tools_data):
|
|
if not isinstance(tool, dict):
|
|
continue
|
|
|
|
name = tool.get("name", "")
|
|
description = tool.get("description", "")
|
|
input_schema = {"type": "object", "properties": {}, "required": []}
|
|
|
|
if "inputSchema" in tool and isinstance(tool["inputSchema"], dict):
|
|
schema = tool["inputSchema"]
|
|
input_schema["type"] = schema.get("type", "object")
|
|
input_schema["properties"] = schema.get("properties", {})
|
|
input_schema["required"] = [
|
|
s for s in schema.get("required", []) if isinstance(s, str)
|
|
]
|
|
|
|
new_tool = {
|
|
"name": name,
|
|
"description": description,
|
|
"inputSchema": input_schema,
|
|
}
|
|
await mcp_client.add_tool(new_tool)
|
|
conn.logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}")
|
|
|
|
next_cursor = result.get("nextCursor", "")
|
|
if next_cursor:
|
|
conn.logger.bind(tag=TAG).info(
|
|
f"有更多工具,nextCursor: {next_cursor}"
|
|
)
|
|
await send_mcp_tools_list_continue_request(conn, next_cursor)
|
|
else:
|
|
await mcp_client.set_ready(True)
|
|
conn.logger.bind(tag=TAG).info("所有工具已获取,MCP客户端准备就绪")
|
|
return
|
|
|
|
# Handle method calls (requests from the client)
|
|
elif "method" in payload:
|
|
method = payload["method"]
|
|
conn.logger.bind(tag=TAG).info(f"收到MCP客户端请求: {method}")
|
|
|
|
elif "error" in payload:
|
|
error_data = payload["error"]
|
|
error_msg = error_data.get("message", "未知错误")
|
|
conn.logger.bind(tag=TAG).error(f"收到MCP错误响应: {error_msg}")
|
|
|
|
msg_id = int(payload.get("id", 0))
|
|
if msg_id in mcp_client.call_results:
|
|
await mcp_client.reject_call_result(
|
|
msg_id, Exception(f"MCP错误: {error_msg}")
|
|
)
|
|
|
|
|
|
# --- Outgoing MCP Messages ---
|
|
|
|
|
|
async def send_mcp_initialize_message(conn):
|
|
"""发送MCP初始化消息"""
|
|
|
|
vision_url = get_vision_url(conn.config)
|
|
|
|
# 密钥生成token
|
|
auth = AuthToken(conn.config["server"]["auth_key"])
|
|
token = auth.generate_token(conn.headers.get("device-id"))
|
|
|
|
vision = {
|
|
"url": vision_url,
|
|
"token": token,
|
|
}
|
|
|
|
conn.logger.bind(tag=TAG).info(f"视觉服务信息: {vision}")
|
|
|
|
payload = {
|
|
"jsonrpc": "2.0",
|
|
"id": 1, # mcpInitializeID
|
|
"method": "initialize",
|
|
"params": {
|
|
"protocolVersion": "2024-11-05",
|
|
"capabilities": {
|
|
"roots": {"listChanged": True},
|
|
"sampling": {},
|
|
"vision": vision,
|
|
},
|
|
"clientInfo": {
|
|
"name": "XiaozhiClient",
|
|
"version": "1.0.0",
|
|
},
|
|
},
|
|
}
|
|
conn.logger.bind(tag=TAG).info("发送MCP初始化消息")
|
|
await send_mcp_message(conn, payload)
|
|
|
|
|
|
async def send_mcp_tools_list_request(conn):
|
|
"""发送MCP工具列表请求"""
|
|
payload = {
|
|
"jsonrpc": "2.0",
|
|
"id": 2, # mcpToolsListID
|
|
"method": "tools/list",
|
|
}
|
|
conn.logger.bind(tag=TAG).debug("发送MCP工具列表请求")
|
|
await send_mcp_message(conn, payload)
|
|
|
|
|
|
async def send_mcp_tools_list_continue_request(conn, cursor: str):
|
|
"""发送带有cursor的MCP工具列表请求"""
|
|
payload = {
|
|
"jsonrpc": "2.0",
|
|
"id": 2, # mcpToolsListID (same ID for continuation)
|
|
"method": "tools/list",
|
|
"params": {"cursor": cursor},
|
|
}
|
|
conn.logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
|
|
await send_mcp_message(conn, payload)
|
|
|
|
|
|
async def call_mcp_tool(
|
|
conn, mcp_client: MCPClient, tool_name: str, args: str = "{}", timeout: int = 30
|
|
):
|
|
"""
|
|
调用指定的工具,并等待响应
|
|
"""
|
|
if not await mcp_client.is_ready():
|
|
raise RuntimeError("MCP客户端尚未准备就绪")
|
|
|
|
if not mcp_client.has_tool(tool_name):
|
|
raise ValueError(f"工具 {tool_name} 不存在")
|
|
|
|
tool_call_id = await mcp_client.get_next_id()
|
|
result_future = asyncio.Future()
|
|
await mcp_client.register_call_result_future(tool_call_id, result_future)
|
|
|
|
# 处理参数
|
|
try:
|
|
if isinstance(args, str):
|
|
# 确保字符串是有效的JSON
|
|
if not args.strip():
|
|
arguments = {}
|
|
else:
|
|
try:
|
|
# 尝试直接解析
|
|
arguments = json.loads(args)
|
|
except json.JSONDecodeError:
|
|
# 如果解析失败,尝试合并多个JSON对象
|
|
try:
|
|
# 使用正则表达式匹配所有JSON对象
|
|
import re
|
|
|
|
json_objects = re.findall(r"\{[^{}]*\}", args)
|
|
if len(json_objects) > 1:
|
|
# 合并所有JSON对象
|
|
merged_dict = {}
|
|
for json_str in json_objects:
|
|
try:
|
|
obj = json.loads(json_str)
|
|
if isinstance(obj, dict):
|
|
merged_dict.update(obj)
|
|
except json.JSONDecodeError:
|
|
continue
|
|
if merged_dict:
|
|
arguments = merged_dict
|
|
else:
|
|
raise ValueError(f"无法解析任何有效的JSON对象: {args}")
|
|
else:
|
|
raise ValueError(f"参数JSON解析失败: {args}")
|
|
except Exception as e:
|
|
conn.logger.bind(tag=TAG).error(
|
|
f"参数JSON解析失败: {str(e)}, 原始参数: {args}"
|
|
)
|
|
raise ValueError(f"参数JSON解析失败: {str(e)}")
|
|
elif isinstance(args, dict):
|
|
arguments = args
|
|
else:
|
|
raise ValueError(f"参数类型错误,期望字符串或字典,实际类型: {type(args)}")
|
|
|
|
# 确保参数是字典类型
|
|
if not isinstance(arguments, dict):
|
|
raise ValueError(f"参数必须是字典类型,实际类型: {type(arguments)}")
|
|
|
|
except Exception as e:
|
|
if not isinstance(e, ValueError):
|
|
raise ValueError(f"参数处理失败: {str(e)}")
|
|
raise e
|
|
|
|
payload = {
|
|
"jsonrpc": "2.0",
|
|
"id": tool_call_id,
|
|
"method": "tools/call",
|
|
"params": {"name": tool_name, "arguments": arguments},
|
|
}
|
|
|
|
conn.logger.bind(tag=TAG).info(
|
|
f"发送客户端mcp工具调用请求: {tool_name},参数: {args}"
|
|
)
|
|
await send_mcp_message(conn, payload)
|
|
|
|
try:
|
|
# Wait for response or timeout
|
|
raw_result = await asyncio.wait_for(result_future, timeout=timeout)
|
|
conn.logger.bind(tag=TAG).info(
|
|
f"客户端mcp工具调用 {tool_name} 成功,原始结果: {raw_result}"
|
|
)
|
|
|
|
if isinstance(raw_result, dict):
|
|
if raw_result.get("isError") is True:
|
|
error_msg = raw_result.get(
|
|
"error", "工具调用返回错误,但未提供具体错误信息"
|
|
)
|
|
raise RuntimeError(f"工具调用错误: {error_msg}")
|
|
|
|
content = raw_result.get("content")
|
|
if isinstance(content, list) and len(content) > 0:
|
|
if isinstance(content[0], dict) and "text" in content[0]:
|
|
# 直接返回文本内容,不进行JSON解析
|
|
return content[0]["text"]
|
|
# 如果结果不是预期的格式,将其转换为字符串
|
|
return str(raw_result)
|
|
except asyncio.TimeoutError:
|
|
await mcp_client.cleanup_call_result(tool_call_id)
|
|
raise TimeoutError("工具调用请求超时")
|
|
except Exception as e:
|
|
await mcp_client.cleanup_call_result(tool_call_id)
|
|
raise e
|