mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-28 19:23:52 +08:00
update:【流式tts】增加【非流式】方法,用于测试及生成文件的场景
This commit is contained in:
@@ -1,12 +1,11 @@
|
|||||||
import os
|
|
||||||
import shutil
|
|
||||||
import time
|
import time
|
||||||
import json
|
import json
|
||||||
import random
|
import random
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from core.utils.util import audio_to_data
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
|
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 (
|
from core.handle.mcpHandle import (
|
||||||
MCPClient,
|
MCPClient,
|
||||||
send_mcp_initialize_message,
|
send_mcp_initialize_message,
|
||||||
@@ -59,9 +58,6 @@ async def checkWakeupWords(conn, text):
|
|||||||
if not enable_wakeup_words_response_cache or not conn.tts:
|
if not enable_wakeup_words_response_cache or not conn.tts:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
if conn.tts.interface_type != InterfaceType.NON_STREAM:
|
|
||||||
return False
|
|
||||||
|
|
||||||
_, filtered_text = remove_punctuation_and_length(text)
|
_, filtered_text = remove_punctuation_and_length(text)
|
||||||
if filtered_text not in conn.config.get("wakeup_words"):
|
if filtered_text not in conn.config.get("wakeup_words"):
|
||||||
return False
|
return False
|
||||||
@@ -77,12 +73,9 @@ async def checkWakeupWords(conn, text):
|
|||||||
|
|
||||||
# 播放唤醒词回复
|
# 播放唤醒词回复
|
||||||
conn.client_abort = False
|
conn.client_abort = False
|
||||||
conn.tts.tts_one_sentence(
|
opus_packets, _ = audio_to_data(response["file_path"])
|
||||||
conn,
|
conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, response["text"]))
|
||||||
ContentType.FILE,
|
conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
content_file=response["file_path"],
|
|
||||||
content_detail=response["text"],
|
|
||||||
)
|
|
||||||
|
|
||||||
# 检查是否需要更新唤醒词回复
|
# 检查是否需要更新唤醒词回复
|
||||||
if time.time() - response["time"] > WAKEUP_CONFIG["refresh_time"]:
|
if time.time() - response["time"] > WAKEUP_CONFIG["refresh_time"]:
|
||||||
|
|||||||
@@ -301,7 +301,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
payload = self.get_payload_bytes(
|
payload = self.get_payload_bytes(
|
||||||
event=EVENT_StartSession, speaker=self.voice
|
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("会话启动请求已发送")
|
logger.bind(tag=TAG).info("会话启动请求已发送")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
|
||||||
@@ -334,7 +334,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
event=EVENT_FinishSession, sessionId=session_id
|
event=EVENT_FinishSession, sessionId=session_id
|
||||||
).as_bytes()
|
).as_bytes()
|
||||||
payload = str.encode("{}")
|
payload = str.encode("{}")
|
||||||
await self.send_event(header, optional, payload)
|
await self.send_event(self.ws, header, optional, payload)
|
||||||
logger.bind(tag=TAG).info("会话结束请求已发送")
|
logger.bind(tag=TAG).info("会话结束请求已发送")
|
||||||
|
|
||||||
# 等待监听任务完成
|
# 等待监听任务完成
|
||||||
@@ -451,7 +451,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.ws = None
|
self.ws = None
|
||||||
|
|
||||||
async def send_event(
|
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:
|
try:
|
||||||
full_client_request = bytearray(header)
|
full_client_request = bytearray(header)
|
||||||
@@ -461,7 +465,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
payload_size = len(payload).to_bytes(4, "big", signed=True)
|
payload_size = len(payload).to_bytes(4, "big", signed=True)
|
||||||
full_client_request.extend(payload_size)
|
full_client_request.extend(payload_size)
|
||||||
full_client_request.extend(payload)
|
full_client_request.extend(payload)
|
||||||
await self.ws.send(full_client_request)
|
await ws.send(full_client_request)
|
||||||
except websockets.ConnectionClosed:
|
except websockets.ConnectionClosed:
|
||||||
logger.bind(tag=TAG).error(f"ConnectionClosed")
|
logger.bind(tag=TAG).error(f"ConnectionClosed")
|
||||||
raise
|
raise
|
||||||
@@ -476,7 +480,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
payload = self.get_payload_bytes(
|
payload = self.get_payload_bytes(
|
||||||
event=EVENT_TaskRequest, text=text, speaker=speaker
|
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 数组某段 字符串内容
|
# 读取 res 数组某段 字符串内容
|
||||||
def read_res_content(self, res: bytes, offset: int):
|
def read_res_content(self, res: bytes, offset: int):
|
||||||
@@ -554,7 +558,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
).as_bytes()
|
).as_bytes()
|
||||||
optional = Optional(event=EVENT_Start_Connection).as_bytes()
|
optional = Optional(event=EVENT_Start_Connection).as_bytes()
|
||||||
payload = str.encode("{}")
|
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):
|
def print_response(self, res, tag_msg: str):
|
||||||
logger.bind(tag=TAG).debug(f"===>{tag_msg} header:{res.header.__dict__}")
|
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):
|
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)
|
opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end)
|
||||||
return opus_datas
|
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 []
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ import queue
|
|||||||
import asyncio
|
import asyncio
|
||||||
import traceback
|
import traceback
|
||||||
import aiohttp
|
import aiohttp
|
||||||
|
import requests
|
||||||
|
import time
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
@@ -55,7 +57,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
segment_text = self._get_segment_text()
|
segment_text = self._get_segment_text()
|
||||||
if segment_text:
|
if segment_text:
|
||||||
self.to_tts(segment_text)
|
self.to_tts_single_stream(segment_text)
|
||||||
|
|
||||||
elif ContentType.FILE == message.content_type:
|
elif ContentType.FILE == message.content_type:
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
@@ -87,12 +89,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
if remaining_text:
|
if remaining_text:
|
||||||
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
||||||
if segment_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)
|
self.processed_chars += len(full_text)
|
||||||
else:
|
else:
|
||||||
self._process_before_stop_play_files()
|
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:
|
try:
|
||||||
max_repeat_time = 5
|
max_repeat_time = 5
|
||||||
text = MarkdownCleaner.clean_markdown(text)
|
text = MarkdownCleaner.clean_markdown(text)
|
||||||
@@ -225,3 +227,73 @@ class TTSProvider(TTSProviderBase):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"TTS请求异常: {e}")
|
logger.error(f"TTS请求异常: {e}")
|
||||||
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
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 []
|
||||||
|
|||||||
Reference in New Issue
Block a user