diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 42271f91..abe09d02 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -315,7 +315,9 @@ class ConnectionHandler: self.asr = self._asr if self.tts is None: self.tts = self._tts - self.tts.startSession(self) + # 使用事件循环运行异步方法 + asyncio.run_coroutine_threadsafe(self.tts.open_audio_channels(self), self.loop) + """加载记忆""" self._initialize_memory() """加载意图识别""" diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 67087dbe..fbcd6355 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -4,6 +4,7 @@ import queue import os import json import threading +from enum import Enum from core.utils import p3 from core.handle.sendAudioHandle import sendAudioMessage from core.handle.reportHandle import enqueue_tts_report @@ -16,6 +17,14 @@ TAG = __name__ logger = setup_logging() +class TTSImplementationType(Enum): + """TTS实现类型枚举""" + + NON_STREAMING = "non_streaming" # 非流式实现 + SINGLE_STREAMING = "single_streaming" # 单流式实现 + DOUBLE_STREAMING = "double_streaming" # 双流式实现 + + class TTSProviderBase(ABC): def __init__(self, config, delete_audio_file): self.conn = None @@ -27,12 +36,20 @@ class TTSProviderBase(ABC): self.opus_encoder = opus_encoder_utils.OpusEncoderUtils( sample_rate=16000, channels=1, frame_size_ms=60 ) + # 添加实现类型属性,默认为非流式 + self.interface_type = TTSImplementationType.NON_STREAMING @abstractmethod def generate_filename(self): pass def to_tts(self, text): + """如果是流式实现,一般没有文件生成,我们返回枚举值""" + if self.interface_type != TTSImplementationType.SINGLE_STREAMING: + asyncio.run(self.text_to_speak(text, None)) + return self.interface_type.value + + """以下是非流式实现,会返回文件""" tmp_file = self.generate_filename() try: max_repeat_time = 5 @@ -75,7 +92,7 @@ class TTSProviderBase(ABC): """音频文件转换为Opus编码""" return audio_to_data(audio_file_path, is_opus=True) - def startSession(self, conn): + async def open_audio_channels(self, conn): self.conn = conn self.tts_timeout = conn.config.get("tts_timeout", 10) # tts 消化线程 @@ -110,10 +127,18 @@ class TTSProviderBase(ABC): try: logger.bind(tag=TAG).debug("正在处理TTS任务...") tts_file, text, _ = future.result(timeout=self.tts_timeout) + + # 如果tts_file返回流式标识,则不继续处理 if tts_file is None: logger.bind(tag=TAG).error( f"TTS出错: file is empty: {text_index}: {text}" ) + elif ( + tts_file == TTSImplementationType.SINGLE_STREAMING.value + or tts_file == TTSImplementationType.DOUBLE_STREAMING.value + ): + logger.bind(tag=TAG).debug(f"TTS生成:流式标识: {tts_file}") + continue else: logger.bind(tag=TAG).debug(f"TTS生成:文件路径: {tts_file}") if os.path.exists(tts_file): @@ -134,7 +159,21 @@ class TTSProviderBase(ABC): except TimeoutError: logger.bind(tag=TAG).error("TTS超时") except Exception as e: - logger.bind(tag=TAG).error(f"TTS出错: {e}") + import traceback + + error_info = { + "error_type": type(e).__name__, + "error_message": str(e), + "stack_trace": traceback.format_exc(), + "text_index": text_index, + "text": text, + "tts_file": tts_file, + "audio_format": getattr(self.conn, "audio_format", None), + "interface_type": self.interface_type.value, + } + logger.bind(tag=TAG).error( + f"TTS处理出错: {json.dumps(error_info, ensure_ascii=False)}" + ) if not self.conn.client_abort: # 如果没有中途打断就发送语音 self.audio_play_queue.put((audio_datas, text, text_index)) @@ -154,7 +193,7 @@ class TTSProviderBase(ABC): { "type": "tts", "state": "stop", - "session_id": self.session_id, + "session_id": self.conn.session_id, } ) ), @@ -185,3 +224,14 @@ class TTSProviderBase(ABC): def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False): opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end) return opus_datas + + async def start_session(self, session_id): + pass + + async def finish_session(self, session_id): + pass + + async def close(self): + """资源清理方法""" + if hasattr(self, "ws") and self.ws: + await self.ws.close() diff --git a/main/xiaozhi-server/core/providers/tts/huoshan.py b/main/xiaozhi-server/core/providers/tts/huoshan.py index 31138104..54978bea 100644 --- a/main/xiaozhi-server/core/providers/tts/huoshan.py +++ b/main/xiaozhi-server/core/providers/tts/huoshan.py @@ -1,5 +1,4 @@ import asyncio -import io import os import threading import traceback @@ -10,7 +9,7 @@ from datetime import datetime import websockets from config.logger import setup_logging -from core.providers.tts.base import TTSProviderBase +from core.providers.tts.base import TTSProviderBase, TTSImplementationType TAG = __name__ logger = setup_logging() @@ -151,34 +150,24 @@ class TTSProvider(TTSProviderBase): self.enable_two_way = True self.start_connection_flag = False self.tts_text = "" + self.interface_type = TTSImplementationType.DOUBLE_STREAMING - def startSession(self, conn): - self.conn = conn - self.tts_timeout = conn.config.get("tts_timeout", 10) - # tts 消化线程 - self.tts_priority_thread = threading.Thread( - target=self._tts_priority_thread, daemon=True - ) - self.tts_priority_thread.start() - - # 音频播放 消化线程 - self.audio_play_priority_thread = threading.Thread( - target=self._audio_play_priority_thread, daemon=True - ) - self.audio_play_priority_thread.start() + async def open_audio_channels(self, conn): + await super().open_audio_channels(conn) ws_header = { "X-Api-App-Key": self.appId, "X-Api-Access-Key": self.access_token, "X-Api-Resource-Id": self.resource_id, "X-Api-Connect-Id": uuid.uuid4(), } - self.ws = websockets.connect( + self.ws = await websockets.connect( self.ws_url, additional_headers=ws_header, max_size=1000000000 ) tts_priority = threading.Thread( target=self._start_monitor_tts_response_thread(), daemon=True ) tts_priority.start() + await self.start_session(conn.session_id) def generate_filename(self, extension=".wav"): return os.path.join( @@ -381,7 +370,7 @@ class TTSProvider(TTSProviderBase): async def _start_monitor_tts_response(self): chunk_total = b"" - while not self.stop_event.is_set(): + while not self.conn.stop_event.is_set(): try: msg = await self.ws.recv() # 确保 `recv()` 运行在同一个 event loop res = self.parser_response(msg)