diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index 6dc0b25d..71c13ec7 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -104,7 +104,8 @@ wakeup_words: - "喵喵同学" - "小滨小滨" - "小冰小冰" - +# MCP接入点地址 +mcp_endpoint: 你的接入点 websocket地址 # 插件的基础配置 plugins: # 获取天气插件的配置,这里填写你的api_key diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 3ad94592..7826a1e9 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -10,7 +10,6 @@ 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, @@ -18,7 +17,6 @@ from core.utils.util import ( filter_sensitive_info, ) from typing import Dict, Any -from core.mcp.manager import MCPManager from core.utils.modules_initialize import ( initialize_modules, initialize_tts, @@ -30,7 +28,7 @@ from concurrent.futures import ThreadPoolExecutor from core.utils.dialogue import Message, Dialogue from core.providers.asr.dto.dto import InterfaceType from core.handle.textHandle import handleTextMessage -from core.handle.functionHandler import FunctionHandler +from core.providers.tools.unified_tool_handler import UnifiedToolHandler from plugins_func.loadplugins import auto_import_modules from plugins_func.register import Action, ActionResponse from core.auth import AuthMiddleware, AuthenticationError @@ -586,14 +584,12 @@ class ConnectionHandler: self.intent.set_llm(self.llm) self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型") - """加载插件""" - self.func_handler = FunctionHandler(self) - self.mcp_manager = MCPManager(self) + """加载统一工具处理器""" + self.func_handler = UnifiedToolHandler(self) - """加载MCP工具""" - asyncio.run_coroutine_threadsafe( - self.mcp_manager.initialize_servers(), self.loop - ) + # 异步初始化工具处理器 + if hasattr(self, "loop") and self.loop: + asyncio.run_coroutine_threadsafe(self.func_handler._initialize(), self.loop) def change_system_prompt(self, prompt): self.prompt = prompt @@ -611,12 +607,6 @@ 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: @@ -630,7 +620,6 @@ class ConnectionHandler: self.sentence_id = str(uuid.uuid4().hex) - if self.intent_type == "function_call" and functions is not None: # 使用支持functions的streaming接口 llm_responses = self.llm.response_with_functions( @@ -734,59 +723,13 @@ class ConnectionHandler: "arguments": function_arguments, } - # 处理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}") - - resultJson = None - if isinstance(result, str): - try: - resultJson = json.loads(result) - except Exception as e: - self.logger.bind(tag=TAG).error( - f"解析MCP工具返回结果失败: {e}" - ) - - # 视觉大模型不经过二次LLM处理 - if ( - resultJson is not None - and isinstance(resultJson, dict) - and "action" in resultJson - ): - result = ActionResponse( - action=Action[resultJson["action"]], - result=None, - response=resultJson.get("response", ""), - ) - else: - 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( + # 使用统一工具处理器处理所有工具调用 + result = asyncio.run_coroutine_threadsafe( + self.func_handler.handle_llm_function_call( self, function_call_data - ) + ), + self.loop, + ).result() self._handle_function_result(result, function_call_data) # 存储对话内容 @@ -809,48 +752,6 @@ class ConnectionHandler: return True - def _handle_mcp_tool_call(self, function_call_data): - function_arguments = function_call_data["arguments"] - function_name = function_call_data["name"] - try: - args_dict = function_arguments - if isinstance(function_arguments, str): - try: - args_dict = json.loads(function_arguments) - except json.JSONDecodeError: - self.logger.bind(tag=TAG).error( - f"无法解析 function_arguments: {function_arguments}" - ) - return ActionResponse( - action=Action.REQLLM, result="参数解析失败", response="" - ) - - tool_result = asyncio.run_coroutine_threadsafe( - self.mcp_manager.execute_tool(function_name, args_dict), self.loop - ).result() - # meta=None content=[TextContent(type='text', text='北京当前天气:\n温度: 21°C\n天气: 晴\n湿度: 6%\n风向: 西北 风\n风力等级: 5级', annotations=None)] isError=False - content_text = "" - if tool_result is not None and tool_result.content is not None: - for content in tool_result.content: - content_type = content.type - if content_type == "text": - content_text = content.text - elif content_type == "image": - pass - - if len(content_text) > 0: - return ActionResponse( - action=Action.REQLLM, result=content_text, response="" - ) - - except Exception as e: - self.logger.bind(tag=TAG).error(f"MCP工具调用错误: {e}") - return ActionResponse( - action=Action.REQLLM, result="工具调用出错", response="" - ) - - return ActionResponse(action=Action.REQLLM, result="工具调用出错", response="") - def _handle_function_result(self, result, function_call_data): if result.action == Action.RESPONSE: # 直接回复前端 text = result.response @@ -890,7 +791,7 @@ class ConnectionHandler: ) self.chat(text, tool_call=True) elif result.action == Action.NOTFOUND or result.action == Action.ERROR: - text = result.result + text = result.response if result.response else result.result self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text) self.dialogue.put(Message(role="assistant", content=text)) else: @@ -945,9 +846,9 @@ class ConnectionHandler: self.timeout_task.cancel() self.timeout_task = None - # 清理MCP资源 - if hasattr(self, "mcp_manager") and self.mcp_manager: - await self.mcp_manager.cleanup_all() + # 清理工具处理器资源 + if hasattr(self, "func_handler") and self.func_handler: + await self.func_handler.cleanup() # 触发停止事件 if self.stop_event: diff --git a/main/xiaozhi-server/core/handle/functionHandler.py b/main/xiaozhi-server/core/handle/functionHandler.py deleted file mode 100644 index e1292553..00000000 --- a/main/xiaozhi-server/core/handle/functionHandler.py +++ /dev/null @@ -1,103 +0,0 @@ -from config.logger import setup_logging -import json -from plugins_func.register import ( - FunctionRegistry, - ActionResponse, - Action, - ToolType, - DeviceTypeRegistry, -) -from plugins_func.functions.hass_init import append_devices_to_prompt - -TAG = __name__ - - -class FunctionHandler: - def __init__(self, conn): - self.conn = conn - self.config = conn.config - self.device_type_registry = DeviceTypeRegistry() - self.function_registry = FunctionRegistry() - self.register_nessary_functions() - self.register_config_functions() - self.functions_desc = self.function_registry.get_all_function_desc() - self.finish_init = True - - def upload_functions_desc(self): - self.functions_desc = self.function_registry.get_all_function_desc() - - - def current_support_functions(self): - func_names = [] - for func in self.functions_desc: - func_names.append(func["function"]["name"]) - # 打印当前支持的函数列表 - self.conn.logger.bind(tag=TAG, session_id=self.conn.session_id).info( - f"当前支持的函数列表: {func_names}" - ) - return func_names - - def get_functions(self): - """获取功能调用配置""" - return self.functions_desc - - def register_nessary_functions(self): - """注册必要的函数""" - self.function_registry.register_function("handle_exit_intent") - self.function_registry.register_function("get_time") - self.function_registry.register_function("get_lunar") - - def register_config_functions(self): - """注册配置中的函数,可以不同客户端使用不同的配置""" - for func in self.config["Intent"][self.config["selected_module"]["Intent"]].get( - "functions", [] - ): - self.function_registry.register_function(func) - - """home assistant需要初始化提示词""" - append_devices_to_prompt(self.conn) - - def get_function(self, name): - return self.function_registry.get_function(name) - - def handle_llm_function_call(self, conn, function_call_data): - # 多函数调用处理 - if "function_calls" in function_call_data: - responses = [] - for call in function_call_data["function_calls"]: - func = self.get_function(call["name"]) - if func: - # 执行函数并收集响应 - response = func(conn, **call.get("arguments", {})) - responses.append(response) - return self._combine_responses(responses) # 合并响应 - try: - function_name = function_call_data["name"] - funcItem = self.get_function(function_name) - if not funcItem: - return ActionResponse( - action=Action.NOTFOUND, result="没有找到对应的函数", response="" - ) - func = funcItem.func - arguments = function_call_data["arguments"] - arguments = json.loads(arguments) if arguments else {} - self.conn.logger.bind(tag=TAG).debug( - f"调用函数: {function_name}, 参数: {arguments}" - ) - if ( - funcItem.type == ToolType.SYSTEM_CTL - or funcItem.type == ToolType.IOT_CTL - ): - return func(conn, **arguments) - elif funcItem.type == ToolType.WAIT: - return func(**arguments) - elif funcItem.type == ToolType.CHANGE_SYS_PROMPT: - return func(conn, **arguments) - else: - return ActionResponse( - action=Action.NOTFOUND, result="没有找到对应的函数", response="" - ) - except Exception as e: - self.conn.logger.bind(tag=TAG).error(f"处理function call错误: {e}") - - return None diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py index f5a3b0cd..e4e836b4 100644 --- a/main/xiaozhi-server/core/handle/helloHandle.py +++ b/main/xiaozhi-server/core/handle/helloHandle.py @@ -7,7 +7,7 @@ from core.utils.util import audio_to_data from core.handle.sendAudioHandle import sendAudioMessage, send_stt_message from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes from core.providers.tts.dto.dto import ContentType, SentenceType -from core.handle.mcpHandle import ( +from core.providers.tools.device_mcp import ( MCPClient, send_mcp_initialize_message, send_mcp_tools_list_request, diff --git a/main/xiaozhi-server/core/handle/intentHandler.py b/main/xiaozhi-server/core/handle/intentHandler.py index de4c36c1..2d0baca4 100644 --- a/main/xiaozhi-server/core/handle/intentHandler.py +++ b/main/xiaozhi-server/core/handle/intentHandler.py @@ -6,7 +6,7 @@ from core.handle.helloHandle import checkWakeupWords from core.utils.util import remove_punctuation_and_length from core.providers.tts.dto.dto import ContentType from core.utils.dialogue import Message -from core.handle.mcpHandle import call_mcp_tool +from core.providers.tools.device_mcp import call_mcp_tool from plugins_func.register import Action, ActionResponse from loguru import logger @@ -106,36 +106,18 @@ async def process_intent_result(conn, intent_result, original_text): def process_function_call(): conn.dialogue.put(Message(role="user", content=original_text)) - # 处理Server端MCP工具调用 - if conn.mcp_manager.is_mcp_tool(function_name): - result = conn._handle_mcp_tool_call(function_call_data) - elif hasattr(conn, "mcp_client") and conn.mcp_client.has_tool( - function_name - ): - # 如果是小智端MCP工具调用 - conn.logger.bind(tag=TAG).debug( - f"调用小智端MCP工具: {function_name}, 参数: {function_args}" - ) - try: - result = asyncio.run_coroutine_threadsafe( - call_mcp_tool( - conn, conn.mcp_client, function_name, function_args - ), - conn.loop, - ).result() - conn.logger.bind(tag=TAG).debug(f"MCP工具调用结果: {result}") - result = ActionResponse( - action=Action.REQLLM, result=result, response="" - ) - except Exception as e: - conn.logger.bind(tag=TAG).error(f"MCP工具调用失败: {e}") - result = ActionResponse( - action=Action.REQLLM, result="MCP工具调用失败", response="" - ) - else: - # 处理系统函数 - result = conn.func_handler.handle_llm_function_call( - conn, function_call_data + # 使用统一工具处理器处理所有工具调用 + try: + result = asyncio.run_coroutine_threadsafe( + conn.func_handler.handle_llm_function_call( + conn, function_call_data + ), + conn.loop, + ).result() + except Exception as e: + conn.logger.bind(tag=TAG).error(f"工具调用失败: {e}") + result = ActionResponse( + action=Action.ERROR, result=str(e), response=str(e) ) if result: diff --git a/main/xiaozhi-server/core/handle/iotHandle.py b/main/xiaozhi-server/core/handle/iotHandle.py deleted file mode 100644 index bebb44fc..00000000 --- a/main/xiaozhi-server/core/handle/iotHandle.py +++ /dev/null @@ -1,427 +0,0 @@ -import json -import asyncio -from plugins_func.register import ( - FunctionItem, - register_device_function, - ActionResponse, - Action, - ToolType, -) - -TAG = __name__ - - -def wrap_async_function(async_func): - """包装异步函数为同步函数""" - - def wrapper(*args, **kwargs): - try: - # 获取连接对象(第一个参数) - conn = args[0] - if not hasattr(conn, "loop"): - conn.logger.bind(tag=TAG).error("Connection对象没有loop属性") - return ActionResponse( - Action.ERROR, - "Connection对象没有loop属性", - "执行操作时出错: Connection对象没有loop属性", - ) - - # 使用conn对象中的事件循环 - loop = conn.loop - # 在conn的事件循环中运行异步函数 - future = asyncio.run_coroutine_threadsafe(async_func(*args, **kwargs), loop) - # 等待结果返回 - return future.result() - except Exception as e: - conn.logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}") - return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}") - - return wrapper - - -def create_iot_function(device_name, method_name, method_info): - """ - 根据IOT设备描述生成通用的控制函数 - """ - - async def iot_control_function( - conn, response_success=None, response_failure=None, **params - ): - try: - # 设置默认响应消息 - if not response_success: - response_success = "操作成功" - if not response_failure: - response_failure = "操作失败" - - # 打印响应参数 - conn.logger.bind(tag=TAG).debug( - f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'" - ) - - # 发送控制命令 - await send_iot_conn(conn, device_name, method_name, params) - # 等待一小段时间让状态更新 - await asyncio.sleep(0.1) - - # 生成结果信息 - result = f"{device_name}的{method_name}操作执行成功" - - # 处理响应中可能的占位符 - response = response_success - # 替换{value}占位符 - for param_name, param_value in params.items(): - # 先尝试直接替换参数值 - if "{" + param_name + "}" in response: - response = response.replace( - "{" + param_name + "}", str(param_value) - ) - - # 如果有{value}占位符,用相关参数替换 - if "{value}" in response: - response = response.replace("{value}", str(param_value)) - break - - return ActionResponse( - Action.REQLLM, - result=f"{device_name}操作执行成功,请继续处理剩余指令", - response=response_success # 保留成功提示 - ) - except Exception as e: - conn.logger.bind(tag=TAG).error( - f"执行{device_name}的{method_name}操作失败: {e}" - ) - - # 操作失败时使用大模型提供的失败响应 - response = response_failure - - return ActionResponse(Action.ERROR, str(e), response) - - return wrap_async_function(iot_control_function) - - -def create_iot_query_function(device_name, prop_name, prop_info): - """ - 根据IOT设备属性创建查询函数 - """ - - async def iot_query_function(conn, response_success=None, response_failure=None): - try: - # 打印响应参数 - conn.logger.bind(tag=TAG).info( - f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'" - ) - - value = await get_iot_status(conn, device_name, prop_name) - - # 查询成功,生成结果 - if value is not None: - # 使用大模型提供的成功响应,并替换其中的占位符 - response = response_success.replace("{value}", str(value)) - - return ActionResponse(Action.RESPONSE, str(value), response) - else: - # 查询失败,使用大模型提供的失败响应 - response = response_failure - - return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response) - except Exception as e: - conn.logger.bind(tag=TAG).error( - f"查询{device_name}的{prop_name}时出错: {e}" - ) - - # 查询出错时使用大模型提供的失败响应 - response = response_failure - - return ActionResponse(Action.ERROR, str(e), response) - - return wrap_async_function(iot_query_function) - - -class IotDescriptor: - """ - A class to represent an IoT descriptor. - """ - - def __init__(self, name, description, properties, methods): - self.name = name - self.description = description - self.properties = [] - self.methods = [] - - # 根据描述创建属性 - if properties is not None: - for key, value in properties.items(): - property_item = {} - property_item["name"] = key - property_item["description"] = value["description"] - if value["type"] == "number": - property_item["value"] = 0 - elif value["type"] == "boolean": - property_item["value"] = False - else: - property_item["value"] = "" - self.properties.append(property_item) - - # 根据描述创建方法 - if methods is not None: - for key, value in methods.items(): - method = {} - method["description"] = value["description"] - method["name"] = key - # 检查方法是否有参数 - if "parameters" in value: - method["parameters"] = {} - for k, v in value["parameters"].items(): - method["parameters"][k] = { - "description": v["description"], - "type": v["type"], - } - self.methods.append(method) - - -def register_device_type(descriptor, device_type_registry): - """注册设备类型及其功能""" - device_name = descriptor["name"] - type_id = device_type_registry.generate_device_type_id(descriptor) - - # 如果该类型已注册,直接返回类型ID - if type_id in device_type_registry.type_functions: - return type_id - - functions = {} - - # 为每个属性创建查询函数 - for prop_name, prop_info in descriptor["properties"].items(): - func_name = f"get_{device_name.lower()}_{prop_name.lower()}" - func_desc = { - "type": "function", - "function": { - "name": func_name, - "description": f"查询{descriptor['description']}的{prop_info['description']}", - "parameters": { - "type": "object", - "properties": { - "response_success": { - "type": "string", - "description": f"查询成功时的友好回复,必须使用{{value}}作为占位符表示查询到的值", - }, - "response_failure": { - "type": "string", - "description": f"查询失败时的友好回复,例如:'无法获取{device_name}的{prop_info['description']}'", - }, - }, - "required": ["response_success", "response_failure"], - }, - }, - } - query_func = create_iot_query_function(device_name, prop_name, prop_info) - decorated_func = register_device_function( - func_name, func_desc, ToolType.IOT_CTL - )(query_func) - functions[func_name] = FunctionItem( - func_name, func_desc, decorated_func, ToolType.IOT_CTL - ) - - # 为每个方法创建控制函数 - for method_name, method_info in descriptor["methods"].items(): - func_name = f"{device_name.lower()}_{method_name.lower()}" - - # 创建参数字典,添加原有参数 - parameters = {} - required_params = [] - - # 如果方法有参数,则添加参数信息 - if "parameters" in method_info: - parameters = { - param_name: { - "type": param_info["type"], - "description": param_info["description"], - } - for param_name, param_info in method_info["parameters"].items() - } - required_params = list(method_info["parameters"].keys()) - - # 添加响应参数 - parameters.update( - { - "response_success": { - "type": "string", - "description": "操作成功时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称", - }, - "response_failure": { - "type": "string", - "description": "操作失败时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称", - }, - } - ) - - # 构建必须参数列表(原有参数 + 响应参数) - required_params.extend(["response_success", "response_failure"]) - - func_desc = { - "type": "function", - "function": { - "name": func_name, - "description": f"{descriptor['description']} - {method_info['description']}", - "parameters": { - "type": "object", - "properties": parameters, - "required": required_params, - }, - }, - } - control_func = create_iot_function(device_name, method_name, method_info) - decorated_func = register_device_function( - func_name, func_desc, ToolType.IOT_CTL - )(control_func) - functions[func_name] = FunctionItem( - func_name, func_desc, decorated_func, ToolType.IOT_CTL - ) - - device_type_registry.register_device_type(type_id, functions) - return type_id - - -# 用于接受前端设备推送的搜索iot描述 -async def handleIotDescriptors(conn, descriptors): - wait_max_time = 5 - while conn.func_handler is None or not conn.func_handler.finish_init: - await asyncio.sleep(1) - wait_max_time -= 1 - if wait_max_time <= 0: - conn.logger.bind(tag=TAG).debug("连接对象没有func_handler") - return - """处理物联网描述""" - functions_changed = False - - for descriptor in descriptors: - # 如果descriptor没有properties和methods,则直接跳过 - if "properties" not in descriptor and "methods" not in descriptor: - continue - - # 处理缺失properties的情况 - if "properties" not in descriptor: - descriptor["properties"] = {} - # 从methods中提取所有参数作为properties - if "methods" in descriptor: - for method_name, method_info in descriptor["methods"].items(): - if "parameters" in method_info: - for param_name, param_info in method_info["parameters"].items(): - # 将参数信息转换为属性信息 - descriptor["properties"][param_name] = { - "description": param_info["description"], - "type": param_info["type"], - } - - # 创建IOT设备描述符 - iot_descriptor = IotDescriptor( - descriptor["name"], - descriptor["description"], - descriptor["properties"], - descriptor["methods"], - ) - conn.iot_descriptors[descriptor["name"]] = iot_descriptor - - if conn.load_function_plugin: - # 注册或获取设备类型 - device_type_registry = conn.func_handler.device_type_registry - type_id = register_device_type(descriptor, device_type_registry) - device_functions = device_type_registry.get_device_functions(type_id) - - # 在连接级注册设备函数 - if hasattr(conn, "func_handler"): - for func_name, func_item in device_functions.items(): - conn.func_handler.function_registry.register_function( - func_name, func_item - ) - conn.logger.bind(tag=TAG).info( - f"注册IOT函数到function handler: {func_name}" - ) - functions_changed = True - - # 如果注册了新函数,更新function描述列表 - if functions_changed and hasattr(conn, "func_handler"): - conn.func_handler.upload_functions_desc() - - func_names = conn.func_handler.current_support_functions() - conn.logger.bind(tag=TAG).info(f"设备类型: {type_id}") - conn.logger.bind(tag=TAG).info( - f"更新function描述列表完成,当前支持的函数: {func_names}" - ) - - -async def handleIotStatus(conn, states): - """处理物联网状态""" - for state in states: - for key, value in conn.iot_descriptors.items(): - if key == state["name"]: - for property_item in value.properties: - for k, v in state["state"].items(): - if property_item["name"] == k: - if type(v) != type(property_item["value"]): - conn.logger.bind(tag=TAG).error( - f"属性{property_item['name']}的值类型不匹配" - ) - break - else: - property_item["value"] = v - conn.logger.bind(tag=TAG).info( - f"物联网状态更新: {key} , {property_item['name']} = {v}" - ) - break - break - - -async def get_iot_status(conn, name, property_name): - """获取物联网状态""" - for key, value in conn.iot_descriptors.items(): - if key == name: - for property_item in value.properties: - if property_item["name"] == property_name: - return property_item["value"] - conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}") - return None - - -async def set_iot_status(conn, name, property_name, value): - """设置物联网状态""" - for key, iot_descriptor in conn.iot_descriptors.items(): - if key == name: - for property_item in iot_descriptor.properties: - if property_item["name"] == property_name: - if type(value) != type(property_item["value"]): - conn.logger.bind(tag=TAG).error( - f"属性{property_item['name']}的值类型不匹配" - ) - return - property_item["value"] = value - conn.logger.bind(tag=TAG).info( - f"物联网状态更新: {name} , {property_name} = {value}" - ) - return - conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}") - - -async def send_iot_conn(conn, name, method_name, parameters): - """发送物联网指令""" - for key, value in conn.iot_descriptors.items(): - if key == name: - # 找到了设备 - for method in value.methods: - # 找到了方法 - if method["name"] == method_name: - # 构建命令对象 - command = { - "name": name, - "method": method_name, - } - - # 只有当参数不为空时才添加parameters字段 - if parameters: - command["parameters"] = parameters - send_message = json.dumps({"type": "iot", "commands": [command]}) - await conn.websocket.send(send_message) - conn.logger.bind(tag=TAG).info(f"发送物联网指令: {send_message}") - return - conn.logger.bind(tag=TAG).error(f"未找到方法{method_name}") diff --git a/main/xiaozhi-server/core/handle/textHandle.py b/main/xiaozhi-server/core/handle/textHandle.py index 676e33ab..cdb20ac5 100644 --- a/main/xiaozhi-server/core/handle/textHandle.py +++ b/main/xiaozhi-server/core/handle/textHandle.py @@ -1,11 +1,11 @@ import json from core.handle.abortHandle import handleAbortMessage from core.handle.helloHandle import handleHelloMessage -from core.handle.mcpHandle import handle_mcp_message +from core.providers.tools.device_mcp 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 -from core.handle.iotHandle import handleIotDescriptors, handleIotStatus +from core.providers.tools.device_iot import handleIotDescriptors, handleIotStatus from core.handle.reportHandle import enqueue_asr_report import asyncio @@ -77,7 +77,7 @@ async def handleTextMessage(conn, message): 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}") + conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message[:100]}") if "payload" in msg_json: asyncio.create_task( handle_mcp_message(conn, conn.mcp_client, msg_json["payload"]) diff --git a/main/xiaozhi-server/core/providers/tools/__init__.py b/main/xiaozhi-server/core/providers/tools/__init__.py new file mode 100644 index 00000000..0519ecba --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/__init__.py @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/main/xiaozhi-server/core/providers/tools/base/__init__.py b/main/xiaozhi-server/core/providers/tools/base/__init__.py new file mode 100644 index 00000000..7476ddd5 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/base/__init__.py @@ -0,0 +1,6 @@ +"""基础工具定义模块""" + +from .tool_types import ToolType, ToolDefinition +from .tool_executor import ToolExecutor + +__all__ = ["ToolType", "ToolDefinition", "ToolExecutor"] diff --git a/main/xiaozhi-server/core/providers/tools/base/tool_executor.py b/main/xiaozhi-server/core/providers/tools/base/tool_executor.py new file mode 100644 index 00000000..7685e7c2 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/base/tool_executor.py @@ -0,0 +1,27 @@ +"""工具执行器基类定义""" + +from abc import ABC, abstractmethod +from typing import Dict, Any +from .tool_types import ToolDefinition +from plugins_func.register import ActionResponse + + +class ToolExecutor(ABC): + """工具执行器抽象基类""" + + @abstractmethod + async def execute( + self, conn, tool_name: str, arguments: Dict[str, Any] + ) -> ActionResponse: + """执行工具调用""" + pass + + @abstractmethod + def get_tools(self) -> Dict[str, ToolDefinition]: + """获取该执行器管理的所有工具""" + pass + + @abstractmethod + def has_tool(self, tool_name: str) -> bool: + """检查是否有指定工具""" + pass diff --git a/main/xiaozhi-server/core/providers/tools/base/tool_types.py b/main/xiaozhi-server/core/providers/tools/base/tool_types.py new file mode 100644 index 00000000..91466fad --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/base/tool_types.py @@ -0,0 +1,27 @@ +"""工具系统的类型定义""" + +from enum import Enum + +from dataclasses import dataclass +from typing import Any, Dict, Optional +from plugins_func.register import Action + + +class ToolType(Enum): + """工具类型枚举""" + + SERVER_PLUGIN = "server_plugin" # 服务端插件 + SERVER_MCP = "server_mcp" # 服务端MCP + DEVICE_IOT = "device_iot" # 设备端IoT + DEVICE_MCP = "device_mcp" # 设备端MCP + MCP_ENDPOINT = "mcp_endpoint" # MCP接入点 + + +@dataclass +class ToolDefinition: + """工具定义""" + + name: str # 工具名称 + description: Dict[str, Any] # 工具描述(OpenAI函数调用格式) + tool_type: ToolType # 工具类型 + parameters: Optional[Dict[str, Any]] = None # 额外参数 diff --git a/main/xiaozhi-server/core/providers/tools/device_iot/__init__.py b/main/xiaozhi-server/core/providers/tools/device_iot/__init__.py new file mode 100644 index 00000000..844c5c28 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/device_iot/__init__.py @@ -0,0 +1,12 @@ +"""设备端IoT工具模块""" + +from .iot_descriptor import IotDescriptor +from .iot_handler import handleIotDescriptors, handleIotStatus +from .iot_executor import DeviceIoTExecutor + +__all__ = [ + "IotDescriptor", + "handleIotDescriptors", + "handleIotStatus", + "DeviceIoTExecutor", +] diff --git a/main/xiaozhi-server/core/providers/tools/device_iot/iot_descriptor.py b/main/xiaozhi-server/core/providers/tools/device_iot/iot_descriptor.py new file mode 100644 index 00000000..81df02ad --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/device_iot/iot_descriptor.py @@ -0,0 +1,46 @@ +"""IoT设备描述符定义""" + +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + + +class IotDescriptor: + """IoT设备描述符""" + + def __init__(self, name, description, properties, methods): + self.name = name + self.description = description + self.properties = [] + self.methods = [] + + # 根据描述创建属性 + if properties is not None: + for key, value in properties.items(): + property_item = {} + property_item["name"] = key + property_item["description"] = value["description"] + if value["type"] == "number": + property_item["value"] = 0 + elif value["type"] == "boolean": + property_item["value"] = False + else: + property_item["value"] = "" + self.properties.append(property_item) + + # 根据描述创建方法 + if methods is not None: + for key, value in methods.items(): + method = {} + method["description"] = value["description"] + method["name"] = key + # 检查方法是否有参数 + if "parameters" in value: + method["parameters"] = {} + for k, v in value["parameters"].items(): + method["parameters"][k] = { + "description": v["description"], + "type": v["type"], + } + self.methods.append(method) diff --git a/main/xiaozhi-server/core/providers/tools/device_iot/iot_executor.py b/main/xiaozhi-server/core/providers/tools/device_iot/iot_executor.py new file mode 100644 index 00000000..7be9e112 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/device_iot/iot_executor.py @@ -0,0 +1,238 @@ +"""设备端IoT工具执行器""" + +import json +import asyncio +from typing import Dict, Any +from ..base import ToolType, ToolDefinition, ToolExecutor +from plugins_func.register import Action, ActionResponse + + +class DeviceIoTExecutor(ToolExecutor): + """设备端IoT工具执行器""" + + def __init__(self, conn): + self.conn = conn + self.iot_tools: Dict[str, ToolDefinition] = {} + + async def execute( + self, conn, tool_name: str, arguments: Dict[str, Any] + ) -> ActionResponse: + """执行设备端IoT工具""" + if not self.has_tool(tool_name): + return ActionResponse( + action=Action.NOTFOUND, response=f"IoT工具 {tool_name} 不存在" + ) + + try: + # 解析工具名称,获取设备名和操作类型 + if tool_name.startswith("get_"): + # 查询操作:get_devicename_property + parts = tool_name.split("_", 2) + if len(parts) >= 3: + device_name = parts[1] + property_name = parts[2] + + value = await self._get_iot_status(device_name, property_name) + if value is not None: + # 处理响应模板 + response_success = arguments.get( + "response_success", "查询成功:{value}" + ) + response = response_success.replace("{value}", str(value)) + + return ActionResponse( + action=Action.RESPONSE, + response=response, + ) + else: + response_failure = arguments.get( + "response_failure", f"无法获取{device_name}的状态" + ) + return ActionResponse( + action=Action.ERROR, response=response_failure + ) + else: + # 控制操作:devicename_method + parts = tool_name.split("_", 1) + if len(parts) >= 2: + device_name = parts[0] + method_name = parts[1] + + # 提取控制参数(排除响应参数) + control_params = { + k: v + for k, v in arguments.items() + if k not in ["response_success", "response_failure"] + } + + # 发送IoT控制命令 + await self._send_iot_command( + device_name, method_name, control_params + ) + + # 等待状态更新 + await asyncio.sleep(0.1) + + response_success = arguments.get("response_success", "操作成功") + + # 处理响应中的占位符 + for param_name, param_value in control_params.items(): + placeholder = "{" + param_name + "}" + if placeholder in response_success: + response_success = response_success.replace( + placeholder, str(param_value) + ) + if "{value}" in response_success: + response_success = response_success.replace( + "{value}", str(param_value) + ) + break + + return ActionResponse( + action=Action.REQLLM, + result=response_success, + ) + + return ActionResponse(action=Action.ERROR, response="无法解析IoT工具名称") + + except Exception as e: + response_failure = arguments.get("response_failure", "操作失败") + return ActionResponse(action=Action.ERROR, response=response_failure) + + async def _get_iot_status(self, device_name: str, property_name: str): + """获取IoT设备状态""" + for key, value in self.conn.iot_descriptors.items(): + if key.lower() == device_name.lower(): + for property_item in value.properties: + if property_item["name"].lower() == property_name.lower(): + return property_item["value"] + return None + + async def _send_iot_command( + self, device_name: str, method_name: str, parameters: Dict[str, Any] + ): + """发送IoT控制命令""" + for key, value in self.conn.iot_descriptors.items(): + if key.lower() == device_name.lower(): + for method in value.methods: + if method["name"].lower() == method_name.lower(): + command = { + "name": key, + "method": method["name"], + } + + if parameters: + command["parameters"] = parameters + + send_message = json.dumps( + {"type": "iot", "commands": [command]} + ) + await self.conn.websocket.send(send_message) + return + + raise Exception(f"未找到设备{device_name}的方法{method_name}") + + def register_iot_tools(self, descriptors: list): + """注册IoT工具""" + for descriptor in descriptors: + device_name = descriptor["name"] + device_desc = descriptor["description"] + + # 注册查询工具 + if "properties" in descriptor: + for prop_name, prop_info in descriptor["properties"].items(): + tool_name = f"get_{device_name.lower()}_{prop_name.lower()}" + + tool_desc = { + "type": "function", + "function": { + "name": tool_name, + "description": f"查询{device_desc}的{prop_info['description']}", + "parameters": { + "type": "object", + "properties": { + "response_success": { + "type": "string", + "description": f"查询成功时的友好回复,必须使用{{value}}作为占位符表示查询到的值", + }, + "response_failure": { + "type": "string", + "description": f"查询失败时的友好回复", + }, + }, + "required": ["response_success", "response_failure"], + }, + }, + } + + self.iot_tools[tool_name] = ToolDefinition( + name=tool_name, + description=tool_desc, + tool_type=ToolType.DEVICE_IOT, + ) + + # 注册控制工具 + if "methods" in descriptor: + for method_name, method_info in descriptor["methods"].items(): + tool_name = f"{device_name.lower()}_{method_name.lower()}" + + # 构建参数 + parameters = {} + required_params = [] + + # 添加方法的原始参数 + if "parameters" in method_info: + parameters.update( + { + param_name: { + "type": param_info["type"], + "description": param_info["description"], + } + for param_name, param_info in method_info[ + "parameters" + ].items() + } + ) + required_params.extend(method_info["parameters"].keys()) + + # 添加响应参数 + parameters.update( + { + "response_success": { + "type": "string", + "description": "操作成功时的友好回复", + }, + "response_failure": { + "type": "string", + "description": "操作失败时的友好回复", + }, + } + ) + required_params.extend(["response_success", "response_failure"]) + + tool_desc = { + "type": "function", + "function": { + "name": tool_name, + "description": f"{device_desc} - {method_info['description']}", + "parameters": { + "type": "object", + "properties": parameters, + "required": required_params, + }, + }, + } + + self.iot_tools[tool_name] = ToolDefinition( + name=tool_name, + description=tool_desc, + tool_type=ToolType.DEVICE_IOT, + ) + + def get_tools(self) -> Dict[str, ToolDefinition]: + """获取所有设备端IoT工具""" + return self.iot_tools.copy() + + def has_tool(self, tool_name: str) -> bool: + """检查是否有指定的设备端IoT工具""" + return tool_name in self.iot_tools diff --git a/main/xiaozhi-server/core/providers/tools/device_iot/iot_handler.py b/main/xiaozhi-server/core/providers/tools/device_iot/iot_handler.py new file mode 100644 index 00000000..ffea57a4 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/device_iot/iot_handler.py @@ -0,0 +1,86 @@ +"""IoT设备支持模块,提供IoT设备描述符和状态处理""" + +import asyncio +from config.logger import setup_logging +from .iot_descriptor import IotDescriptor + +TAG = __name__ +logger = setup_logging() + + +async def handleIotDescriptors(conn, descriptors): + """处理物联网描述""" + wait_max_time = 5 + while ( + not hasattr(conn, "func_handler") + or conn.func_handler is None + or not conn.func_handler.finish_init + ): + await asyncio.sleep(1) + wait_max_time -= 1 + if wait_max_time <= 0: + logger.bind(tag=TAG).debug("连接对象没有func_handler") + return + + functions_changed = False + + for descriptor in descriptors: + # 如果descriptor没有properties和methods,则直接跳过 + if "properties" not in descriptor and "methods" not in descriptor: + continue + + # 处理缺失properties的情况 + if "properties" not in descriptor: + descriptor["properties"] = {} + # 从methods中提取所有参数作为properties + if "methods" in descriptor: + for method_name, method_info in descriptor["methods"].items(): + if "parameters" in method_info: + for param_name, param_info in method_info["parameters"].items(): + # 将参数信息转换为属性信息 + descriptor["properties"][param_name] = { + "description": param_info["description"], + "type": param_info["type"], + } + + # 创建IOT设备描述符 + iot_descriptor = IotDescriptor( + descriptor["name"], + descriptor["description"], + descriptor["properties"], + descriptor["methods"], + ) + conn.iot_descriptors[descriptor["name"]] = iot_descriptor + functions_changed = True + + # 如果注册了新函数,更新function描述列表 + if functions_changed and hasattr(conn, "func_handler"): + # 注册IoT工具到统一工具处理器 + await conn.func_handler.register_iot_tools(descriptors) + + func_names = conn.func_handler.current_support_functions() + logger.bind(tag=TAG).info( + f"更新function描述列表完成,当前支持的函数: {func_names}" + ) + + +async def handleIotStatus(conn, states): + """处理物联网状态""" + for state in states: + for key, value in conn.iot_descriptors.items(): + if key == state["name"]: + for property_item in value.properties: + for k, v in state["state"].items(): + if property_item["name"] == k: + if type(v) != type(property_item["value"]): + logger.bind(tag=TAG).error( + f"属性{property_item['name']}的值类型不匹配" + ) + break + else: + property_item["value"] = v + logger.bind(tag=TAG).info( + f"物联网状态更新: {key} , {property_item['name']} = {v}" + ) + break + break diff --git a/main/xiaozhi-server/core/providers/tools/device_mcp/__init__.py b/main/xiaozhi-server/core/providers/tools/device_mcp/__init__.py new file mode 100644 index 00000000..266d60c8 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/device_mcp/__init__.py @@ -0,0 +1,21 @@ +"""设备端MCP工具模块""" + +from .mcp_client import MCPClient +from .mcp_handler import ( + send_mcp_message, + handle_mcp_message, + send_mcp_initialize_message, + send_mcp_tools_list_request, + call_mcp_tool, +) +from .mcp_executor import DeviceMCPExecutor + +__all__ = [ + "MCPClient", + "send_mcp_message", + "handle_mcp_message", + "send_mcp_initialize_message", + "send_mcp_tools_list_request", + "call_mcp_tool", + "DeviceMCPExecutor", +] diff --git a/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_client.py b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_client.py new file mode 100644 index 00000000..75aa22b5 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_client.py @@ -0,0 +1,93 @@ +"""设备端MCP客户端定义""" + +import asyncio +from concurrent.futures import Future +from core.utils.util import sanitize_tool_name +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + + +class MCPClient: + """设备端MCP客户端,用于管理MCP状态和工具""" + + def __init__(self): + self.tools = {} # sanitized_name -> tool_data + self.name_mapping = {} + 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_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: + sanitized_name = sanitize_tool_name(tool_data["name"]) + self.tools[sanitized_name] = tool_data + self.name_mapping[sanitized_name] = tool_data["name"] + 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) diff --git a/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_executor.py b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_executor.py new file mode 100644 index 00000000..9f4be061 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_executor.py @@ -0,0 +1,89 @@ +"""设备端MCP工具执行器""" + +from typing import Dict, Any +from ..base import ToolType, ToolDefinition, ToolExecutor +from plugins_func.register import Action, ActionResponse +from .mcp_handler import call_mcp_tool + + +class DeviceMCPExecutor(ToolExecutor): + """设备端MCP工具执行器""" + + def __init__(self, conn): + self.conn = conn + + async def execute( + self, conn, tool_name: str, arguments: Dict[str, Any] + ) -> ActionResponse: + """执行设备端MCP工具""" + if not hasattr(conn, "mcp_client") or not conn.mcp_client: + return ActionResponse( + action=Action.ERROR, + response="设备端MCP客户端未初始化", + ) + + if not await conn.mcp_client.is_ready(): + return ActionResponse( + action=Action.ERROR, + response="设备端MCP客户端未准备就绪", + ) + + try: + # 转换参数为JSON字符串 + import json + + args_str = json.dumps(arguments) if arguments else "{}" + + # 调用设备端MCP工具 + result = await call_mcp_tool(conn, conn.mcp_client, tool_name, args_str) + + resultJson = None + if isinstance(result, str): + try: + resultJson = json.loads(result) + except Exception as e: + pass + + # 视觉大模型不经过二次LLM处理 + if ( + resultJson is not None + and isinstance(resultJson, dict) + and "action" in resultJson + ): + return ActionResponse( + action=Action[resultJson["action"]], + response=resultJson.get("response", ""), + ) + + return ActionResponse(action=Action.REQLLM, result=str(result)) + + except ValueError as e: + return ActionResponse(action=Action.NOTFOUND, response=str(e)) + except Exception as e: + return ActionResponse(action=Action.ERROR, response=str(e)) + + def get_tools(self) -> Dict[str, ToolDefinition]: + """获取所有设备端MCP工具""" + if not hasattr(self.conn, "mcp_client") or not self.conn.mcp_client: + return {} + + tools = {} + mcp_tools = self.conn.mcp_client.get_available_tools() + + for tool in mcp_tools: + func_def = tool.get("function", {}) + tool_name = func_def.get("name", "") + + if tool_name: + tools[tool_name] = ToolDefinition( + name=tool_name, description=tool, tool_type=ToolType.DEVICE_MCP + ) + + return tools + + def has_tool(self, tool_name: str) -> bool: + """检查是否有指定的设备端MCP工具""" + if not hasattr(self.conn, "mcp_client") or not self.conn.mcp_client: + return False + + return self.conn.mcp_client.has_tool(tool_name) diff --git a/main/xiaozhi-server/core/handle/mcpHandle.py b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_handler.py similarity index 87% rename from main/xiaozhi-server/core/handle/mcpHandle.py rename to main/xiaozhi-server/core/providers/tools/device_mcp/mcp_handler.py index 3f3216ac..d5f42833 100644 --- a/main/xiaozhi-server/core/handle/mcpHandle.py +++ b/main/xiaozhi-server/core/providers/tools/device_mcp/mcp_handler.py @@ -1,14 +1,19 @@ +"""设备端MCP客户端支持模块""" + import json import asyncio +import re from concurrent.futures import Future from core.utils.util import get_vision_url, sanitize_tool_name from core.utils.auth import AuthToken +from config.logger import setup_logging TAG = __name__ +logger = setup_logging() class MCPClient: - """MCPClient,用于管理MCP状态和工具""" + """设备端MCP客户端,用于管理MCP状态和工具""" def __init__(self): self.tools = {} # sanitized_name -> tool_data @@ -94,24 +99,24 @@ class MCPClient: 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消息") + 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}") + logger.bind(tag=TAG).info(f"成功发送MCP消息: {message}") except Exception as e: - conn.logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}") + 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}") + logger.bind(tag=TAG).info(f"处理MCP消息: {str(payload)[:100]}") if not isinstance(payload, dict): - conn.logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误") + logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误") return # Handle result @@ -121,32 +126,32 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict): # Check for tool call response first if msg_id in mcp_client.call_results: - conn.logger.bind(tag=TAG).debug( + 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初始化响应") + 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( + logger.bind(tag=TAG).info( f"客户端MCP服务器信息: name={name}, version={version}" ) return elif msg_id == 2: # mcpToolsListID - conn.logger.bind(tag=TAG).debug("收到MCP工具列表响应") + 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("工具列表格式错误") + logger.bind(tag=TAG).error("工具列表格式错误") return - conn.logger.bind(tag=TAG).info( + logger.bind(tag=TAG).info( f"客户端设备支持的工具数量: {len(tools_data)}" ) @@ -172,7 +177,7 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict): "inputSchema": input_schema, } await mcp_client.add_tool(new_tool) - conn.logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}") + logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}") # 替换所有工具描述中的工具名称 for tool_data in mcp_client.tools.values(): @@ -190,24 +195,22 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict): next_cursor = result.get("nextCursor", "") if next_cursor: - conn.logger.bind(tag=TAG).info( - f"有更多工具,nextCursor: {next_cursor}" - ) + 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客户端准备就绪") + 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}") + 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}") + logger.bind(tag=TAG).error(f"收到MCP错误响应: {error_msg}") msg_id = int(payload.get("id", 0)) if msg_id in mcp_client.call_results: @@ -216,9 +219,6 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict): ) -# --- Outgoing MCP Messages --- - - async def send_mcp_initialize_message(conn): """发送MCP初始化消息""" @@ -250,7 +250,7 @@ async def send_mcp_initialize_message(conn): }, }, } - conn.logger.bind(tag=TAG).info("发送MCP初始化消息") + logger.bind(tag=TAG).info("发送MCP初始化消息") await send_mcp_message(conn, payload) @@ -261,7 +261,7 @@ async def send_mcp_tools_list_request(conn): "id": 2, # mcpToolsListID "method": "tools/list", } - conn.logger.bind(tag=TAG).debug("发送MCP工具列表请求") + logger.bind(tag=TAG).debug("发送MCP工具列表请求") await send_mcp_message(conn, payload) @@ -273,7 +273,7 @@ async def send_mcp_tools_list_continue_request(conn, cursor: str): "method": "tools/list", "params": {"cursor": cursor}, } - conn.logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}") + logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}") await send_mcp_message(conn, payload) @@ -307,8 +307,6 @@ async def call_mcp_tool( # 如果解析失败,尝试合并多个JSON对象 try: # 使用正则表达式匹配所有JSON对象 - import re - json_objects = re.findall(r"\{[^{}]*\}", args) if len(json_objects) > 1: # 合并所有JSON对象 @@ -327,7 +325,7 @@ async def call_mcp_tool( else: raise ValueError(f"参数JSON解析失败: {args}") except Exception as e: - conn.logger.bind(tag=TAG).error( + logger.bind(tag=TAG).error( f"参数JSON解析失败: {str(e)}, 原始参数: {args}" ) raise ValueError(f"参数JSON解析失败: {str(e)}") @@ -353,15 +351,13 @@ async def call_mcp_tool( "params": {"name": actual_name, "arguments": arguments}, } - conn.logger.bind(tag=TAG).info( - f"发送客户端mcp工具调用请求: {actual_name},参数: {args}" - ) + logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {actual_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( + logger.bind(tag=TAG).info( f"客户端mcp工具调用 {actual_name} 成功,原始结果: {raw_result}" ) diff --git a/main/xiaozhi-server/core/providers/tools/mcp_endpoint/__init__.py b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/__init__.py new file mode 100644 index 00000000..de2d9b1c --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/__init__.py @@ -0,0 +1,21 @@ +"""MCP接入点工具模块""" + +from .mcp_endpoint_executor import MCPEndpointExecutor +from .mcp_endpoint_client import MCPEndpointClient +from .mcp_endpoint_handler import ( + connect_mcp_endpoint, + send_mcp_endpoint_initialize, + send_mcp_endpoint_notification, + send_mcp_endpoint_tools_list, + call_mcp_endpoint_tool, +) + +__all__ = [ + "MCPEndpointExecutor", + "MCPEndpointClient", + "connect_mcp_endpoint", + "send_mcp_endpoint_initialize", + "send_mcp_endpoint_notification", + "send_mcp_endpoint_tools_list", + "call_mcp_endpoint_tool", +] diff --git a/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_client.py b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_client.py new file mode 100644 index 00000000..da8241d5 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_client.py @@ -0,0 +1,111 @@ +"""MCP接入点客户端定义""" + +import asyncio +from concurrent.futures import Future +from core.utils.util import sanitize_tool_name +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + + +class MCPEndpointClient: + """MCP接入点客户端,用于管理MCP接入点状态和工具""" + + def __init__(self): + self.tools = {} # sanitized_name -> tool_data + self.name_mapping = {} + 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 + self.websocket = None # WebSocket连接 + + 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_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: + sanitized_name = sanitize_tool_name(tool_data["name"]) + self.tools[sanitized_name] = tool_data + self.name_mapping[sanitized_name] = tool_data["name"] + 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) + + def set_websocket(self, websocket): + """设置WebSocket连接""" + self.websocket = websocket + + async def send_message(self, message: str): + """发送消息到MCP接入点""" + if self.websocket: + await self.websocket.send(message) + else: + raise RuntimeError("WebSocket连接未建立") + + async def close(self): + """关闭WebSocket连接""" + if self.websocket: + await self.websocket.close() + self.websocket = None diff --git a/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_executor.py b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_executor.py new file mode 100644 index 00000000..f2c9bac1 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_executor.py @@ -0,0 +1,97 @@ +"""MCP接入点工具执行器""" + +from typing import Dict, Any +from ..base import ToolType, ToolDefinition, ToolExecutor +from plugins_func.register import Action, ActionResponse +from .mcp_endpoint_handler import call_mcp_endpoint_tool + + +class MCPEndpointExecutor(ToolExecutor): + """MCP接入点工具执行器""" + + def __init__(self, conn): + self.conn = conn + + async def execute( + self, conn, tool_name: str, arguments: Dict[str, Any] + ) -> ActionResponse: + """执行MCP接入点工具""" + if not hasattr(conn, "mcp_endpoint_client") or not conn.mcp_endpoint_client: + return ActionResponse( + action=Action.ERROR, + response="MCP接入点客户端未初始化", + ) + + if not await conn.mcp_endpoint_client.is_ready(): + return ActionResponse( + action=Action.ERROR, + response="MCP接入点客户端未准备就绪", + ) + + try: + # 转换参数为JSON字符串 + import json + + args_str = json.dumps(arguments) if arguments else "{}" + + # 调用MCP接入点工具 + result = await call_mcp_endpoint_tool( + conn.mcp_endpoint_client, tool_name, args_str + ) + + resultJson = None + if isinstance(result, str): + try: + resultJson = json.loads(result) + except Exception as e: + pass + + # 视觉大模型不经过二次LLM处理 + if ( + resultJson is not None + and isinstance(resultJson, dict) + and "action" in resultJson + ): + return ActionResponse( + action=Action[resultJson["action"]], + response=resultJson.get("response", ""), + ) + + return ActionResponse(action=Action.REQLLM, result=str(result)) + + except ValueError as e: + return ActionResponse(action=Action.NOTFOUND, response=str(e)) + except Exception as e: + return ActionResponse(action=Action.ERROR, response=str(e)) + + def get_tools(self) -> Dict[str, ToolDefinition]: + """获取所有MCP接入点工具""" + if ( + not hasattr(self.conn, "mcp_endpoint_client") + or not self.conn.mcp_endpoint_client + ): + return {} + + tools = {} + mcp_tools = self.conn.mcp_endpoint_client.get_available_tools() + + for tool in mcp_tools: + func_def = tool.get("function", {}) + tool_name = func_def.get("name", "") + + if tool_name: + tools[tool_name] = ToolDefinition( + name=tool_name, description=tool, tool_type=ToolType.MCP_ENDPOINT + ) + + return tools + + def has_tool(self, tool_name: str) -> bool: + """检查是否有指定的MCP接入点工具""" + if ( + not hasattr(self.conn, "mcp_endpoint_client") + or not self.conn.mcp_endpoint_client + ): + return False + + return self.conn.mcp_endpoint_client.has_tool(tool_name) diff --git a/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_handler.py b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_handler.py new file mode 100644 index 00000000..bf3253f0 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/mcp_endpoint/mcp_endpoint_handler.py @@ -0,0 +1,365 @@ +"""MCP接入点处理器""" + +import json +import asyncio +import re +import websockets +from concurrent.futures import Future +from core.utils.util import sanitize_tool_name +from config.logger import setup_logging +from .mcp_endpoint_client import MCPEndpointClient + +TAG = __name__ +logger = setup_logging() + + +async def connect_mcp_endpoint(mcp_endpoint_url: str) -> MCPEndpointClient: + """连接到MCP接入点""" + if not mcp_endpoint_url: + logger.bind(tag=TAG).warning("MCP接入点URL为空,跳过连接") + return None + + try: + logger.bind(tag=TAG).info(f"正在连接到MCP接入点: {mcp_endpoint_url}") + websocket = await websockets.connect(mcp_endpoint_url) + + mcp_client = MCPEndpointClient() + mcp_client.set_websocket(websocket) + + # 启动消息监听器 + asyncio.create_task(_message_listener(mcp_client)) + + # 发送初始化消息 + await send_mcp_endpoint_initialize(mcp_client) + + # 发送初始化完成通知 + await send_mcp_endpoint_notification(mcp_client, "notifications/initialized") + + # 获取工具列表 + await send_mcp_endpoint_tools_list(mcp_client) + + logger.bind(tag=TAG).info("MCP接入点连接成功") + return mcp_client + + except Exception as e: + logger.bind(tag=TAG).error(f"连接MCP接入点失败: {e}") + return None + + +async def _message_listener(mcp_client: MCPEndpointClient): + """监听MCP接入点消息""" + try: + async for message in mcp_client.websocket: + await handle_mcp_endpoint_message(mcp_client, message) + except websockets.exceptions.ConnectionClosed: + logger.bind(tag=TAG).info("MCP接入点连接已关闭") + except Exception as e: + logger.bind(tag=TAG).error(f"MCP接入点消息监听器错误: {e}") + finally: + await mcp_client.set_ready(False) + + +async def handle_mcp_endpoint_message(mcp_client: MCPEndpointClient, message: str): + """处理MCP接入点消息""" + try: + payload = json.loads(message) + logger.bind(tag=TAG).debug(f"收到MCP接入点消息: {payload}") + + if not isinstance(payload, dict): + logger.bind(tag=TAG).error("MCP接入点消息格式错误") + return + + # Handle result + if "result" in payload: + result = payload["result"] + # 安全地获取消息ID,如果为None则使用0 + msg_id_raw = payload.get("id") + msg_id = int(msg_id_raw) if msg_id_raw is not None else 0 + + # Check for tool call response first + if msg_id in mcp_client.call_results: + 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 + 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") + logger.bind(tag=TAG).info( + f"MCP接入点服务器信息: name={name}, version={version}" + ) + return + + elif msg_id == 2: # mcpToolsListID + logger.bind(tag=TAG).debug("收到MCP接入点工具列表响应") + if isinstance(result, dict) and "tools" in result: + tools_data = result["tools"] + if not isinstance(tools_data, list): + logger.bind(tag=TAG).error("工具列表格式错误") + return + + logger.bind(tag=TAG).info( + f"MCP接入点支持的工具数量: {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) + logger.bind(tag=TAG).debug(f"MCP接入点工具 #{i+1}: {name}") + + # 替换所有工具描述中的工具名称 + for tool_data in mcp_client.tools.values(): + if "description" in tool_data: + description = tool_data["description"] + # 遍历所有工具名称进行替换 + for ( + sanitized_name, + original_name, + ) in mcp_client.name_mapping.items(): + description = description.replace( + original_name, sanitized_name + ) + tool_data["description"] = description + + next_cursor = result.get("nextCursor", "") + if next_cursor: + logger.bind(tag=TAG).info( + f"有更多工具,nextCursor: {next_cursor}" + ) + await send_mcp_endpoint_tools_list_continue( + mcp_client, next_cursor + ) + else: + await mcp_client.set_ready(True) + logger.bind(tag=TAG).info( + "所有MCP接入点工具已获取,客户端准备就绪" + ) + return + + # Handle method calls (requests from the endpoint) + elif "method" in payload: + method = payload["method"] + logger.bind(tag=TAG).info(f"收到MCP接入点请求: {method}") + + elif "error" in payload: + error_data = payload["error"] + error_msg = error_data.get("message", "未知错误") + logger.bind(tag=TAG).error(f"收到MCP接入点错误响应: {error_msg}") + + # 安全地获取消息ID,如果为None则使用0 + msg_id_raw = payload.get("id") + msg_id = int(msg_id_raw) if msg_id_raw is not None else 0 + + if msg_id in mcp_client.call_results: + await mcp_client.reject_call_result( + msg_id, Exception(f"MCP接入点错误: {error_msg}") + ) + + except json.JSONDecodeError as e: + logger.bind(tag=TAG).error(f"MCP接入点消息JSON解析失败: {e}") + except Exception as e: + logger.bind(tag=TAG).error(f"处理MCP接入点消息时出错: {e}") + import traceback + + logger.bind(tag=TAG).error(f"错误详情: {traceback.format_exc()}") + + +async def send_mcp_endpoint_initialize(mcp_client: MCPEndpointClient): + """发送MCP接入点初始化消息""" + payload = { + "jsonrpc": "2.0", + "id": 1, # mcpInitializeID + "method": "initialize", + "params": { + "protocolVersion": "2024-11-05", + "capabilities": { + "roots": {"listChanged": True}, + "sampling": {}, + }, + "clientInfo": { + "name": "XiaozhiMCPEndpointClient", + "version": "1.0.0", + }, + }, + } + message = json.dumps(payload) + logger.bind(tag=TAG).info("发送MCP接入点初始化消息") + await mcp_client.send_message(message) + + +async def send_mcp_endpoint_notification(mcp_client: MCPEndpointClient, method: str): + """发送MCP接入点通知消息""" + payload = { + "jsonrpc": "2.0", + "method": method, + "params": {}, + } + message = json.dumps(payload) + logger.bind(tag=TAG).debug(f"发送MCP接入点通知: {method}") + await mcp_client.send_message(message) + + +async def send_mcp_endpoint_tools_list(mcp_client: MCPEndpointClient): + """发送MCP接入点工具列表请求""" + payload = { + "jsonrpc": "2.0", + "id": 2, # mcpToolsListID + "method": "tools/list", + } + message = json.dumps(payload) + logger.bind(tag=TAG).debug("发送MCP接入点工具列表请求") + await mcp_client.send_message(message) + + +async def send_mcp_endpoint_tools_list_continue( + mcp_client: MCPEndpointClient, cursor: str +): + """发送带有cursor的MCP接入点工具列表请求""" + payload = { + "jsonrpc": "2.0", + "id": 2, # mcpToolsListID (same ID for continuation) + "method": "tools/list", + "params": {"cursor": cursor}, + } + message = json.dumps(payload) + logger.bind(tag=TAG).info(f"发送带cursor的MCP接入点工具列表请求: {cursor}") + await mcp_client.send_message(message) + + +async def call_mcp_endpoint_tool( + mcp_client: MCPEndpointClient, tool_name: str, args: str = "{}", timeout: int = 30 +): + """ + 调用指定的MCP接入点工具,并等待响应 + """ + 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对象 + 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: + 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 + + actual_name = mcp_client.name_mapping.get(tool_name, tool_name) + payload = { + "jsonrpc": "2.0", + "id": tool_call_id, + "method": "tools/call", + "params": {"name": actual_name, "arguments": arguments}, + } + + message = json.dumps(payload) + logger.bind(tag=TAG).info(f"发送MCP接入点工具调用请求: {actual_name},参数: {args}") + await mcp_client.send_message(message) + + try: + # Wait for response or timeout + raw_result = await asyncio.wait_for(result_future, timeout=timeout) + logger.bind(tag=TAG).info( + f"MCP接入点工具调用 {actual_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/providers/tools/server_mcp/__init__.py b/main/xiaozhi-server/core/providers/tools/server_mcp/__init__.py new file mode 100644 index 00000000..aecd80b8 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/__init__.py @@ -0,0 +1,7 @@ +"""服务端MCP工具模块""" + +from .mcp_manager import ServerMCPManager +from .mcp_executor import ServerMCPExecutor +from .mcp_client import ServerMCPClient + +__all__ = ["ServerMCPManager", "ServerMCPExecutor", "ServerMCPClient"] diff --git a/main/xiaozhi-server/core/mcp/MCPClient.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_client.py similarity index 74% rename from main/xiaozhi-server/core/mcp/MCPClient.py rename to main/xiaozhi-server/core/providers/tools/server_mcp/mcp_client.py index 5b7bab55..8b60ac1b 100644 --- a/main/xiaozhi-server/core/mcp/MCPClient.py +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_client.py @@ -1,7 +1,12 @@ +"""服务端MCP客户端""" + from __future__ import annotations from datetime import timedelta -import asyncio, os, shutil, concurrent.futures +import asyncio +import os +import shutil +import concurrent.futures from contextlib import AsyncExitStack from typing import Optional, List, Dict, Any @@ -14,8 +19,15 @@ from core.utils.util import sanitize_tool_name TAG = __name__ -class MCPClient: +class ServerMCPClient: + """服务端MCP客户端,用于连接和管理MCP服务""" + def __init__(self, config: Dict[str, Any]): + """初始化服务端MCP客户端 + + Args: + config: MCP服务配置字典 + """ self.logger = setup_logging() self.config = config @@ -24,21 +36,26 @@ class MCPClient: self._shutdown_evt = asyncio.Event() self.session: Optional[ClientSession] = None - self.tools: List = [] # original tool objects + 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="MCPClientWorker") + + self._worker_task = asyncio.create_task( + self._worker(), name="ServerMCPClientWorker" + ) await self._ready_evt.wait() self.logger.bind(tag=TAG).info( - f"Connected, tools = {[name for name in self.name_mapping.values()]}" + f"服务端MCP客户端已连接,可用工具: {[name for name in self.name_mapping.values()]}" ) async def cleanup(self): + """清理MCP客户端资源""" if not self._worker_task: return @@ -46,14 +63,27 @@ class MCPClient: try: await asyncio.wait_for(self._worker_task, timeout=20) except (asyncio.TimeoutError, Exception) as e: - self.logger.bind(tag=TAG).error(f"worker shutdown err: {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): + def get_available_tools(self) -> List[Dict[str, Any]]: + """获取所有可用工具的定义 + + Returns: + List[Dict[str, Any]]: 工具定义列表 + """ return [ { "type": "function", @@ -66,9 +96,21 @@ class MCPClient: for name, tool in self.tools_dict.items() ] - async def call_tool(self, name: str, args: dict): + async def call_tool(self, name: str, args: dict) -> Any: + """调用指定工具 + + Args: + name: 工具名称 + args: 工具参数 + + Returns: + Any: 工具执行结果 + + Raises: + RuntimeError: 客户端未初始化时抛出 + """ if not self.session: - raise RuntimeError("MCPClient not initialized") + raise RuntimeError("服务端MCP客户端未初始化") real_name = self.name_mapping.get(name, name) loop = self._worker_task.get_loop() @@ -89,19 +131,20 @@ class MCPClient: # 检查工作任务是否存在 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 @@ -121,6 +164,7 @@ class MCPClient: stdio_client(params) ) read_stream, write_stream = stdio_r, stdio_w + # 建立SSEClient elif "url" in self.config: if "API_ACCESS_TOKEN" in self.config: @@ -135,7 +179,7 @@ class MCPClient: read_stream, write_stream = sse_r, sse_w else: - raise ValueError("MCPClient config must include 'command' or 'url'") + raise ValueError("MCP客户端配置必须包含'command'或'url'") self.session = await stack.enter_async_context( ClientSession( @@ -159,6 +203,6 @@ class MCPClient: await self._shutdown_evt.wait() except Exception as e: - self.logger.bind(tag=TAG).error(f"worker error: {e}") + self.logger.bind(tag=TAG).error(f"服务端MCP客户端工作协程错误: {e}") self._ready_evt.set() raise diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py new file mode 100644 index 00000000..9ae15d4b --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py @@ -0,0 +1,89 @@ +"""服务端MCP工具执行器""" + +from typing import Dict, Any, Optional +from ..base import ToolType, ToolDefinition, ToolExecutor +from plugins_func.register import Action, ActionResponse +from .mcp_manager import ServerMCPManager + + +class ServerMCPExecutor(ToolExecutor): + """服务端MCP工具执行器""" + + def __init__(self, conn): + self.conn = conn + self.mcp_manager: Optional[ServerMCPManager] = None + self._initialized = False + + async def initialize(self): + """初始化MCP管理器""" + if not self._initialized: + self.mcp_manager = ServerMCPManager(self.conn) + await self.mcp_manager.initialize_servers() + self._initialized = True + + async def execute( + self, conn, tool_name: str, arguments: Dict[str, Any] + ) -> ActionResponse: + """执行服务端MCP工具""" + if not self._initialized or not self.mcp_manager: + return ActionResponse( + action=Action.ERROR, + response="MCP管理器未初始化", + ) + + try: + # 移除mcp_前缀(如果有) + actual_tool_name = tool_name + if tool_name.startswith("mcp_"): + actual_tool_name = tool_name[4:] + + result = await self.mcp_manager.execute_tool(actual_tool_name, arguments) + + return ActionResponse(action=Action.REQLLM, result=str(result)) + + except ValueError as e: + return ActionResponse( + action=Action.NOTFOUND, + response=str(e), + ) + except Exception as e: + return ActionResponse( + action=Action.ERROR, + response=str(e), + ) + + def get_tools(self) -> Dict[str, ToolDefinition]: + """获取所有服务端MCP工具""" + if not self._initialized or not self.mcp_manager: + return {} + + tools = {} + mcp_tools = self.mcp_manager.get_all_tools() + + for tool in mcp_tools: + func_def = tool.get("function", {}) + tool_name = func_def.get("name", "") + if tool_name == "": + continue + tools[tool_name] = ToolDefinition( + name=tool_name, description=tool, tool_type=ToolType.SERVER_MCP + ) + + return tools + + def has_tool(self, tool_name: str) -> bool: + """检查是否有指定的服务端MCP工具""" + if not self._initialized or not self.mcp_manager: + return False + + # 移除mcp_前缀(如果有) + actual_tool_name = tool_name + if tool_name.startswith("mcp_"): + actual_tool_name = tool_name[4:] + + return self.mcp_manager.is_mcp_tool(actual_tool_name) + + async def cleanup(self): + """清理MCP连接""" + if self.mcp_manager: + await self.mcp_manager.cleanup_all() diff --git a/main/xiaozhi-server/core/mcp/manager.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py similarity index 53% rename from main/xiaozhi-server/core/mcp/manager.py rename to main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py index 1c677e0e..6589c302 100644 --- a/main/xiaozhi-server/core/mcp/manager.py +++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py @@ -1,37 +1,34 @@ -"""MCP服务管理器""" +"""服务端MCP管理器""" import asyncio -import os, json +import os +import json from typing import Dict, Any, List -from .MCPClient import MCPClient -from plugins_func.register import register_function, ToolType from config.config_loader import get_project_dir +from config.logger import setup_logging +from .mcp_client import ServerMCPClient TAG = __name__ +logger = setup_logging() -class MCPManager: - """管理多个MCP服务的集中管理器""" +class ServerMCPManager: + """管理多个服务端MCP服务的集中管理器""" def __init__(self, conn) -> None: - """ - 初始化MCP管理器 - """ + """初始化MCP管理器""" self.conn = conn self.config_path = get_project_dir() + "data/.mcp_server_settings.json" - if os.path.exists(self.config_path) == False: + if not os.path.exists(self.config_path): self.config_path = "" - self.conn.logger.bind(tag=TAG).warning( + logger.bind(tag=TAG).warning( f"请检查mcp服务配置文件:data/.mcp_server_settings.json" ) - self.client: Dict[str, MCPClient] = {} + self.clients: Dict[str, ServerMCPClient] = {} self.tools = [] def load_config(self) -> Dict[str, Any]: - """加载MCP服务配置 - Returns: - Dict[str, Any]: 服务配置字典 - """ + """加载MCP服务配置""" if len(self.config_path) == 0: return {} @@ -40,7 +37,7 @@ class MCPManager: config = json.load(f) return config.get("mcpServers", {}) except Exception as e: - self.conn.logger.bind(tag=TAG).error( + logger.bind(tag=TAG).error( f"Error loading MCP config from {self.config_path}: {e}" ) return {} @@ -50,84 +47,58 @@ class MCPManager: config = self.load_config() for name, srv_config in config.items(): if not srv_config.get("command") and not srv_config.get("url"): - self.conn.logger.bind(tag=TAG).warning( + logger.bind(tag=TAG).warning( f"Skipping server {name}: neither command nor url specified" ) continue try: - client = MCPClient(srv_config) + # 初始化服务端MCP客户端 + logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}") + client = ServerMCPClient(srv_config) await client.initialize() - self.client[name] = client - self.conn.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}") + self.clients[name] = client client_tools = client.get_available_tools() self.tools.extend(client_tools) - for tool in client_tools: - func_name = "mcp_" + tool["function"]["name"] - register_function(func_name, tool, ToolType.MCP_CLIENT)( - self.execute_tool - ) - self.conn.func_handler.function_registry.register_function( - func_name - ) except Exception as e: - self.conn.logger.bind(tag=TAG).error( + logger.bind(tag=TAG).error( f"Failed to initialize MCP server {name}: {e}" ) - self.conn.func_handler.upload_functions_desc() def get_all_tools(self) -> List[Dict[str, Any]]: - """获取所有服务的工具function定义 - Returns: - List[Dict[str, Any]]: 所有工具的function定义列表 - """ + """获取所有服务的工具function定义""" return self.tools def is_mcp_tool(self, tool_name: str) -> bool: - """检查是否是MCP工具 - Args: - tool_name: 工具名称 - Returns: - bool: 是否是MCP工具 - """ + """检查是否是MCP工具""" for tool in self.tools: if ( - tool.get("function") != None + tool.get("function") is not None and tool["function"].get("name") == tool_name ): return True return False async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Any: - """执行工具调用,失败时会尝试重新连接 - Args: - tool_name: 工具名称 - arguments: 工具参数 - Returns: - Any: 工具执行结果 - Raises: - ValueError: 工具未找到时抛出 - """ - self.conn.logger.bind(tag=TAG).info( - f"Executing tool {tool_name} with arguments: {arguments}" - ) - + """执行工具调用,失败时会尝试重新连接""" + logger.bind(tag=TAG).info(f"执行服务端MCP工具 {tool_name},参数: {arguments}") + max_retries = 3 # 最大重试次数 retry_interval = 2 # 重试间隔(秒) - + # 找到对应的客户端 client_name = None target_client = None - for name, client in self.client.items(): + 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 {tool_name} not found in any MCP server") - + raise ValueError(f"工具 {tool_name} 在任意MCP服务中未找到") + # 带重试机制的工具调用 for attempt in range(max_retries): try: @@ -136,49 +107,48 @@ class MCPManager: # 最后一次尝试失败时直接抛出异常 if attempt == max_retries - 1: raise - - self.conn.logger.bind(tag=TAG).warning( + + logger.bind(tag=TAG).warning( f"执行工具 {tool_name} 失败 (尝试 {attempt+1}/{max_retries}): {e}" ) - + # 尝试重新连接 - self.conn.logger.bind(tag=TAG).info( + logger.bind(tag=TAG).info( f"重试前尝试重新连接 MCP 客户端 {client_name}" ) try: # 关闭旧的连接 await target_client.cleanup() - + # 重新初始化客户端 config = self.load_config() if client_name in config: - client = MCPClient(config[client_name]) + client = ServerMCPClient(config[client_name]) await client.initialize() - self.client[client_name] = client - target_client = client - self.conn.logger.bind(tag=TAG).info( + self.clients[client_name] = client + target_client = client + logger.bind(tag=TAG).info( f"成功重新连接 MCP 客户端: {client_name}" ) else: - self.conn.logger.bind(tag=TAG).error( + logger.bind(tag=TAG).error( f"Cannot reconnect MCP client {client_name}: config not found" ) except Exception as reconnect_error: - self.conn.logger.bind(tag=TAG).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: - """依次关闭所有 MCPClient,不让异常阻断整体流程。""" - for name, client in list(self.client.items()): + """关闭所有 MCP客户端""" + for name, client in list(self.clients.items()): try: - await asyncio.wait_for(client.cleanup(), timeout=20) - self.conn.logger.bind(tag=TAG).info(f"MCP client closed: {name}") + if hasattr(client, "cleanup"): + await asyncio.wait_for(client.cleanup(), timeout=20) + logger.bind(tag=TAG).info(f"服务端MCP客户端已关闭: {name}") except (asyncio.TimeoutError, Exception) as e: - self.conn.logger.bind(tag=TAG).error( - f"Error closing MCP client {name}: {e}" - ) - self.client.clear() + logger.bind(tag=TAG).error(f"关闭服务端MCP客户端 {name} 时出错: {e}") + self.clients.clear() diff --git a/main/xiaozhi-server/core/providers/tools/server_plugins/__init__.py b/main/xiaozhi-server/core/providers/tools/server_plugins/__init__.py new file mode 100644 index 00000000..232b50d2 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/server_plugins/__init__.py @@ -0,0 +1,5 @@ +"""服务端插件工具模块""" + +from .plugin_executor import ServerPluginExecutor + +__all__ = ["ServerPluginExecutor"] diff --git a/main/xiaozhi-server/core/providers/tools/server_plugins/plugin_executor.py b/main/xiaozhi-server/core/providers/tools/server_plugins/plugin_executor.py new file mode 100644 index 00000000..e9194e5a --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/server_plugins/plugin_executor.py @@ -0,0 +1,77 @@ +"""服务端插件工具执行器""" + +from typing import Dict, Any +from ..base import ToolType, ToolDefinition, ToolExecutor +from plugins_func.register import all_function_registry, Action, ActionResponse + + +class ServerPluginExecutor(ToolExecutor): + """服务端插件工具执行器""" + + def __init__(self, conn): + self.conn = conn + self.config = conn.config + + async def execute( + self, conn, tool_name: str, arguments: Dict[str, Any] + ) -> ActionResponse: + """执行服务端插件工具""" + func_item = all_function_registry.get(tool_name) + if not func_item: + return ActionResponse( + action=Action.NOTFOUND, response=f"插件函数 {tool_name} 不存在" + ) + + try: + # 根据工具类型决定如何调用 + if hasattr(func_item, "type"): + func_type = func_item.type + if func_type.code in [4, 5]: # SYSTEM_CTL, IOT_CTL (需要conn参数) + result = func_item.func(conn, **arguments) + elif func_type.code == 2: # WAIT + result = func_item.func(**arguments) + elif func_type.code == 3: # CHANGE_SYS_PROMPT + result = func_item.func(conn, **arguments) + else: + result = func_item.func(**arguments) + else: + # 默认不传conn参数 + result = func_item.func(**arguments) + + return result + + except Exception as e: + return ActionResponse( + action=Action.ERROR, + response=str(e), + ) + + def get_tools(self) -> Dict[str, ToolDefinition]: + """获取所有注册的服务端插件工具""" + tools = {} + + # 获取必要的函数 + necessary_functions = ["handle_exit_intent", "get_time", "get_lunar"] + + # 获取配置中的函数 + config_functions = self.config["Intent"][ + self.config["selected_module"]["Intent"] + ].get("functions", []) + + # 合并所有需要的函数 + all_required_functions = list(set(necessary_functions + config_functions)) + + for func_name in all_required_functions: + func_item = all_function_registry.get(func_name) + if func_item: + tools[func_name] = ToolDefinition( + name=func_name, + description=func_item.description, + tool_type=ToolType.SERVER_PLUGIN, + ) + + return tools + + def has_tool(self, tool_name: str) -> bool: + """检查是否有指定的服务端插件工具""" + return tool_name in all_function_registry diff --git a/main/xiaozhi-server/core/providers/tools/unified_tool_handler.py b/main/xiaozhi-server/core/providers/tools/unified_tool_handler.py new file mode 100644 index 00000000..e65aa791 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/unified_tool_handler.py @@ -0,0 +1,226 @@ +"""统一工具处理器""" + +import json +from typing import Dict, List, Any, Optional +from config.logger import setup_logging +from plugins_func.loadplugins import auto_import_modules + +from .base import ToolType +from plugins_func.register import Action, ActionResponse +from .unified_tool_manager import ToolManager +from .server_plugins import ServerPluginExecutor +from .server_mcp import ServerMCPExecutor +from .device_iot import DeviceIoTExecutor +from .device_mcp import DeviceMCPExecutor +from .mcp_endpoint import MCPEndpointExecutor + + +class UnifiedToolHandler: + """统一工具处理器""" + + def __init__(self, conn): + self.conn = conn + self.config = conn.config + self.logger = setup_logging() + + # 创建工具管理器 + self.tool_manager = ToolManager(conn) + + # 创建各类执行器 + self.server_plugin_executor = ServerPluginExecutor(conn) + self.server_mcp_executor = ServerMCPExecutor(conn) + self.device_iot_executor = DeviceIoTExecutor(conn) + self.device_mcp_executor = DeviceMCPExecutor(conn) + self.mcp_endpoint_executor = MCPEndpointExecutor(conn) + + # 注册执行器 + self.tool_manager.register_executor( + ToolType.SERVER_PLUGIN, self.server_plugin_executor + ) + self.tool_manager.register_executor( + ToolType.SERVER_MCP, self.server_mcp_executor + ) + self.tool_manager.register_executor( + ToolType.DEVICE_IOT, self.device_iot_executor + ) + self.tool_manager.register_executor( + ToolType.DEVICE_MCP, self.device_mcp_executor + ) + self.tool_manager.register_executor( + ToolType.MCP_ENDPOINT, self.mcp_endpoint_executor + ) + + # 初始化标志 + self.finish_init = False + + async def _initialize(self): + """异步初始化""" + try: + # 自动导入插件模块 + auto_import_modules("plugins_func.functions") + + # 初始化服务端MCP + await self.server_mcp_executor.initialize() + + # 初始化MCP接入点 + await self._initialize_mcp_endpoint() + + # 初始化Home Assistant(如果需要) + self._initialize_home_assistant() + + self.finish_init = True + self.logger.info("统一工具处理器初始化完成") + + except Exception as e: + self.logger.error(f"统一工具处理器初始化失败: {e}") + + async def _initialize_mcp_endpoint(self): + """初始化MCP接入点""" + try: + from .mcp_endpoint import connect_mcp_endpoint + + # 从配置中获取MCP接入点URL + mcp_endpoint_url = self.config.get("mcp_endpoint", "") + + if mcp_endpoint_url and "你的" not in mcp_endpoint_url: + self.logger.info(f"正在初始化MCP接入点: {mcp_endpoint_url}") + mcp_endpoint_client = await connect_mcp_endpoint(mcp_endpoint_url) + + if mcp_endpoint_client: + # 将MCP接入点客户端保存到连接对象中 + self.conn.mcp_endpoint_client = mcp_endpoint_client + self.logger.info("MCP接入点初始化成功") + else: + self.logger.warning("MCP接入点初始化失败") + + except Exception as e: + self.logger.error(f"初始化MCP接入点失败: {e}") + + def _initialize_home_assistant(self): + """初始化Home Assistant提示词""" + try: + from plugins_func.functions.hass_init import append_devices_to_prompt + + append_devices_to_prompt(self.conn) + except ImportError: + pass # 忽略导入错误 + except Exception as e: + self.logger.error(f"初始化Home Assistant失败: {e}") + + def get_functions(self) -> List[Dict[str, Any]]: + """获取所有工具的函数描述""" + return self.tool_manager.get_function_descriptions() + + def current_support_functions(self) -> List[str]: + """获取当前支持的函数名称列表""" + func_names = self.tool_manager.get_supported_tool_names() + self.logger.info(f"当前支持的函数列表: {func_names}") + return func_names + + def upload_functions_desc(self): + """刷新函数描述列表""" + self.tool_manager.refresh_tools() + self.logger.info("函数描述列表已刷新") + + def has_tool(self, tool_name: str) -> bool: + """检查是否有指定工具""" + return self.tool_manager.has_tool(tool_name) + + async def handle_llm_function_call( + self, conn, function_call_data: Dict[str, Any] + ) -> Optional[ActionResponse]: + """处理LLM函数调用""" + try: + # 处理多函数调用 + if "function_calls" in function_call_data: + responses = [] + for call in function_call_data["function_calls"]: + result = await self.tool_manager.execute_tool( + call["name"], call.get("arguments", {}) + ) + responses.append(result) + return self._combine_responses(responses) + + # 处理单函数调用 + function_name = function_call_data["name"] + arguments = function_call_data.get("arguments", {}) + + # 如果arguments是字符串,尝试解析为JSON + if isinstance(arguments, str): + try: + arguments = json.loads(arguments) if arguments else {} + except json.JSONDecodeError: + self.logger.error(f"无法解析函数参数: {arguments}") + return ActionResponse( + action=Action.ERROR, + response="无法解析函数参数", + ) + + self.logger.debug(f"调用函数: {function_name}, 参数: {arguments}") + + # 执行工具调用 + result = await self.tool_manager.execute_tool(function_name, arguments) + return result + + except Exception as e: + self.logger.error(f"处理function call错误: {e}") + return ActionResponse(action=Action.ERROR, response=str(e)) + + def _combine_responses(self, responses: List[ActionResponse]) -> ActionResponse: + """合并多个函数调用的响应""" + if not responses: + return ActionResponse(action=Action.NONE, response="无响应") + + # 如果有任何错误,返回第一个错误 + for response in responses: + if response.action == Action.ERROR: + return response + + # 合并所有成功的响应 + contents = [] + responses_text = [] + + for response in responses: + if response.content: + contents.append(response.content) + if response.response: + responses_text.append(response.response) + + # 确定最终的动作类型 + final_action = Action.RESPONSE + for response in responses: + if response.action == Action.REQLLM: + final_action = Action.REQLLM + break + + return ActionResponse( + action=final_action, + result="; ".join(contents) if contents else None, + response="; ".join(responses_text) if responses_text else None, + ) + + async def register_iot_tools(self, descriptors: List[Dict[str, Any]]): + """注册IoT设备工具""" + self.device_iot_executor.register_iot_tools(descriptors) + self.tool_manager.refresh_tools() + self.logger.info(f"注册了{len(descriptors)}个IoT设备的工具") + + def get_tool_statistics(self) -> Dict[str, int]: + """获取工具统计信息""" + return self.tool_manager.get_tool_statistics() + + async def cleanup(self): + """清理资源""" + try: + await self.server_mcp_executor.cleanup() + + # 清理MCP接入点连接 + if ( + hasattr(self.conn, "mcp_endpoint_client") + and self.conn.mcp_endpoint_client + ): + await self.conn.mcp_endpoint_client.close() + + self.logger.info("工具处理器清理完成") + except Exception as e: + self.logger.error(f"工具处理器清理失败: {e}") diff --git a/main/xiaozhi-server/core/providers/tools/unified_tool_manager.py b/main/xiaozhi-server/core/providers/tools/unified_tool_manager.py new file mode 100644 index 00000000..e2a91869 --- /dev/null +++ b/main/xiaozhi-server/core/providers/tools/unified_tool_manager.py @@ -0,0 +1,124 @@ +"""统一工具管理器""" + +from typing import Dict, List, Optional, Any +from config.logger import setup_logging +from plugins_func.register import Action, ActionResponse +from .base import ToolType, ToolDefinition, ToolExecutor + + +class ToolManager: + """统一工具管理器,管理所有类型的工具""" + + def __init__(self, conn): + self.conn = conn + self.logger = setup_logging() + self.executors: Dict[ToolType, ToolExecutor] = {} + self._cached_tools: Optional[Dict[str, ToolDefinition]] = None + self._cached_function_descriptions: Optional[List[Dict[str, Any]]] = None + + def register_executor(self, tool_type: ToolType, executor: ToolExecutor): + """注册工具执行器""" + self.executors[tool_type] = executor + self._invalidate_cache() + self.logger.info(f"注册工具执行器: {tool_type.value}") + + def _invalidate_cache(self): + """使缓存失效""" + self._cached_tools = None + self._cached_function_descriptions = None + + def get_all_tools(self) -> Dict[str, ToolDefinition]: + """获取所有工具定义""" + if self._cached_tools is not None: + return self._cached_tools + + all_tools = {} + for tool_type, executor in self.executors.items(): + try: + tools = executor.get_tools() + for name, definition in tools.items(): + if name in all_tools: + self.logger.warning(f"工具名称冲突: {name}") + all_tools[name] = definition + except Exception as e: + self.logger.error(f"获取{tool_type.value}工具时出错: {e}") + + self._cached_tools = all_tools + return all_tools + + def get_function_descriptions(self) -> List[Dict[str, Any]]: + """获取所有工具的函数描述(OpenAI格式)""" + if self._cached_function_descriptions is not None: + return self._cached_function_descriptions + + descriptions = [] + tools = self.get_all_tools() + for tool_definition in tools.values(): + descriptions.append(tool_definition.description) + + self._cached_function_descriptions = descriptions + return descriptions + + def has_tool(self, tool_name: str) -> bool: + """检查是否存在指定工具""" + tools = self.get_all_tools() + return tool_name in tools + + def get_tool_type(self, tool_name: str) -> Optional[ToolType]: + """获取工具类型""" + tools = self.get_all_tools() + tool_def = tools.get(tool_name) + return tool_def.tool_type if tool_def else None + + async def execute_tool( + self, tool_name: str, arguments: Dict[str, Any] + ) -> ActionResponse: + """执行工具调用""" + try: + # 查找工具类型 + tool_type = self.get_tool_type(tool_name) + if not tool_type: + return ActionResponse( + action=Action.NOTFOUND, + response=f"工具 {tool_name} 不存在", + ) + + # 获取对应的执行器 + executor = self.executors.get(tool_type) + if not executor: + return ActionResponse( + action=Action.ERROR, + response=f"工具类型 {tool_type.value} 的执行器未注册", + ) + + # 执行工具 + self.logger.info(f"执行工具: {tool_name},参数: {arguments}") + result = await executor.execute(self.conn, tool_name, arguments) + self.logger.debug(f"工具执行结果: {result}") + return result + + except Exception as e: + self.logger.error(f"执行工具 {tool_name} 时出错: {e}") + return ActionResponse(action=Action.ERROR, response=str(e)) + + def get_supported_tool_names(self) -> List[str]: + """获取所有支持的工具名称""" + tools = self.get_all_tools() + return list(tools.keys()) + + def refresh_tools(self): + """刷新工具缓存""" + self._invalidate_cache() + self.logger.info("工具缓存已刷新") + + def get_tool_statistics(self) -> Dict[str, int]: + """获取工具统计信息""" + stats = {} + for tool_type, executor in self.executors.items(): + try: + tools = executor.get_tools() + stats[tool_type.value] = len(tools) + except Exception as e: + self.logger.error(f"获取{tool_type.value}工具统计时出错: {e}") + stats[tool_type.value] = 0 + return stats diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 1a9d069c..223cb807 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -43,7 +43,6 @@ class TTSProviderBase(ABC): self.tts_text_buff = [] 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: self.tts_text_queue.put( TTSMessageDTO( diff --git a/main/xiaozhi-server/plugins_func/register.py b/main/xiaozhi-server/plugins_func/register.py index 873c61e1..5c2b0781 100644 --- a/main/xiaozhi-server/plugins_func/register.py +++ b/main/xiaozhi-server/plugins_func/register.py @@ -35,7 +35,7 @@ class Action(Enum): class ActionResponse: - def __init__(self, action: Action, result, response): + def __init__(self, action: Action, result=None, response=None): self.action = action # 动作类型 self.result = result # 动作产生的结果 self.response = response # 直接回复的内容