update:优化对话数据上传

This commit is contained in:
hrz
2025-05-02 00:07:36 +08:00
parent 456b1eddcc
commit f0e353cc1a
22 changed files with 295 additions and 340 deletions
+38 -74
View File
@@ -17,7 +17,6 @@ from core.handle.textHandle import handleTextMessage
from core.utils.util import (
get_string_no_punctuation_or_emoji,
extract_json_from_string,
get_ip_info,
initialize_modules,
)
from concurrent.futures import ThreadPoolExecutor, TimeoutError
@@ -30,7 +29,7 @@ from core.mcp.manager import MCPManager
from config.config_loader import get_private_config_from_api
from config.manage_api_client import DeviceNotFoundException, DeviceBindException
from core.utils.output_counter import add_device_output
from core.handle.ttsReportHandle import enqueue_tts_report
from core.handle.ttsReportHandle import enqueue_tts_report, report_tts
TAG = __name__
@@ -43,7 +42,15 @@ class TTSException(RuntimeError):
class ConnectionHandler:
def __init__(
self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _memory, _intent, server=None
self,
config: Dict[str, Any],
_vad,
_asr,
_llm,
_tts,
_memory,
_intent,
server=None,
):
self.config = config
self.server = server
@@ -52,6 +59,7 @@ class ConnectionHandler:
self.need_bind = False
self.bind_code = None
self.read_config_from_api = self.config.get("read_config_from_api", False)
self.websocket = None
self.headers = None
@@ -74,11 +82,8 @@ class ConnectionHandler:
self.audio_play_queue = queue.Queue()
self.executor = ThreadPoolExecutor(max_workers=10)
# 上报线程标志
self.session_open_time = time.time()
# 上报线程
self.tts_report_queue = queue.Queue()
self.asr_report_queue = queue.Queue()
self.asr_report_thread = None
self.tts_report_thread = None
# 依赖的组件
@@ -275,19 +280,27 @@ class ConnectionHandler:
)
except Exception as e:
self.logger.bind(tag=TAG).error(f"模块初始化失败: {str(e)}")
await self.websocket.send(json.dumps({
"type": "config_update_response",
"status": "error",
"message": f"模块初始化失败: {str(e)}"
}))
await self.websocket.send(
json.dumps(
{
"type": "config_update_response",
"status": "error",
"message": f"模块初始化失败: {str(e)}",
}
)
)
return
# 返回成功响应
await self.websocket.send(json.dumps({
"type": "config_update_response",
"status": "success",
"message": f"已更新配置: {', '.join(updated_modules)}"
}))
await self.websocket.send(
json.dumps(
{
"type": "config_update_response",
"status": "success",
"message": f"已更新配置: {', '.join(updated_modules)}",
}
)
)
def _initialize_components(self, private_config):
"""初始化组件"""
@@ -305,26 +318,18 @@ class ConnectionHandler:
def _init_report_threads(self):
"""初始化ASR和TTS上报线程"""
if self.asr_report_thread is None or not self.asr_report_thread.is_alive():
self.asr_report_thread = threading.Thread(
target=self._asr_report_worker,
daemon=True
)
self.asr_report_thread.start()
self.logger.bind(tag=TAG).info("ASR上报线程已启动")
if not self.read_config_from_api:
return
if self.tts_report_thread is None or not self.tts_report_thread.is_alive():
self.tts_report_thread = threading.Thread(
target=self._tts_report_worker,
daemon=True
target=self._tts_report_worker, daemon=True
)
self.tts_report_thread.start()
self.logger.bind(tag=TAG).info("TTS上报线程已启动")
def _initialize_private_config(self):
read_config_from_api = self.config.get("read_config_from_api", False)
"""如果是从配置文件获取,则进行二次实例化"""
if not read_config_from_api:
if not self.read_config_from_api:
return
"""从接口获取差异化的配置进行二次实例化,非全量重新实例化"""
try:
@@ -880,10 +885,9 @@ class ConnectionHandler:
f"TTS生成:文件路径: {tts_file}"
)
if os.path.exists(tts_file):
opus_datas, _ = self.tts.audio_to_opus_data(tts_file)
# 在这里上报TTS数据(使用文件路径)
enqueue_tts_report(self, text, tts_file)
opus_datas, duration = self.tts.audio_to_opus_data(tts_file)
enqueue_tts_report(self, 2, text, opus_datas)
else:
self.logger.bind(tag=TAG).error(
f"TTS出错:文件不存在{tts_file}"
@@ -939,43 +943,8 @@ class ConnectionHandler:
f"audio_play_priority priority_thread: {text} {e}"
)
def _asr_report_worker(self):
"""ASR上报工作线程"""
# 提前导入避免循环引用问题
from core.handle.asrReportHandle import report_asr
while not self.stop_event.is_set():
try:
# 从队列获取数据,设置超时以便定期检查停止事件
item = self.asr_report_queue.get(timeout=1)
if item is None: # 检测毒丸对象
break
text, file_path = item
try:
# 执行上报(传入文件路径)
await_result = report_asr(self, text, file_path)
# 使用asyncio.run_coroutine_threadsafe执行异步操作
future = asyncio.run_coroutine_threadsafe(await_result, self.loop)
future.result()
except Exception as e:
self.logger.bind(tag=TAG).error(f"ASR上报线程异常: {e}")
finally:
# 标记任务完成
self.asr_report_queue.task_done()
except queue.Empty:
continue
except Exception as e:
self.logger.bind(tag=TAG).error(f"ASR上报工作线程异常: {e}")
self.logger.bind(tag=TAG).info("ASR上报线程已退出")
def _tts_report_worker(self):
"""TTS上报工作线程"""
# 提前导入避免循环引用问题
from core.handle.ttsReportHandle import report_tts
while not self.stop_event.is_set():
try:
@@ -984,15 +953,11 @@ class ConnectionHandler:
if item is None: # 检测毒丸对象
break
text, audio_data = item
type, text, audio_data = item
try:
# 执行上报(传入二进制数据)
await_result = report_tts(self, text, audio_data)
# 使用asyncio.run_coroutine_threadsafe执行异步操作
future = asyncio.run_coroutine_threadsafe(await_result, self.loop)
future.result()
report_tts(self, type, text, audio_data)
except Exception as e:
self.logger.bind(tag=TAG).error(f"TTS上报线程异常: {e}")
finally:
@@ -1051,7 +1016,6 @@ class ConnectionHandler:
self.executor = None
# 添加毒丸对象到上报队列确保线程退出
self.asr_report_queue.put(None)
self.tts_report_queue.put(None)
# 清空任务队列