diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index b8f06f70..04f79970 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, @@ -145,6 +146,8 @@ class ConnectionHandler: ) # 在原来第一道关闭的基础上加60秒,进行二道关闭 self.audio_format = "opus" + # {"mcp":true} 表示启用MCP功能 + self.features = None async def handle_connection(self, ws): try: @@ -578,6 +581,12 @@ class ConnectionHandler: functions = None if self.intent_type == "function_call" and hasattr(self, "func_handler"): functions = self.func_handler.get_functions() + if self.mcp_client is not None: + 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: @@ -698,6 +707,14 @@ class ConnectionHandler: # 处理MCP工具调用 if self.mcp_manager.is_mcp_tool(function_name): result = self._handle_mcp_tool_call(function_call_data) + elif self.mcp_client is not None and self.mcp_client.has_tool(function_name): + # 如果是MCP工具调用 + self.logger.bind(tag=TAG).debug( + f"调用MCP工具: {function_name}, 参数: {function_arguments}" + ) + 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="") else: # 处理系统函数 result = self.func_handler.handle_llm_function_call( diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py index 6fd40401..a38a33ce 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..838fb6bd --- /dev/null +++ b/main/xiaozhi-server/core/handle/mcpHandle.py @@ -0,0 +1,277 @@ +import json +import asyncio +from concurrent.futures import Future + +TAG = __name__ + +class MCPClient: + """MCPClient,用于管理MCP状态和工具""" + def __init__(self): + self.tools = [] + self.ready = False + self.call_results = {} # To store Futures for tool call responses + self.next_id = 1 + self.lock = asyncio.Lock() + + async def has_tool(self, name: str) -> bool: + async with self.lock: + for tool in self.tools: + if tool["name"] == name: + return True + return False + + def get_available_tools(self) -> list: + # async with self.lock: + result = [] + for tool in self.tools: + function_def = { + "name": tool["name"], + "description": tool["description"], + "parameters": { + "type": tool["inputSchema"].get("type", "object"), + "properties": tool["inputSchema"].get("properties", {}), + "required": tool["inputSchema"].get("required", []) + } + } + result.append({"type": "function", "function": function_def}) + 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.append(tool_data) + + 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客户端准备就绪") + tool_result = await call_mcp_tool(conn, mcp_client, "self.get_device_status", {}) + print(f"Tool call result: {tool_result}") + 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 await 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 diff --git a/main/xiaozhi-server/core/handle/textHandle.py b/main/xiaozhi-server/core/handle/textHandle.py index 602f72aa..ace995e4 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 @@ -73,6 +74,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( @@ -152,5 +157,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/plugins_func/functions/handle_speaker_or_screen.py b/main/xiaozhi-server/plugins_func/functions/handle_speaker_or_screen.py index f89f30b0..8547c985 100644 --- a/main/xiaozhi-server/plugins_func/functions/handle_speaker_or_screen.py +++ b/main/xiaozhi-server/plugins_func/functions/handle_speaker_or_screen.py @@ -102,11 +102,11 @@ handle_device_function_desc = { } -@register_function( - "handle_speaker_volume_or_screen_brightness", - handle_device_function_desc, - ToolType.IOT_CTL, -) +# @register_function( +# "handle_speaker_volume_or_screen_brightness", +# handle_device_function_desc, +# ToolType.IOT_CTL, +# ) def handle_speaker_volume_or_screen_brightness( conn, device_type: str, action: str, value: int = None ):