update:优化工具回复

This commit is contained in:
hrz
2025-06-26 11:27:23 +08:00
parent 2f78acaf4d
commit 0f36daf1fd
12 changed files with 108 additions and 195 deletions
+2 -39
View File
@@ -17,7 +17,6 @@ from core.utils.util import (
filter_sensitive_info,
)
from typing import Dict, Any
from core.providers.tools.base import ToolAction
from core.utils.modules_initialize import (
initialize_modules,
initialize_tts,
@@ -725,15 +724,12 @@ class ConnectionHandler:
}
# 使用统一工具处理器处理所有工具调用
tool_result = asyncio.run_coroutine_threadsafe(
result = asyncio.run_coroutine_threadsafe(
self.func_handler.handle_llm_function_call(
self, function_call_data
),
self.loop,
).result()
# 转换ToolResult为ActionResponse
result = self._convert_tool_result_to_action_response(tool_result)
self._handle_function_result(result, function_call_data)
# 存储对话内容
@@ -756,39 +752,6 @@ class ConnectionHandler:
return True
def _convert_tool_result_to_action_response(self, tool_result):
"""转换ToolResult为ActionResponse"""
if tool_result.action == ToolAction.ERROR:
return ActionResponse(
action=Action.ERROR,
result=tool_result.error or tool_result.content,
response=tool_result.response or tool_result.content,
)
elif tool_result.action == ToolAction.NOT_FOUND:
return ActionResponse(
action=Action.NOTFOUND,
result=tool_result.content,
response=tool_result.response or tool_result.content,
)
elif tool_result.action == ToolAction.RESPONSE:
return ActionResponse(
action=Action.RESPONSE,
result=tool_result.content,
response=tool_result.response or tool_result.content,
)
elif tool_result.action == ToolAction.REQUEST_LLM:
return ActionResponse(
action=Action.REQLLM,
result=tool_result.content,
response=tool_result.response or "",
)
else:
return ActionResponse(
action=Action.NONE,
result=tool_result.content,
response=tool_result.response or "",
)
def _handle_function_result(self, result, function_call_data):
if result.action == Action.RESPONSE: # 直接回复前端
text = result.response
@@ -828,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:
@@ -108,15 +108,12 @@ async def process_intent_result(conn, intent_result, original_text):
# 使用统一工具处理器处理所有工具调用
try:
tool_result = asyncio.run_coroutine_threadsafe(
result = asyncio.run_coroutine_threadsafe(
conn.func_handler.handle_llm_function_call(
conn, function_call_data
),
conn.loop,
).result()
# 转换ToolResult为ActionResponse
result = conn._convert_tool_result_to_action_response(tool_result)
except Exception as e:
conn.logger.bind(tag=TAG).error(f"工具调用失败: {e}")
result = ActionResponse(
@@ -1,6 +1,6 @@
"""基础工具定义模块"""
from .tool_types import ToolType, ToolAction, ToolResult, ToolDefinition
from .tool_types import ToolType, ToolDefinition
from .tool_executor import ToolExecutor
__all__ = ["ToolType", "ToolAction", "ToolResult", "ToolDefinition", "ToolExecutor"]
__all__ = ["ToolType", "ToolDefinition", "ToolExecutor"]
@@ -2,7 +2,8 @@
from abc import ABC, abstractmethod
from typing import Dict, Any
from .tool_types import ToolDefinition, ToolResult
from .tool_types import ToolDefinition
from plugins_func.register import ActionResponse
class ToolExecutor(ABC):
@@ -11,7 +12,7 @@ class ToolExecutor(ABC):
@abstractmethod
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ToolResult:
) -> ActionResponse:
"""执行工具调用"""
pass
@@ -1,9 +1,10 @@
"""工具系统的类型定义"""
from enum import Enum
from dataclasses import dataclass
from typing import Any, Dict, Optional, Callable, Awaitable
from abc import ABC, abstractmethod
from typing import Any, Dict, Optional
from plugins_func.register import Action
class ToolType(Enum):
@@ -15,26 +16,6 @@ class ToolType(Enum):
DEVICE_MCP = "device_mcp" # 设备端MCP
class ToolAction(Enum):
"""工具执行后的动作类型"""
ERROR = "error" # 错误
NOT_FOUND = "not_found" # 工具未找到
RESPONSE = "response" # 直接回复
REQUEST_LLM = "request_llm" # 需要LLM处理
NONE = "none" # 无需特殊处理
@dataclass
class ToolResult:
"""工具执行结果"""
action: ToolAction
content: str # 结果内容
response: Optional[str] = None # 直接回复内容
error: Optional[str] = None # 错误信息
@dataclass
class ToolDefinition:
"""工具定义"""
@@ -3,7 +3,8 @@
import json
import asyncio
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import Action, ActionResponse
class DeviceIoTExecutor(ToolExecutor):
@@ -15,11 +16,11 @@ class DeviceIoTExecutor(ToolExecutor):
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ToolResult:
) -> ActionResponse:
"""执行设备端IoT工具"""
if not self.has_tool(tool_name):
return ToolResult(
action=ToolAction.NOT_FOUND, content=f"IoT工具 {tool_name} 不存在"
return ActionResponse(
action=Action.NOTFOUND, response=f"IoT工具 {tool_name} 不存在"
)
try:
@@ -39,19 +40,16 @@ class DeviceIoTExecutor(ToolExecutor):
)
response = response_success.replace("{value}", str(value))
return ToolResult(
action=ToolAction.RESPONSE,
content=str(value),
return ActionResponse(
action=Action.RESPONSE,
response=response,
)
else:
response_failure = arguments.get(
"response_failure", f"无法获取{device_name}的状态"
)
return ToolResult(
action=ToolAction.ERROR,
content=f"属性{property_name}不存在",
error=response_failure,
return ActionResponse(
action=Action.ERROR, response=response_failure
)
else:
# 控制操作:devicename_method
@@ -90,19 +88,16 @@ class DeviceIoTExecutor(ToolExecutor):
)
break
return ToolResult(
action=ToolAction.REQUEST_LLM,
content=f"{device_name}操作执行成功,请继续处理剩余指令",
response=response_success,
return ActionResponse(
action=Action.REQLLM,
result=response_success,
)
return ToolResult(action=ToolAction.ERROR, content="无法解析IoT工具名称")
return ActionResponse(action=Action.ERROR, response="无法解析IoT工具名称")
except Exception as e:
response_failure = arguments.get("response_failure", "操作失败")
return ToolResult(
action=ToolAction.ERROR, content=str(e), error=response_failure
)
return ActionResponse(action=Action.ERROR, response=response_failure)
async def _get_iot_status(self, device_name: str, property_name: str):
"""获取IoT设备状态"""
@@ -1,7 +1,8 @@
"""设备端MCP工具执行器"""
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import Action, ActionResponse
from .mcp_handler import call_mcp_tool
@@ -13,20 +14,18 @@ class DeviceMCPExecutor(ToolExecutor):
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ToolResult:
) -> ActionResponse:
"""执行设备端MCP工具"""
if not hasattr(conn, "mcp_client") or not conn.mcp_client:
return ToolResult(
action=ToolAction.ERROR,
content="设备端MCP客户端未初始化",
error="设备端MCP客户端未初始化",
return ActionResponse(
action=Action.ERROR,
response="设备端MCP客户端未初始化",
)
if not await conn.mcp_client.is_ready():
return ToolResult(
action=ToolAction.ERROR,
content="设备端MCP客户端未准备就绪",
error="设备端MCP客户端未准备就绪",
return ActionResponse(
action=Action.ERROR,
response="设备端MCP客户端未准备就绪",
)
try:
@@ -38,12 +37,12 @@ class DeviceMCPExecutor(ToolExecutor):
# 调用设备端MCP工具
result = await call_mcp_tool(conn, conn.mcp_client, tool_name, args_str)
return ToolResult(action=ToolAction.RESPONSE, content=str(result))
return ActionResponse(action=Action.REQLLM, result=str(result))
except ValueError as e:
return ToolResult(action=ToolAction.NOT_FOUND, content=str(e), error=str(e))
return ActionResponse(action=Action.NOTFOUND, response=str(e))
except Exception as e:
return ToolResult(action=ToolAction.ERROR, content=str(e), error=str(e))
return ActionResponse(action=Action.ERROR, response=str(e))
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有设备端MCP工具"""
@@ -1,7 +1,8 @@
"""服务端MCP工具执行器"""
from typing import Dict, Any, Optional
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import Action, ActionResponse
from .mcp_manager import ServerMCPManager
@@ -22,13 +23,12 @@ class ServerMCPExecutor(ToolExecutor):
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ToolResult:
) -> ActionResponse:
"""执行服务端MCP工具"""
if not self._initialized or not self.mcp_manager:
return ToolResult(
action=ToolAction.ERROR,
content="MCP管理器未初始化",
error="MCP管理器未初始化",
return ActionResponse(
action=Action.ERROR,
response="MCP管理器未初始化",
)
try:
@@ -39,12 +39,18 @@ class ServerMCPExecutor(ToolExecutor):
result = await self.mcp_manager.execute_tool(actual_tool_name, arguments)
return ToolResult(action=ToolAction.RESPONSE, content=str(result))
return ActionResponse(action=Action.REQLLM, result=str(result))
except ValueError as e:
return ToolResult(action=ToolAction.NOT_FOUND, content=str(e), error=str(e))
return ActionResponse(
action=Action.NOTFOUND,
response=str(e),
)
except Exception as e:
return ToolResult(action=ToolAction.ERROR, content=str(e), error=str(e))
return ActionResponse(
action=Action.ERROR,
response=str(e),
)
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有服务端MCP工具"""
@@ -1,7 +1,7 @@
"""服务端插件工具执行器"""
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import all_function_registry, Action, ActionResponse
@@ -14,12 +14,12 @@ class ServerPluginExecutor(ToolExecutor):
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ToolResult:
) -> ActionResponse:
"""执行服务端插件工具"""
func_item = all_function_registry.get(tool_name)
if not func_item:
return ToolResult(
action=ToolAction.NOT_FOUND, content=f"插件函数 {tool_name} 不存在"
return ActionResponse(
action=Action.NOTFOUND, response=f"插件函数 {tool_name} 不存在"
)
try:
@@ -38,38 +38,13 @@ class ServerPluginExecutor(ToolExecutor):
# 默认不传conn参数
result = func_item.func(**arguments)
# 转换ActionResponse到ToolResult
if isinstance(result, ActionResponse):
if result.action == Action.ERROR:
return ToolResult(
action=ToolAction.ERROR,
content=result.result,
error=result.response,
)
elif result.action == Action.RESPONSE:
return ToolResult(
action=ToolAction.RESPONSE,
content=result.result,
response=result.response,
)
elif result.action == Action.REQLLM:
return ToolResult(
action=ToolAction.REQUEST_LLM,
content=result.result,
response=result.response,
)
else:
return ToolResult(
action=ToolAction.NONE,
content=result.result,
response=result.response,
)
else:
# 直接返回结果
return ToolResult(action=ToolAction.RESPONSE, content=str(result))
return result
except Exception as e:
return ToolResult(action=ToolAction.ERROR, content=str(e), error=str(e))
return ActionResponse(
action=Action.ERROR,
response=str(e),
)
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有注册的服务端插件工具"""
@@ -5,7 +5,8 @@ from typing import Dict, List, Any, Optional
from config.logger import setup_logging
from plugins_func.loadplugins import auto_import_modules
from .base import ToolType, ToolResult, ToolAction
from .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
@@ -97,7 +98,7 @@ class UnifiedToolHandler:
async def handle_llm_function_call(
self, conn, function_call_data: Dict[str, Any]
) -> Optional[ToolResult]:
) -> Optional[ActionResponse]:
"""处理LLM函数调用"""
try:
# 处理多函数调用
@@ -120,10 +121,9 @@ class UnifiedToolHandler:
arguments = json.loads(arguments) if arguments else {}
except json.JSONDecodeError:
self.logger.error(f"无法解析函数参数: {arguments}")
return ToolResult(
action=ToolAction.ERROR,
content="无法解析函数参数",
error="无法解析函数参数",
return ActionResponse(
action=Action.ERROR,
response="无法解析函数参数",
)
self.logger.debug(f"调用函数: {function_name}, 参数: {arguments}")
@@ -134,16 +134,16 @@ class UnifiedToolHandler:
except Exception as e:
self.logger.error(f"处理function call错误: {e}")
return ToolResult(action=ToolAction.ERROR, content=str(e), error=str(e))
return ActionResponse(action=Action.ERROR, response=str(e))
def _combine_responses(self, responses: List[ToolResult]) -> ToolResult:
def _combine_responses(self, responses: List[ActionResponse]) -> ActionResponse:
"""合并多个函数调用的响应"""
if not responses:
return ToolResult(action=ToolAction.NONE, content="无响应")
return ActionResponse(action=Action.NONE, response="无响应")
# 如果有任何错误,返回第一个错误
for response in responses:
if response.action == ToolAction.ERROR:
if response.action == Action.ERROR:
return response
# 合并所有成功的响应
@@ -157,15 +157,15 @@ class UnifiedToolHandler:
responses_text.append(response.response)
# 确定最终的动作类型
final_action = ToolAction.RESPONSE
final_action = Action.RESPONSE
for response in responses:
if response.action == ToolAction.REQUEST_LLM:
final_action = ToolAction.REQUEST_LLM
if response.action == Action.REQLLM:
final_action = Action.REQLLM
break
return ToolResult(
return ActionResponse(
action=final_action,
content="; ".join(contents),
result="; ".join(contents) if contents else None,
response="; ".join(responses_text) if responses_text else None,
)
@@ -1,37 +1,37 @@
"""统一工具管理器"""
import asyncio
from typing import Dict, List, Optional, Any
from config.logger import setup_logging
from .base import ToolType, ToolDefinition, ToolResult, ToolExecutor, ToolAction
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:
@@ -42,79 +42,75 @@ class ToolManager:
all_tools[name] = definition
except Exception as e:
self.logger.error(f"获取{tool_type.value}工具时出错: {e}")
self._cached_tools = all_tools
return all_tools
def get_function_descriptions(self) -> List[Dict[str, Any]]:
"""获取所有工具的函数描述(OpenAI格式)"""
if self._cached_function_descriptions is not None:
return self._cached_function_descriptions
descriptions = []
tools = self.get_all_tools()
for tool_definition in tools.values():
descriptions.append(tool_definition.description)
self._cached_function_descriptions = descriptions
return descriptions
def has_tool(self, tool_name: str) -> bool:
"""检查是否存在指定工具"""
tools = self.get_all_tools()
return tool_name in tools
def get_tool_type(self, tool_name: str) -> Optional[ToolType]:
"""获取工具类型"""
tools = self.get_all_tools()
tool_def = tools.get(tool_name)
return tool_def.tool_type if tool_def else None
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> ToolResult:
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 ToolResult(
action=ToolAction.NOT_FOUND,
content=f"工具 {tool_name} 不存在",
error=f"工具 {tool_name} 不存在"
return ActionResponse(
action=Action.NOTFOUND,
response=f"工具 {tool_name} 不存在",
)
# 获取对应的执行器
executor = self.executors.get(tool_type)
if not executor:
return ToolResult(
action=ToolAction.ERROR,
content=f"工具类型 {tool_type.value} 的执行器未注册",
error=f"工具类型 {tool_type.value} 的执行器未注册"
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 ToolResult(
action=ToolAction.ERROR,
content=str(e),
error=str(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 = {}
@@ -125,4 +121,4 @@ class ToolManager:
except Exception as e:
self.logger.error(f"获取{tool_type.value}工具统计时出错: {e}")
stats[tool_type.value] = 0
return stats
return stats
+1 -1
View File
@@ -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 # 直接回复的内容