Merge pull request #1414 from myifeng/dev

修复 has_tool错误;优化检索复杂度
This commit is contained in:
hrz
2025-05-30 16:03:36 +08:00
committed by GitHub
6 changed files with 344 additions and 3 deletions
+24 -1
View File
@@ -10,6 +10,7 @@ import threading
import traceback
import subprocess
import websockets
from core.handle.mcpHandle import call_mcp_tool
from core.utils.util import (
extract_json_from_string,
check_vad_update,
@@ -146,6 +147,8 @@ class ConnectionHandler:
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
self.audio_format = "opus"
# {"mcp":true} 表示启用MCP功能
self.features = None
async def handle_connection(self, ws):
try:
@@ -579,6 +582,12 @@ class ConnectionHandler:
functions = None
if self.intent_type == "function_call" and hasattr(self, "func_handler"):
functions = self.func_handler.get_functions()
if hasattr(self, "mcp_client"):
mcp_tools = self.mcp_client.get_available_tools()
if mcp_tools is not None and len(mcp_tools) > 0:
if functions is None:
functions = []
functions.extend(mcp_tools)
response_message = []
try:
@@ -696,9 +705,23 @@ class ConnectionHandler:
"arguments": function_arguments,
}
# 处理MCP工具调用
# 处理Server端MCP工具调用
if self.mcp_manager.is_mcp_tool(function_name):
result = self._handle_mcp_tool_call(function_call_data)
elif hasattr(self, "mcp_client") and self.mcp_client.has_tool(function_name):
# 如果是小智端MCP工具调用
self.logger.bind(tag=TAG).debug(
f"调用小智端MCP工具: {function_name}, 参数: {function_arguments}"
)
try:
result = asyncio.run_coroutine_threadsafe(call_mcp_tool(self, self.mcp_client, function_name, function_arguments), self.loop).result()
self.logger.bind(tag=TAG).debug(f"MCP工具调用结果: {result}")
result = ActionResponse(action=Action.REQLLM, result=result, response="")
except Exception as e:
self.logger.bind(tag=TAG).error(f"MCP工具调用失败: {e}")
result = ActionResponse(
action=Action.REQLLM, result="MCP工具调用失败", response=""
)
else:
# 处理系统函数
result = self.func_handler.handle_llm_function_call(
@@ -61,7 +61,7 @@ class FunctionHandler:
self.function_registry.register_function("plugin_loader")
self.function_registry.register_function("get_time")
self.function_registry.register_function("get_lunar")
self.function_registry.register_function("handle_speaker_volume_or_screen_brightness")
# self.function_registry.register_function("handle_speaker_volume_or_screen_brightness")
def register_config_functions(self):
"""注册配置中的函数,可以不同客户端使用不同的配置"""
@@ -4,6 +4,7 @@ import json
import random
import shutil
import asyncio
from core.handle.mcpHandle import MCPClient, send_mcp_initialize_message, send_mcp_tools_list_request
from core.handle.sendAudioHandle import send_stt_message
from core.utils.util import remove_punctuation_and_length
from core.providers.tts.dto.dto import ContentType, InterfaceType
@@ -31,6 +32,17 @@ async def handleHelloMessage(conn, msg_json):
if conn.asr is not None:
conn.asr.set_audio_format(format)
conn.welcome_msg["audio_params"] = audio_params
features = msg_json.get("features")
if features:
conn.logger.bind(tag=TAG).info(f"客户端特性: {features}")
conn.features = features
if features.get("mcp"):
conn.logger.bind(tag=TAG).info("客户端支持MCP")
conn.mcp_client = MCPClient()
# 发送初始化
asyncio.create_task(send_mcp_initialize_message(conn))
# 发送mcp消息,获取tools列表
asyncio.create_task(send_mcp_tools_list_request(conn))
await conn.websocket.send(json.dumps(conn.welcome_msg))
@@ -0,0 +1,279 @@
import json
import asyncio
from concurrent.futures import Future
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初始化消息"""
payload = {
"jsonrpc": "2.0",
"id": 1, # mcpInitializeID
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {
"roots": {"listChanged": True},
"sampling": {},
},
"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)
payload = {
"jsonrpc": "2.0",
"id": tool_call_id,
"method": "tools/call",
"params": {
"name": tool_name,
"arguments": json.loads(args) if isinstance(args, str) else args
},
}
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]:
return content[0]["text"]
return 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
@@ -1,6 +1,7 @@
import json
from core.handle.abortHandle import handleAbortMessage
from core.handle.helloHandle import handleHelloMessage
from core.handle.mcpHandle import handle_mcp_message
from core.utils.util import remove_punctuation_and_length, filter_sensitive_info
from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
from core.handle.sendAudioHandle import send_stt_message, send_tts_message
@@ -74,6 +75,10 @@ async def handleTextMessage(conn, message):
asyncio.create_task(handleIotDescriptors(conn, msg_json["descriptors"]))
if "states" in msg_json:
asyncio.create_task(handleIotStatus(conn, msg_json["states"]))
elif msg_json["type"] == "mcp":
conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message}")
if "payload" in msg_json:
asyncio.create_task(handle_mcp_message(conn, conn.mcp_client, msg_json["payload"]))
elif msg_json["type"] == "server":
# 记录日志时过滤敏感信息
conn.logger.bind(tag=TAG).info(
@@ -153,5 +158,7 @@ async def handleTextMessage(conn, message):
# 重启服务器
elif msg_json["action"] == "restart":
await conn.handle_restart(msg_json)
else:
conn.logger.bind(tag=TAG).error(f"收到未知类型消息:{message}")
except json.JSONDecodeError:
await conn.websocket.send(message)
+21 -1
View File
@@ -1484,6 +1484,23 @@
if (message.text && message.text !== '😊') {
addMessage(message.text);
}
}else if (message.type === 'mcp') {
const payload = message.payload || {};
log(`服务器下发: ${JSON.stringify(message)}`, 'info');
if (payload) {
// 模拟小智客户端行为
if(payload.method === 'tools/list'){
const replay_message = JSON.stringify({"session_id":"","type":"mcp","payload":{"jsonrpc":"2.0","id":2,"result":{"tools":[{"name":"self.get_device_status","description":"Provides the real-time information of the device, including the current status of the audio speaker, screen, battery, network, etc.\nUse this tool for: \n1. Answering questions about current condition (e.g. what is the current volume of the audio speaker?)\n2. As the first step to control the device (e.g. turn up / down the volume of the audio speaker, etc.)","inputSchema":{"type":"object","properties":{}}},{"name":"self.audio_speaker.set_volume","description":"Set the volume of the audio speaker. If the current volume is unknown, you must call `self.get_device_status` tool first and then call this tool.","inputSchema":{"type":"object","properties":{"volume":{"type":"integer","minimum":0,"maximum":100}},"required":["volume"]}},{"name":"self.screen.set_brightness","description":"Set the brightness of the screen.","inputSchema":{"type":"object","properties":{"brightness":{"type":"integer","minimum":0,"maximum":100}},"required":["brightness"]}},{"name":"self.screen.set_theme","description":"Set the theme of the screen. The theme can be 'light' or 'dark'.","inputSchema":{"type":"object","properties":{"theme":{"type":"string"}},"required":["theme"]}}]}}})
websocket.send(replay_message);
log(`回复MCP消息: ${replay_message}`, 'info');
} else if(payload.method === 'tools/call'){
// 模拟回复
const replay_message = JSON.stringify({"session_id":"9f261599","type":"mcp","payload":{"jsonrpc":"2.0","id": payload.id,"result":{"content":[{"type":"text","text":"true"}],"isError":false}}})
websocket.send(replay_message);
log(`回复MCP消息: ${replay_message}`, 'info');
}
}
} else {
// 未知消息类型
log(`未知消息类型: ${message.type}`, 'info');
@@ -1523,7 +1540,10 @@
device_id: config.deviceId,
device_name: config.deviceName,
device_mac: config.deviceMac,
token: config.token
token: config.token,
features: {
mcp: true
}
};
log('发送hello握手消息', 'info');