diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 8dbecb34..e5e47ccc 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -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( diff --git a/main/xiaozhi-server/core/handle/functionHandler.py b/main/xiaozhi-server/core/handle/functionHandler.py index 5f620880..c5e438ad 100644 --- a/main/xiaozhi-server/core/handle/functionHandler.py +++ b/main/xiaozhi-server/core/handle/functionHandler.py @@ -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): """注册配置中的函数,可以不同客户端使用不同的配置""" diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py index 45ef941a..e0c0f4e3 100644 --- a/main/xiaozhi-server/core/handle/helloHandle.py +++ b/main/xiaozhi-server/core/handle/helloHandle.py @@ -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)) diff --git a/main/xiaozhi-server/core/handle/mcpHandle.py b/main/xiaozhi-server/core/handle/mcpHandle.py new file mode 100644 index 00000000..1c587dd2 --- /dev/null +++ b/main/xiaozhi-server/core/handle/mcpHandle.py @@ -0,0 +1,357 @@ +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) + + # 处理参数 + 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 diff --git a/main/xiaozhi-server/core/handle/textHandle.py b/main/xiaozhi-server/core/handle/textHandle.py index 043a8b16..0fbae11e 100644 --- a/main/xiaozhi-server/core/handle/textHandle.py +++ b/main/xiaozhi-server/core/handle/textHandle.py @@ -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) diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index dbd1c339..707f80bb 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -47,6 +47,7 @@ class ASRProviderBase(ABC): if self.conn.client_voice_stop: asr_audio_task = copy.deepcopy(self.conn.asr_audio) self.conn.asr_audio.clear() + # 音频太短了,无法识别 self.conn.reset_vad_states() if len(asr_audio_task) > 15: diff --git a/main/xiaozhi-server/test/test_page.html b/main/xiaozhi-server/test/test_page.html index fbad025d..567332b8 100644 --- a/main/xiaozhi-server/test/test_page.html +++ b/main/xiaozhi-server/test/test_page.html @@ -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');