add: 语音流式输出

This commit is contained in:
lizhongxiang
2025-03-05 17:26:32 +08:00
parent 8b151d07c2
commit fa75f56ffb
5 changed files with 305 additions and 16 deletions
+79 -10
View File
@@ -13,7 +13,7 @@ from core.utils.dialogue import Message, Dialogue
from core.handle.textHandle import handleTextMessage
from core.utils.util import get_string_no_punctuation_or_emoji
from concurrent.futures import ThreadPoolExecutor, TimeoutError
from core.handle.sendAudioHandle import sendAudioMessage
from core.handle.sendAudioHandle import sendAudioMessage, sendAudioMessageStream
from core.handle.receiveAudioHandle import handleAudioMessage
from config.private_config import PrivateConfig
from core.auth import AuthMiddleware, AuthenticationError
@@ -32,6 +32,8 @@ class ConnectionHandler:
self.logger = setup_logging()
self.auth = AuthMiddleware(config)
self.tts_stream = self.config.get("TTS_SET", {}).get("TTS_STREAM", False)
self.websocket = None
self.headers = None
self.session_id = None
@@ -46,8 +48,10 @@ class ConnectionHandler:
self.loop = asyncio.get_event_loop()
self.stop_event = threading.Event()
self.tts_queue = queue.Queue()
self.tts_queue_stream = queue.Queue()
self.audio_play_queue = queue.Queue()
self.executor = ThreadPoolExecutor(max_workers=10)
max_workers = self.config.get("TTS_SET", {}).get("MAX_WORKERS", 10)
self.executor = ThreadPoolExecutor(max_workers=max_workers)
# 依赖的组件
self.vad = _vad
@@ -74,6 +78,7 @@ class ConnectionHandler:
# tts相关变量
self.tts_first_text_index = -1
self.tts_last_text_index = -1
self.tts_duration = 0
# iot相关变量
self.iot_descriptors = {}
@@ -262,8 +267,17 @@ class ConnectionHandler:
# segment_text = " "
text_index += 1
self.recode_first_last_text(segment_text, text_index)
future = self.executor.submit(self.speak_and_play, segment_text, text_index)
self.tts_queue.put(future)
if self.tts_stream:
stream_queue = queue.Queue()
self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue)
self.tts_queue_stream.put({
"text": segment_text,
"chunk_queque": stream_queue,
"text_index": text_index
})
else:
future = self.executor.submit(self.speak_and_play, segment_text, text_index)
self.tts_queue.put(future)
processed_chars += len(segment_text_raw) # 更新已处理字符位置
# 处理最后剩余的文本
@@ -274,8 +288,17 @@ class ConnectionHandler:
if segment_text:
text_index += 1
self.recode_first_last_text(segment_text, text_index)
future = self.executor.submit(self.speak_and_play, segment_text, text_index)
self.tts_queue.put(future)
if self.tts_stream:
stream_queue = queue.Queue()
self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue, text_index)
self.tts_queue_stream.put({
"text": segment_text,
"chunk_queque": stream_queue,
"text_index": text_index
})
else:
future = self.executor.submit(self.speak_and_play, segment_text, text_index)
self.tts_queue.put(future)
self.llm_finish_task = True
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
@@ -283,6 +306,12 @@ class ConnectionHandler:
return True
def _tts_priority_thread(self):
if self.tts_stream:
self._tts_priority_thread_stream()
else:
self._tts_priority_thread_no_stream()
def _tts_priority_thread_no_stream(self):
while not self.stop_event.is_set():
text = None
try:
@@ -322,14 +351,47 @@ class ConnectionHandler:
)
self.logger.bind(tag=TAG).error(f"tts_priority priority_thread: {text} {e}")
def _tts_priority_thread_stream(self):
while not self.stop_event.is_set():
text = None
try:
tts_stream_queue_msg = self.tts_queue_stream.get()
try:
text = tts_stream_queue_msg["text"]
chunk_queque = tts_stream_queue_msg["chunk_queque"]
text_index = tts_stream_queue_msg["text_index"]
except TimeoutError:
self.logger.error("TTS 任务超时")
continue
except Exception as e:
self.logger.error(f"TTS 任务出错: {e}")
continue
if not self.client_abort:
# 如果没有中途打断就发送语音
self.audio_play_queue.put((chunk_queque, text, text_index))
except Exception as e:
self.clearSpeakStatus()
asyncio.run_coroutine_threadsafe(
self.websocket.send(json.dumps({"type": "tts", "state": "stop", "session_id": self.session_id})),
self.loop
)
self.logger.error(f"tts_priority priority_thread: {text}{e}")
def _audio_play_priority_thread(self):
while not self.stop_event.is_set():
text = None
try:
opus_datas, text, text_index = self.audio_play_queue.get()
future = asyncio.run_coroutine_threadsafe(sendAudioMessage(self, opus_datas, text, text_index),
self.loop)
future.result()
if self.tts_stream:
chunk_queque, text, text_index = self.audio_play_queue.get()
future = asyncio.run_coroutine_threadsafe(
sendAudioMessageStream(self, chunk_queque, text, text_index),
self.loop)
future.result()
else:
opus_datas, text, text_index = self.audio_play_queue.get()
future = asyncio.run_coroutine_threadsafe(sendAudioMessage(self, opus_datas, text, text_index),
self.loop)
future.result()
except Exception as e:
self.logger.bind(tag=TAG).error(f"audio_play_priority priority_thread: {text} {e}")
@@ -344,11 +406,18 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).debug(f"TTS 文件生成完毕: {tts_file}")
return tts_file, text, text_index
def speak_and_play_stream(self, text, queue: queue.Queue, text_index=0):
if text is None or len(text) <= 0:
self.logger.bind(tag=TAG).info(f"无需tts转换,query为空,{text}")
return None, text
self.tts.to_tts_stream(text, queue, text_index)
def clearSpeakStatus(self):
self.logger.bind(tag=TAG).debug(f"清除服务端讲话状态")
self.asr_server_receive = True
self.tts_last_text_index = -1
self.tts_first_text_index = -1
self.tts_duration = 0
def recode_first_last_text(self, text, text_index=0):
if self.tts_first_text_index == -1: