From 98174bcc162128a9e9258d077f47a3e122812a8a Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Sun, 8 Jun 2025 00:06:40 +0800 Subject: [PATCH] =?UTF-8?q?update:=E3=80=90=E6=B5=81=E5=BC=8Ftts=E3=80=91?= =?UTF-8?q?=E5=A2=9E=E5=8A=A0=E3=80=90=E9=9D=9E=E6=B5=81=E5=BC=8F=E3=80=91?= =?UTF-8?q?=E6=96=B9=E6=B3=95=EF=BC=8C=E7=94=A8=E4=BA=8E=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E5=8F=8A=E7=94=9F=E6=88=90=E6=96=87=E4=BB=B6=E7=9A=84=E5=9C=BA?= =?UTF-8?q?=E6=99=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../xiaozhi-server/core/handle/helloHandle.py | 17 +-- .../providers/tts/huoshan_double_stream.py | 120 +++++++++++++++++- .../core/providers/tts/linkerai.py | 78 +++++++++++- 3 files changed, 194 insertions(+), 21 deletions(-) diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py index b19dc054..bb15382d 100644 --- a/main/xiaozhi-server/core/handle/helloHandle.py +++ b/main/xiaozhi-server/core/handle/helloHandle.py @@ -1,12 +1,11 @@ -import os -import shutil import time import json import random import asyncio +from core.utils.util import audio_to_data from core.handle.sendAudioHandle import send_stt_message from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes -from core.providers.tts.dto.dto import ContentType, InterfaceType +from core.providers.tts.dto.dto import ContentType, SentenceType from core.handle.mcpHandle import ( MCPClient, send_mcp_initialize_message, @@ -59,9 +58,6 @@ async def checkWakeupWords(conn, text): if not enable_wakeup_words_response_cache or not conn.tts: return False - if conn.tts.interface_type != InterfaceType.NON_STREAM: - return False - _, filtered_text = remove_punctuation_and_length(text) if filtered_text not in conn.config.get("wakeup_words"): return False @@ -77,12 +73,9 @@ async def checkWakeupWords(conn, text): # 播放唤醒词回复 conn.client_abort = False - conn.tts.tts_one_sentence( - conn, - ContentType.FILE, - content_file=response["file_path"], - content_detail=response["text"], - ) + opus_packets, _ = audio_to_data(response["file_path"]) + conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, response["text"])) + conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None)) # 检查是否需要更新唤醒词回复 if time.time() - response["time"] > WAKEUP_CONFIG["refresh_time"]: diff --git a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py index 10a37b51..33a1c8ec 100644 --- a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py +++ b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py @@ -301,7 +301,7 @@ class TTSProvider(TTSProviderBase): payload = self.get_payload_bytes( event=EVENT_StartSession, speaker=self.voice ) - await self.send_event(header, optional, payload) + await self.send_event(self.ws, header, optional, payload) logger.bind(tag=TAG).info("会话启动请求已发送") except Exception as e: logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}") @@ -334,7 +334,7 @@ class TTSProvider(TTSProviderBase): event=EVENT_FinishSession, sessionId=session_id ).as_bytes() payload = str.encode("{}") - await self.send_event(header, optional, payload) + await self.send_event(self.ws, header, optional, payload) logger.bind(tag=TAG).info("会话结束请求已发送") # 等待监听任务完成 @@ -451,7 +451,11 @@ class TTSProvider(TTSProviderBase): self.ws = None async def send_event( - self, header: bytes, optional: bytes | None = None, payload: bytes = None + self, + ws: websockets.WebSocketClientProtocol, + header: bytes, + optional: bytes | None = None, + payload: bytes = None, ): try: full_client_request = bytearray(header) @@ -461,7 +465,7 @@ class TTSProvider(TTSProviderBase): payload_size = len(payload).to_bytes(4, "big", signed=True) full_client_request.extend(payload_size) full_client_request.extend(payload) - await self.ws.send(full_client_request) + await ws.send(full_client_request) except websockets.ConnectionClosed: logger.bind(tag=TAG).error(f"ConnectionClosed") raise @@ -476,7 +480,7 @@ class TTSProvider(TTSProviderBase): payload = self.get_payload_bytes( event=EVENT_TaskRequest, text=text, speaker=speaker ) - return await self.send_event(header, optional, payload) + return await self.send_event(self.ws, header, optional, payload) # 读取 res 数组某段 字符串内容 def read_res_content(self, res: bytes, offset: int): @@ -554,7 +558,7 @@ class TTSProvider(TTSProviderBase): ).as_bytes() optional = Optional(event=EVENT_Start_Connection).as_bytes() payload = str.encode("{}") - return await self.send_event(header, optional, payload) + return await self.send_event(self.ws, header, optional, payload) def print_response(self, res, tag_msg: str): logger.bind(tag=TAG).debug(f"===>{tag_msg} header:{res.header.__dict__}") @@ -590,3 +594,107 @@ class TTSProvider(TTSProviderBase): 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 + + def to_tts(self, text: str) -> list: + """非流式生成音频数据,用于生成音频及测试场景 + + Args: + text: 要转换的文本 + + Returns: + list: 音频数据列表 + """ + try: + # 创建事件循环 + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + # 生成会话ID + session_id = uuid.uuid4().__str__().replace("-", "") + + # 存储音频数据 + audio_data = [] + + async def _generate_audio(): + # 创建新的WebSocket连接 + 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(), + } + ws = await websockets.connect( + self.ws_url, additional_headers=ws_header, max_size=1000000000 + ) + + try: + # 启动会话 + header = Header( + message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON, + ).as_bytes() + optional = Optional( + event=EVENT_StartSession, sessionId=session_id + ).as_bytes() + payload = self.get_payload_bytes( + event=EVENT_StartSession, speaker=self.voice + ) + await self.send_event(ws, header, optional, payload) + + # 发送文本 + header = Header( + message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON, + ).as_bytes() + optional = Optional( + event=EVENT_TaskRequest, sessionId=session_id + ).as_bytes() + payload = self.get_payload_bytes( + event=EVENT_TaskRequest, text=text, speaker=self.voice + ) + await self.send_event(ws, header, optional, payload) + + # 发送结束会话请求 + header = Header( + message_type=FULL_CLIENT_REQUEST, + message_type_specific_flags=MsgTypeFlagWithEvent, + serial_method=JSON, + ).as_bytes() + optional = Optional( + event=EVENT_FinishSession, sessionId=session_id + ).as_bytes() + payload = str.encode("{}") + await self.send_event(ws, header, optional, payload) + + # 接收音频数据 + while True: + msg = await ws.recv() + res = self.parser_response(msg) + + if ( + res.optional.event == EVENT_TTSResponse + and res.header.message_type == AUDIO_ONLY_RESPONSE + ): + opus_datas = self.wav_to_opus_data_audio_raw(res.payload) + audio_data.extend(opus_datas) + elif res.optional.event == EVENT_SessionFinished: + break + + finally: + # 清理资源 + try: + await ws.close() + except: + pass + + # 运行异步任务 + loop.run_until_complete(_generate_audio()) + loop.close() + + return audio_data + + except Exception as e: + logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}") + return [] diff --git a/main/xiaozhi-server/core/providers/tts/linkerai.py b/main/xiaozhi-server/core/providers/tts/linkerai.py index dd8d0e6b..649ed51d 100644 --- a/main/xiaozhi-server/core/providers/tts/linkerai.py +++ b/main/xiaozhi-server/core/providers/tts/linkerai.py @@ -2,6 +2,8 @@ import queue import asyncio import traceback import aiohttp +import requests +import time from config.logger import setup_logging from core.utils.tts import MarkdownCleaner from core.providers.tts.base import TTSProviderBase @@ -55,7 +57,7 @@ class TTSProvider(TTSProviderBase): self.tts_text_buff.append(message.content_detail) segment_text = self._get_segment_text() if segment_text: - self.to_tts(segment_text) + self.to_tts_single_stream(segment_text) elif ContentType.FILE == message.content_type: logger.bind(tag=TAG).info( @@ -87,12 +89,12 @@ class TTSProvider(TTSProviderBase): if remaining_text: segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text) if segment_text: - self.to_tts(segment_text, is_last) + self.to_tts_single_stream(segment_text, is_last) self.processed_chars += len(full_text) else: self._process_before_stop_play_files() - def to_tts(self, text, is_last=False): + def to_tts_single_stream(self, text, is_last=False): try: max_repeat_time = 5 text = MarkdownCleaner.clean_markdown(text) @@ -225,3 +227,73 @@ class TTSProvider(TTSProviderBase): except Exception as e: logger.error(f"TTS请求异常: {e}") self.tts_audio_queue.put((SentenceType.LAST, [], None)) + + def to_tts(self, text: str) -> list: + """非流式TTS处理,用于测试及保存音频文件的场景 + + Args: + text: 要转换的文本 + + Returns: + list: 返回opus编码后的音频数据列表 + """ + start_time = time.time() + text = MarkdownCleaner.clean_markdown(text) + + params = { + "tts_text": text, + "spk_id": self.voice, + "frame_duration": 60, + "stream": False, + "target_sr": 16000, + "audio_format": self.audio_format, + "instruct_text": "请生成一段自然流畅的语音", + } + headers = { + "Authorization": f"Bearer {self.access_token}", + "Content-Type": "application/json", + } + + try: + with requests.get( + self.api_url, params=params, headers=headers, timeout=5 + ) as response: + if response.status_code != 200: + logger.error( + f"TTS请求失败: {response.status_code}, {response.text}" + ) + return [] + + logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒") + + # 使用opus编码器处理PCM数据 + opus_datas = [] + pcm_data = response.content + + # 计算每帧的字节数 + frame_bytes = int( + self.opus_encoder.sample_rate + * self.opus_encoder.channels + * self.opus_encoder.frame_size_ms + / 1000 + * 2 + ) + + # 分帧处理PCM数据 + for i in range(0, len(pcm_data), frame_bytes): + frame = pcm_data[i : i + frame_bytes] + if len(frame) < frame_bytes: + # 最后一帧可能不足,用0填充 + frame = frame + b"\x00" * (frame_bytes - len(frame)) + + opus = self.opus_encoder.encode_pcm_to_opus( + frame, end_of_stream=(i + frame_bytes >= len(pcm_data)) + ) + if opus: + opus_datas.extend(opus) + + return opus_datas + + except Exception as e: + logger.error(f"TTS请求异常: {e}") + return []