mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-22 07:03:53 +08:00
fix:修复 intent_llm 模式下工具调用不上报的问题
This commit is contained in:
@@ -24,7 +24,7 @@ from core.utils.modules_initialize import (
|
|||||||
initialize_tts,
|
initialize_tts,
|
||||||
initialize_asr,
|
initialize_asr,
|
||||||
)
|
)
|
||||||
from core.handle.reportHandle import report
|
from core.handle.reportHandle import report, enqueue_tool_report
|
||||||
from core.providers.tts.default import DefaultTTS
|
from core.providers.tts.default import DefaultTTS
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from core.utils.dialogue import Message, Dialogue
|
from core.utils.dialogue import Message, Dialogue
|
||||||
@@ -1090,20 +1090,9 @@ class ConnectionHandler:
|
|||||||
f"function_name={tool_call_data['name']}, function_id={tool_call_data['id']}, function_arguments={tool_call_data['arguments']}"
|
f"function_name={tool_call_data['name']}, function_id={tool_call_data['id']}, function_arguments={tool_call_data['arguments']}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# 构建工具调用的显示文本,格式如: get_weather({"location": "北京"})
|
# 使用公共方法上报工具调用
|
||||||
tool_input = json.loads(tool_call_data.get("arguments") or "{}")
|
tool_input = json.loads(tool_call_data.get("arguments") or "{}")
|
||||||
tool_text = json.dumps(
|
enqueue_tool_report(self, tool_call_data['name'], tool_input)
|
||||||
[
|
|
||||||
{
|
|
||||||
"type": "tool",
|
|
||||||
"text": f"{tool_call_data['name']}({json.dumps(tool_input, ensure_ascii=False)})",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
# 上报工具调用的内容(使用chatType=3表示工具调用)
|
|
||||||
tool_call_timestamp = int(time.time())
|
|
||||||
self.report_queue.put((3, tool_text, None, tool_call_timestamp))
|
|
||||||
|
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
self.func_handler.handle_llm_function_call(
|
self.func_handler.handle_llm_function_call(
|
||||||
@@ -1111,21 +1100,16 @@ class ConnectionHandler:
|
|||||||
),
|
),
|
||||||
self.loop,
|
self.loop,
|
||||||
)
|
)
|
||||||
futures_with_data.append((future, tool_call_data))
|
futures_with_data.append((future, tool_call_data, tool_input))
|
||||||
|
|
||||||
# 等待协程结束(实际等待时长为最慢的那个)
|
# 等待协程结束(实际等待时长为最慢的那个)
|
||||||
tool_results = []
|
tool_results = []
|
||||||
for future, tool_call_data in futures_with_data:
|
for future, tool_call_data, tool_input in futures_with_data:
|
||||||
result = future.result()
|
result = future.result()
|
||||||
tool_results.append((result, tool_call_data))
|
tool_results.append((result, tool_call_data))
|
||||||
|
|
||||||
# 工具执行完成后,单独上报结果(时间戳+1确保在工具调用之后)
|
# 使用公共方法上报工具调用结果
|
||||||
tool_result_text = str(result.result)
|
enqueue_tool_report(self, tool_call_data['name'], tool_input, str(result.result) if result.result else None, report_tool_call=False)
|
||||||
|
|
||||||
# 格式化为 {"result": ...}
|
|
||||||
result_display = f'{{"result":"{tool_result_text}"}}'
|
|
||||||
result_content = json.dumps([{"type": "tool_result", "text": result_display}], ensure_ascii=False)
|
|
||||||
self.report_queue.put((3, result_content, None, tool_call_timestamp + 1))
|
|
||||||
|
|
||||||
# 统一处理工具调用结果
|
# 统一处理工具调用结果
|
||||||
if tool_results:
|
if tool_results:
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from core.providers.tts.dto.dto import ContentType
|
|||||||
from core.handle.helloHandle import checkWakeupWords
|
from core.handle.helloHandle import checkWakeupWords
|
||||||
from plugins_func.register import Action, ActionResponse
|
from plugins_func.register import Action, ActionResponse
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
|
from core.handle.reportHandle import enqueue_tool_report
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType
|
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType
|
||||||
|
|
||||||
@@ -141,6 +142,17 @@ async def process_intent_result(
|
|||||||
await send_stt_message(conn, original_text)
|
await send_stt_message(conn, original_text)
|
||||||
conn.client_abort = False
|
conn.client_abort = False
|
||||||
|
|
||||||
|
# 准备工具调用参数
|
||||||
|
tool_input = {}
|
||||||
|
if function_args:
|
||||||
|
if isinstance(function_args, str):
|
||||||
|
tool_input = json.loads(function_args) if function_args else {}
|
||||||
|
elif isinstance(function_args, dict):
|
||||||
|
tool_input = function_args
|
||||||
|
|
||||||
|
# 上报工具调用
|
||||||
|
enqueue_tool_report(conn, function_name, tool_input)
|
||||||
|
|
||||||
# 使用executor执行函数调用和结果处理
|
# 使用executor执行函数调用和结果处理
|
||||||
def process_function_call():
|
def process_function_call():
|
||||||
conn.dialogue.put(Message(role="user", content=original_text))
|
conn.dialogue.put(Message(role="user", content=original_text))
|
||||||
@@ -159,7 +171,10 @@ async def process_intent_result(
|
|||||||
action=Action.ERROR, result=str(e), response=str(e)
|
action=Action.ERROR, result=str(e), response=str(e)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 上报工具调用结果
|
||||||
if result:
|
if result:
|
||||||
|
enqueue_tool_report(conn, function_name, tool_input, str(result.result) if result.result else None, report_tool_call=False)
|
||||||
|
|
||||||
if result.action == Action.RESPONSE: # 直接回复前端
|
if result.action == Action.RESPONSE: # 直接回复前端
|
||||||
text = result.response
|
text = result.response
|
||||||
if text is not None:
|
if text is not None:
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ TTS上报功能已集成到ConnectionHandler类中。
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import time
|
import time
|
||||||
|
import json
|
||||||
import opuslib_next
|
import opuslib_next
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
@@ -132,6 +133,45 @@ def enqueue_tts_report(conn: "ConnectionHandler", text, opus_data):
|
|||||||
conn.logger.bind(tag=TAG).error(f"加入TTS上报队列失败: {text}, {e}")
|
conn.logger.bind(tag=TAG).error(f"加入TTS上报队列失败: {text}, {e}")
|
||||||
|
|
||||||
|
|
||||||
|
def enqueue_tool_report(conn: "ConnectionHandler", tool_name: str, tool_input: dict, tool_result: str = None, report_tool_call: bool = True):
|
||||||
|
"""将工具调用数据加入上报队列
|
||||||
|
|
||||||
|
Args:
|
||||||
|
conn: 连接对象
|
||||||
|
tool_name: 工具名称
|
||||||
|
tool_input: 工具输入参数
|
||||||
|
tool_result: 工具执行结果(可选)
|
||||||
|
report_tool_call: 是否上报工具调用本身,默认True;仅上报结果时设为False
|
||||||
|
"""
|
||||||
|
if not conn.read_config_from_api or conn.need_bind:
|
||||||
|
return
|
||||||
|
if conn.chat_history_conf == 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
timestamp = int(time.time())
|
||||||
|
|
||||||
|
# 构建工具调用内容
|
||||||
|
if report_tool_call:
|
||||||
|
tool_text = json.dumps(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"type": "tool",
|
||||||
|
"text": f"{tool_name}({json.dumps(tool_input, ensure_ascii=False)})",
|
||||||
|
}
|
||||||
|
]
|
||||||
|
)
|
||||||
|
conn.report_queue.put((3, tool_text, None, timestamp))
|
||||||
|
|
||||||
|
# 构建工具结果内容
|
||||||
|
if tool_result:
|
||||||
|
result_display = f'{{"result":"{str(tool_result)}"}}'
|
||||||
|
result_content = json.dumps([{"type": "tool_result", "text": result_display}], ensure_ascii=False)
|
||||||
|
conn.report_queue.put((3, result_content, None, timestamp + 1))
|
||||||
|
except Exception as e:
|
||||||
|
conn.logger.bind(tag=TAG).error(f"加入工具上报队列失败: {e}")
|
||||||
|
|
||||||
|
|
||||||
def enqueue_asr_report(conn: "ConnectionHandler", text, opus_data):
|
def enqueue_asr_report(conn: "ConnectionHandler", text, opus_data):
|
||||||
if not conn.read_config_from_api or conn.need_bind or not conn.report_asr_enable:
|
if not conn.read_config_from_api or conn.need_bind or not conn.report_asr_enable:
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -168,10 +168,10 @@ class ASRProviderBase(ABC):
|
|||||||
self.stop_ws_connection()
|
self.stop_ws_connection()
|
||||||
|
|
||||||
if text_len > 0:
|
if text_len > 0:
|
||||||
# 使用自定义模块进行上报
|
|
||||||
await startToChat(conn, enhanced_text)
|
|
||||||
audio_snapshot = asr_audio_task.copy()
|
audio_snapshot = asr_audio_task.copy()
|
||||||
enqueue_asr_report(conn, enhanced_text, audio_snapshot)
|
enqueue_asr_report(conn, enhanced_text, audio_snapshot)
|
||||||
|
# 使用自定义模块进行上报
|
||||||
|
await startToChat(conn, enhanced_text)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"处理语音停止失败: {e}")
|
logger.bind(tag=TAG).error(f"处理语音停止失败: {e}")
|
||||||
import traceback
|
import traceback
|
||||||
|
|||||||
Reference in New Issue
Block a user