mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-22 07:03:53 +08:00
@@ -104,7 +104,8 @@ wakeup_words:
|
||||
- "喵喵同学"
|
||||
- "小滨小滨"
|
||||
- "小冰小冰"
|
||||
|
||||
# MCP接入点地址
|
||||
mcp_endpoint: 你的接入点 websocket地址
|
||||
# 插件的基础配置
|
||||
plugins:
|
||||
# 获取天气插件的配置,这里填写你的api_key
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.providers.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.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:
|
||||
|
||||
@@ -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.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"])
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""基础工具定义模块"""
|
||||
|
||||
from .tool_types import ToolType, ToolDefinition
|
||||
from .tool_executor import ToolExecutor
|
||||
|
||||
__all__ = ["ToolType", "ToolDefinition", "ToolExecutor"]
|
||||
@@ -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
|
||||
@@ -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 # 额外参数
|
||||
@@ -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,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
|
||||
@@ -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,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)
|
||||
+28
-32
@@ -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}"
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -0,0 +1,7 @@
|
||||
"""服务端MCP工具模块"""
|
||||
|
||||
from .mcp_manager import ServerMCPManager
|
||||
from .mcp_executor import ServerMCPExecutor
|
||||
from .mcp_client import ServerMCPClient
|
||||
|
||||
__all__ = ["ServerMCPManager", "ServerMCPExecutor", "ServerMCPClient"]
|
||||
+58
-14
@@ -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
|
||||
@@ -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()
|
||||
+50
-80
@@ -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()
|
||||
@@ -0,0 +1,5 @@
|
||||
"""服务端插件工具模块"""
|
||||
|
||||
from .plugin_executor import ServerPluginExecutor
|
||||
|
||||
__all__ = ["ServerPluginExecutor"]
|
||||
@@ -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
|
||||
@@ -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}")
|
||||
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -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 # 直接回复的内容
|
||||
|
||||
Reference in New Issue
Block a user