update:统一工具注册及调用

This commit is contained in:
hrz
2025-06-25 18:27:08 +08:00
parent 8c1a8f7d55
commit 2348d9ceb3
27 changed files with 1356 additions and 1055 deletions
+49 -111
View File
@@ -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,7 @@ from core.utils.util import (
filter_sensitive_info,
)
from typing import Dict, Any
from core.mcp.manager import MCPManager
from core.tools.base import ToolAction
from core.utils.modules_initialize import (
initialize_modules,
initialize_tts,
@@ -30,7 +29,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.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 +585,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 +608,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 +621,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 +724,16 @@ 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(
# 使用统一工具处理器处理所有工具调用
tool_result = asyncio.run_coroutine_threadsafe(
self.func_handler.handle_llm_function_call(
self, function_call_data
)
),
self.loop,
).result()
# 转换ToolResult为ActionResponse
result = self._convert_tool_result_to_action_response(tool_result)
self._handle_function_result(result, function_call_data)
# 存储对话内容
@@ -809,47 +756,38 @@ 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}")
def _convert_tool_result_to_action_response(self, tool_result):
"""转换ToolResult为ActionResponse"""
if tool_result.action == ToolAction.ERROR:
return ActionResponse(
action=Action.REQLLM, result="工具调用出错", response=""
action=Action.ERROR,
result=tool_result.error or tool_result.content,
response=tool_result.response or tool_result.content,
)
elif tool_result.action == ToolAction.NOT_FOUND:
return ActionResponse(
action=Action.NOTFOUND,
result=tool_result.content,
response=tool_result.response or tool_result.content,
)
elif tool_result.action == ToolAction.RESPONSE:
return ActionResponse(
action=Action.RESPONSE,
result=tool_result.content,
response=tool_result.response or tool_result.content,
)
elif tool_result.action == ToolAction.REQUEST_LLM:
return ActionResponse(
action=Action.REQLLM,
result=tool_result.content,
response=tool_result.response or "",
)
else:
return ActionResponse(
action=Action.NONE,
result=tool_result.content,
response=tool_result.response or "",
)
return ActionResponse(action=Action.REQLLM, result="工具调用出错", response="")
def _handle_function_result(self, result, function_call_data):
if result.action == Action.RESPONSE: # 直接回复前端
@@ -945,9 +883,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:
@@ -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
@@ -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.tools.device_mcp import (
MCPClient,
send_mcp_initialize_message,
send_mcp_tools_list_request,
@@ -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.tools.device_mcp import call_mcp_tool
from plugins_func.register import Action, ActionResponse
from loguru import logger
@@ -106,36 +106,19 @@ 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:
tool_result = asyncio.run_coroutine_threadsafe(
conn.func_handler.handle_llm_function_call(conn, function_call_data),
conn.loop,
).result()
# 转换ToolResult为ActionResponse
result = conn._convert_tool_result_to_action_response(tool_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:
@@ -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}")
@@ -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.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.tools.device_iot import handleIotDescriptors, handleIotStatus
from core.handle.reportHandle import enqueue_asr_report
import asyncio
-164
View File
@@ -1,164 +0,0 @@
from __future__ import annotations
from datetime import timedelta
import asyncio, os, shutil, concurrent.futures
from contextlib import AsyncExitStack
from typing import Optional, List, Dict, Any
from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client
from mcp.client.sse import sse_client
from config.logger import setup_logging
from core.utils.util import sanitize_tool_name
TAG = __name__
class MCPClient:
def __init__(self, config: Dict[str, Any]):
self.logger = setup_logging()
self.config = config
self._worker_task: Optional[asyncio.Task] = None
self._ready_evt = asyncio.Event()
self._shutdown_evt = asyncio.Event()
self.session: Optional[ClientSession] = None
self.tools: List = [] # original tool objects
self.tools_dict: Dict[str, Any] = {}
self.name_mapping: Dict[str, str] = {}
async def initialize(self):
if self._worker_task:
return
self._worker_task = asyncio.create_task(self._worker(), name="MCPClientWorker")
await self._ready_evt.wait()
self.logger.bind(tag=TAG).info(
f"Connected, tools = {[name for name in self.name_mapping.values()]}"
)
async def cleanup(self):
if not self._worker_task:
return
self._shutdown_evt.set()
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}")
finally:
self._worker_task = None
def has_tool(self, name: str) -> bool:
return name in self.tools_dict
def get_available_tools(self):
return [
{
"type": "function",
"function": {
"name": name,
"description": tool.description,
"parameters": tool.inputSchema,
},
}
for name, tool in self.tools_dict.items()
]
async def call_tool(self, name: str, args: dict):
if not self.session:
raise RuntimeError("MCPClient not initialized")
real_name = self.name_mapping.get(name, name)
loop = self._worker_task.get_loop()
coro = self.session.call_tool(real_name, args)
if loop is asyncio.get_running_loop():
return await coro
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
return await asyncio.wrap_future(fut)
def is_connected(self) -> bool:
"""检查MCP客户端是否连接正常
Returns:
bool: 如果客户端已连接并正常工作,返回True,否则返回False
"""
# 检查工作任务是否存在
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):
async with AsyncExitStack() as stack:
try:
# 建立 StdioClient
if "command" in self.config:
cmd = (
shutil.which("npx")
if self.config["command"] == "npx"
else self.config["command"]
)
env = {**os.environ, **self.config.get("env", {})}
params = StdioServerParameters(
command=cmd,
args=self.config.get("args", []),
env=env,
)
stdio_r, stdio_w = await stack.enter_async_context(
stdio_client(params)
)
read_stream, write_stream = stdio_r, stdio_w
# 建立SSEClient
elif "url" in self.config:
if "API_ACCESS_TOKEN" in self.config:
headers = {
"Authorization": f"Bearer {self.config['API_ACCESS_TOKEN']}"
}
else:
headers = {}
sse_r, sse_w = await stack.enter_async_context(
sse_client(self.config["url"], headers=headers)
)
read_stream, write_stream = sse_r, sse_w
else:
raise ValueError("MCPClient config must include 'command' or 'url'")
self.session = await stack.enter_async_context(
ClientSession(
read_stream=read_stream,
write_stream=write_stream,
read_timeout_seconds=timedelta(seconds=15),
)
)
await self.session.initialize()
# 获取工具
self.tools = (await self.session.list_tools()).tools
for t in self.tools:
sanitized = sanitize_tool_name(t.name)
self.tools_dict[sanitized] = t
self.name_mapping[sanitized] = t.name
self._ready_evt.set()
# 挂起等待关闭
await self._shutdown_evt.wait()
except Exception as e:
self.logger.bind(tag=TAG).error(f"worker error: {e}")
self._ready_evt.set()
raise
-184
View File
@@ -1,184 +0,0 @@
"""MCP服务管理器"""
import asyncio
import os, 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
TAG = __name__
class MCPManager:
"""管理多个MCP服务的集中管理器"""
def __init__(self, conn) -> None:
"""
初始化MCP管理器
"""
self.conn = conn
self.config_path = get_project_dir() + "data/.mcp_server_settings.json"
if os.path.exists(self.config_path) == False:
self.config_path = ""
self.conn.logger.bind(tag=TAG).warning(
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
)
self.client: Dict[str, MCPClient] = {}
self.tools = []
def load_config(self) -> Dict[str, Any]:
"""加载MCP服务配置
Returns:
Dict[str, Any]: 服务配置字典
"""
if len(self.config_path) == 0:
return {}
try:
with open(self.config_path, "r", encoding="utf-8") as f:
config = json.load(f)
return config.get("mcpServers", {})
except Exception as e:
self.conn.logger.bind(tag=TAG).error(
f"Error loading MCP config from {self.config_path}: {e}"
)
return {}
async def initialize_servers(self) -> None:
"""初始化所有MCP服务"""
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(
f"Skipping server {name}: neither command nor url specified"
)
continue
try:
client = MCPClient(srv_config)
await client.initialize()
self.client[name] = client
self.conn.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}")
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(
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定义列表
"""
return self.tools
def is_mcp_tool(self, tool_name: str) -> bool:
"""检查是否是MCP工具
Args:
tool_name: 工具名称
Returns:
bool: 是否是MCP工具
"""
for tool in self.tools:
if (
tool.get("function") != 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}"
)
max_retries = 3 # 最大重试次数
retry_interval = 2 # 重试间隔(秒)
# 找到对应的客户端
client_name = None
target_client = None
for name, client in self.client.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")
# 带重试机制的工具调用
for attempt in range(max_retries):
try:
return await target_client.call_tool(tool_name, arguments)
except Exception as e:
# 最后一次尝试失败时直接抛出异常
if attempt == max_retries - 1:
raise
self.conn.logger.bind(tag=TAG).warning(
f"执行工具 {tool_name} 失败 (尝试 {attempt+1}/{max_retries}): {e}"
)
# 尝试重新连接
self.conn.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])
await client.initialize()
self.client[client_name] = client
target_client = client
self.conn.logger.bind(tag=TAG).info(
f"成功重新连接 MCP 客户端: {client_name}"
)
else:
self.conn.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(
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()):
try:
await asyncio.wait_for(client.cleanup(), timeout=20)
self.conn.logger.bind(tag=TAG).info(f"MCP client closed: {name}")
except (asyncio.TimeoutError, Exception) as e:
self.conn.logger.bind(tag=TAG).error(
f"Error closing MCP client {name}: {e}"
)
self.client.clear()
@@ -0,0 +1 @@
@@ -0,0 +1,6 @@
"""基础工具定义模块"""
from .tool_types import ToolType, ToolAction, ToolResult, ToolDefinition
from .tool_executor import ToolExecutor
__all__ = ["ToolType", "ToolAction", "ToolResult", "ToolDefinition", "ToolExecutor"]
@@ -0,0 +1,26 @@
"""工具执行器基类定义"""
from abc import ABC, abstractmethod
from typing import Dict, Any
from .tool_types import ToolDefinition, ToolResult
class ToolExecutor(ABC):
"""工具执行器抽象基类"""
@abstractmethod
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ToolResult:
"""执行工具调用"""
pass
@abstractmethod
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取该执行器管理的所有工具"""
pass
@abstractmethod
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定工具"""
pass
@@ -0,0 +1,45 @@
"""工具系统的类型定义"""
from enum import Enum
from dataclasses import dataclass
from typing import Any, Dict, Optional, Callable, Awaitable
from abc import ABC, abstractmethod
class ToolType(Enum):
"""工具类型枚举"""
SERVER_PLUGIN = "server_plugin" # 服务端插件
SERVER_MCP = "server_mcp" # 服务端MCP
DEVICE_IOT = "device_iot" # 设备端IoT
DEVICE_MCP = "device_mcp" # 设备端MCP
class ToolAction(Enum):
"""工具执行后的动作类型"""
ERROR = "error" # 错误
NOT_FOUND = "not_found" # 工具未找到
RESPONSE = "response" # 直接回复
REQUEST_LLM = "request_llm" # 需要LLM处理
NONE = "none" # 无需特殊处理
@dataclass
class ToolResult:
"""工具执行结果"""
action: ToolAction
content: str # 结果内容
response: Optional[str] = None # 直接回复内容
error: Optional[str] = None # 错误信息
@dataclass
class ToolDefinition:
"""工具定义"""
name: str # 工具名称
description: Dict[str, Any] # 工具描述(OpenAI函数调用格式)
tool_type: ToolType # 工具类型
parameters: Optional[Dict[str, Any]] = None # 额外参数
@@ -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",
]
@@ -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)
@@ -0,0 +1,243 @@
"""设备端IoT工具执行器"""
import json
import asyncio
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
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]
) -> ToolResult:
"""执行设备端IoT工具"""
if not self.has_tool(tool_name):
return ToolResult(
action=ToolAction.NOT_FOUND, content=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 ToolResult(
action=ToolAction.RESPONSE,
content=str(value),
response=response,
)
else:
response_failure = arguments.get(
"response_failure", f"无法获取{device_name}的状态"
)
return ToolResult(
action=ToolAction.ERROR,
content=f"属性{property_name}不存在",
error=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 ToolResult(
action=ToolAction.REQUEST_LLM,
content=f"{device_name}操作执行成功,请继续处理剩余指令",
response=response_success,
)
return ToolResult(action=ToolAction.ERROR, content="无法解析IoT工具名称")
except Exception as e:
response_failure = arguments.get("response_failure", "操作失败")
return ToolResult(
action=ToolAction.ERROR, content=str(e), error=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 == device_name:
for property_item in value.properties:
if property_item["name"] == property_name:
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 == device_name:
for method in value.methods:
if method["name"] == method_name:
command = {
"name": device_name,
"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
@@ -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
@@ -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",
]
@@ -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)
@@ -0,0 +1,72 @@
"""设备端MCP工具执行器"""
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
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]
) -> ToolResult:
"""执行设备端MCP工具"""
if not hasattr(conn, "mcp_client") or not conn.mcp_client:
return ToolResult(
action=ToolAction.ERROR,
content="设备端MCP客户端未初始化",
error="设备端MCP客户端未初始化",
)
if not await conn.mcp_client.is_ready():
return ToolResult(
action=ToolAction.ERROR,
content="设备端MCP客户端未准备就绪",
error="设备端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)
return ToolResult(action=ToolAction.RESPONSE, content=str(result))
except ValueError as e:
return ToolResult(action=ToolAction.NOT_FOUND, content=str(e), error=str(e))
except Exception as e:
return ToolResult(action=ToolAction.ERROR, content=str(e), error=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)
@@ -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消息: {payload}")
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}"
)
@@ -0,0 +1,6 @@
"""服务端MCP工具模块"""
from .mcp_manager import ServerMCPManager
from .mcp_executor import ServerMCPExecutor
__all__ = ["ServerMCPManager", "ServerMCPExecutor"]
@@ -0,0 +1,82 @@
"""服务端MCP工具执行器"""
from typing import Dict, Any, Optional
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
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]
) -> ToolResult:
"""执行服务端MCP工具"""
if not self._initialized or not self.mcp_manager:
return ToolResult(
action=ToolAction.ERROR,
content="MCP管理器未初始化",
error="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 ToolResult(action=ToolAction.RESPONSE, content=str(result))
except ValueError as e:
return ToolResult(action=ToolAction.NOT_FOUND, content=str(e), error=str(e))
except Exception as e:
return ToolResult(action=ToolAction.ERROR, content=str(e), error=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 = f"mcp_{func_def.get('name', '')}"
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()
@@ -0,0 +1,100 @@
"""服务端MCP管理器"""
import asyncio
import os
import json
from typing import Dict, Any, List
from config.config_loader import get_project_dir
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class ServerMCPManager:
"""管理多个服务端MCP服务的集中管理器"""
def __init__(self, conn) -> None:
"""初始化MCP管理器"""
self.conn = conn
self.config_path = get_project_dir() + "data/.mcp_server_settings.json"
if not os.path.exists(self.config_path):
self.config_path = ""
logger.bind(tag=TAG).warning(
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
)
self.clients: Dict[str, Any] = {}
self.tools = []
def load_config(self) -> Dict[str, Any]:
"""加载MCP服务配置"""
if len(self.config_path) == 0:
return {}
try:
with open(self.config_path, "r", encoding="utf-8") as f:
config = json.load(f)
return config.get("mcpServers", {})
except Exception as e:
logger.bind(tag=TAG).error(
f"Error loading MCP config from {self.config_path}: {e}"
)
return {}
async def initialize_servers(self) -> None:
"""初始化所有MCP服务"""
config = self.load_config()
for name, srv_config in config.items():
if not srv_config.get("command") and not srv_config.get("url"):
logger.bind(tag=TAG).warning(
f"Skipping server {name}: neither command nor url specified"
)
continue
try:
# 这里可以添加真正的MCP客户端初始化逻辑
# 暂时使用简化版本
logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}")
# client = MCPClient(srv_config)
# await client.initialize()
# self.clients[name] = client
# client_tools = client.get_available_tools()
# self.tools.extend(client_tools)
except Exception as e:
logger.bind(tag=TAG).error(
f"Failed to initialize MCP server {name}: {e}"
)
def get_all_tools(self) -> List[Dict[str, Any]]:
"""获取所有服务的工具function定义"""
return self.tools
def is_mcp_tool(self, tool_name: str) -> bool:
"""检查是否是MCP工具"""
for tool in self.tools:
if (
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:
"""执行工具调用"""
logger.bind(tag=TAG).info(f"执行服务端MCP工具 {tool_name},参数: {arguments}")
# 这里可以添加真正的工具执行逻辑
# 暂时返回模拟结果
return f"服务端MCP工具 {tool_name} 执行结果"
async def cleanup_all(self) -> None:
"""关闭所有 MCP客户端"""
for name, client in list(self.clients.items()):
try:
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:
logger.bind(tag=TAG).error(f"关闭服务端MCP客户端 {name} 时出错: {e}")
self.clients.clear()
@@ -0,0 +1,5 @@
"""服务端插件工具模块"""
from .plugin_executor import ServerPluginExecutor
__all__ = ["ServerPluginExecutor"]
@@ -0,0 +1,102 @@
"""服务端插件工具执行器"""
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
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]
) -> ToolResult:
"""执行服务端插件工具"""
func_item = all_function_registry.get(tool_name)
if not func_item:
return ToolResult(
action=ToolAction.NOT_FOUND, content=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)
# 转换ActionResponse到ToolResult
if isinstance(result, ActionResponse):
if result.action == Action.ERROR:
return ToolResult(
action=ToolAction.ERROR,
content=result.result,
error=result.response,
)
elif result.action == Action.RESPONSE:
return ToolResult(
action=ToolAction.RESPONSE,
content=result.result,
response=result.response,
)
elif result.action == Action.REQLLM:
return ToolResult(
action=ToolAction.REQUEST_LLM,
content=result.result,
response=result.response,
)
else:
return ToolResult(
action=ToolAction.NONE,
content=result.result,
response=result.response,
)
else:
# 直接返回结果
return ToolResult(action=ToolAction.RESPONSE, content=str(result))
except Exception as e:
return ToolResult(action=ToolAction.ERROR, content=str(e), error=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
@@ -0,0 +1,188 @@
"""统一工具处理器,替换原有的FunctionHandler"""
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, ToolResult, ToolAction
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
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.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.finish_init = False
async def _initialize(self):
"""异步初始化"""
try:
# 自动导入插件模块
auto_import_modules("plugins_func.functions")
# 初始化服务端MCP
await self.server_mcp_executor.initialize()
# 初始化Home Assistant(如果需要)
self._initialize_home_assistant()
self.finish_init = True
self.logger.info("统一工具处理器初始化完成")
except Exception as e:
self.logger.error(f"统一工具处理器初始化失败: {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[ToolResult]:
"""处理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 ToolResult(
action=ToolAction.ERROR,
content="无法解析函数参数",
error="无法解析函数参数",
)
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 ToolResult(action=ToolAction.ERROR, content=str(e), error=str(e))
def _combine_responses(self, responses: List[ToolResult]) -> ToolResult:
"""合并多个函数调用的响应"""
if not responses:
return ToolResult(action=ToolAction.NONE, content="无响应")
# 如果有任何错误,返回第一个错误
for response in responses:
if response.action == ToolAction.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 = ToolAction.RESPONSE
for response in responses:
if response.action == ToolAction.REQUEST_LLM:
final_action = ToolAction.REQUEST_LLM
break
return ToolResult(
action=final_action,
content="; ".join(contents),
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()
self.logger.info("工具处理器清理完成")
except Exception as e:
self.logger.error(f"工具处理器清理失败: {e}")
@@ -0,0 +1,128 @@
"""统一工具管理器"""
import asyncio
from typing import Dict, List, Optional, Any
from config.logger import setup_logging
from .base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
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]) -> ToolResult:
"""执行工具调用"""
try:
# 查找工具类型
tool_type = self.get_tool_type(tool_name)
if not tool_type:
return ToolResult(
action=ToolAction.NOT_FOUND,
content=f"工具 {tool_name} 不存在",
error=f"工具 {tool_name} 不存在"
)
# 获取对应的执行器
executor = self.executors.get(tool_type)
if not executor:
return ToolResult(
action=ToolAction.ERROR,
content=f"工具类型 {tool_type.value} 的执行器未注册",
error=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 ToolResult(
action=ToolAction.ERROR,
content=str(e),
error=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