mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-21 22:53:56 +08:00
add:火山双向tts语音流式输入输出
This commit is contained in:
@@ -156,3 +156,6 @@ main/xiaozhi-server/models/SenseVoiceSmall/model.pt
|
||||
main/xiaozhi-server/models/sherpa-onnx*
|
||||
/main/xiaozhi-server/audio_ref/
|
||||
/audio_ref/
|
||||
/asr-models/iic/SenseVoiceSmall/
|
||||
/main/xiaozhi-server/asr-models/iic/SenseVoiceSmall/
|
||||
/models/SenseVoiceSmall/model.pt
|
||||
|
||||
@@ -256,6 +256,15 @@ TTS:
|
||||
appid: 你的火山引擎语音合成服务appid
|
||||
access_token: 你的火山引擎语音合成服务access_token
|
||||
cluster: volcano_tts
|
||||
#火山tts,支持双向流式tts
|
||||
HuoshanTTS:
|
||||
type: huoshan
|
||||
# 如果是机智云 wss://bytedance.gizwitsapi.com/api/v3/tts/bidirection
|
||||
# 机智云不需要天填 appid
|
||||
ws_url: wss://openspeech.bytedance.com/api/v3/tts/bidirection
|
||||
appid: 你的火山引擎语音合成服务appid
|
||||
access_token: 你的火山引擎语音合成服务access_token
|
||||
speaker: zh_female_meilinvyou_moon_bigtts
|
||||
CosyVoiceSiliconflow:
|
||||
type: siliconflow
|
||||
# 硅基流动TTS
|
||||
|
||||
@@ -11,11 +11,12 @@ import websockets
|
||||
from typing import Dict, Any
|
||||
import plugins_func.loadplugins
|
||||
from config.logger import setup_logging
|
||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType
|
||||
from core.utils.dialogue import Message, Dialogue
|
||||
from core.handle.textHandle import handleTextMessage
|
||||
from core.utils.util import get_string_no_punctuation_or_emoji, extract_json_from_string, get_ip_info
|
||||
from concurrent.futures import ThreadPoolExecutor, TimeoutError
|
||||
from core.handle.sendAudioHandle import sendAudioMessage, sendAudioMessageStream
|
||||
from core.handle.sendAudioHandle import sendAudioMessage
|
||||
from core.handle.receiveAudioHandle import handleAudioMessage
|
||||
from core.handle.functionHandler import FunctionHandler
|
||||
from plugins_func.register import Action
|
||||
@@ -58,6 +59,7 @@ class ConnectionHandler:
|
||||
self.audio_play_queue = queue.Queue()
|
||||
max_workers = self.config.get("TTS_SET", {}).get("MAX_WORKERS", 10)
|
||||
self.executor = ThreadPoolExecutor(max_workers=max_workers)
|
||||
self.start_tts_request_flag = False
|
||||
|
||||
# 依赖的组件
|
||||
self.vad = _vad
|
||||
@@ -161,10 +163,6 @@ class ConnectionHandler:
|
||||
tts_priority = threading.Thread(target=self._tts_priority_thread, daemon=True)
|
||||
tts_priority.start()
|
||||
|
||||
# 音频播放 消化线程
|
||||
audio_play_priority = threading.Thread(target=self._audio_play_priority_thread, daemon=True)
|
||||
audio_play_priority.start()
|
||||
|
||||
try:
|
||||
async for message in self.websocket:
|
||||
await self._route_message(message)
|
||||
@@ -196,9 +194,9 @@ class ConnectionHandler:
|
||||
if self.private_config:
|
||||
self.prompt = self.private_config.private_config.get("prompt", self.prompt)
|
||||
|
||||
self.client_ip_info = get_ip_info(self.client_ip)
|
||||
self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}")
|
||||
self.prompt = self.prompt + f"\n我在:{self.client_ip_info}"
|
||||
# self.client_ip_info = get_ip_info(self.client_ip)
|
||||
# self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}")
|
||||
# self.prompt = self.prompt + f"\n我在:{self.client_ip_info}"
|
||||
self.dialogue.put(Message(role="system", content=self.prompt))
|
||||
|
||||
self.func_handler = FunctionHandler(self.config)
|
||||
@@ -258,6 +256,8 @@ class ConnectionHandler:
|
||||
|
||||
self.llm_finish_task = False
|
||||
text_index = 0
|
||||
uuid_str = str(uuid.uuid4())
|
||||
msg_type = None
|
||||
for content in llm_responses:
|
||||
response_message.append(content)
|
||||
if self.client_abort:
|
||||
@@ -265,62 +265,14 @@ class ConnectionHandler:
|
||||
|
||||
end_time = time.time()
|
||||
self.logger.bind(tag=TAG).debug(f"大模型返回时间: {end_time - start_time} 秒, 生成token={content}")
|
||||
|
||||
# 合并当前全部文本并处理未分割部分
|
||||
full_text = "".join(response_message)
|
||||
current_text = full_text[processed_chars:] # 从未处理的位置开始
|
||||
|
||||
# 查找最后一个有效标点
|
||||
punctuations = ("。", "?", "!", ";", ":", ".", "?", "!", ";", ":", " ")
|
||||
last_punct_pos = -1
|
||||
for punct in punctuations:
|
||||
pos = current_text.rfind(punct)
|
||||
if pos > last_punct_pos:
|
||||
last_punct_pos = pos
|
||||
|
||||
# 找到分割点则处理
|
||||
if last_punct_pos != -1:
|
||||
segment_text_raw = current_text[:last_punct_pos + 1]
|
||||
segment_text = get_string_no_punctuation_or_emoji(segment_text_raw)
|
||||
if segment_text:
|
||||
# 强制设置空字符,测试TTS出错返回语音的健壮性
|
||||
# if text_index % 2 == 0:
|
||||
# segment_text = " "
|
||||
text_index += 1
|
||||
self.recode_first_last_text(segment_text, text_index)
|
||||
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) # 更新已处理字符位置
|
||||
|
||||
# 处理最后剩余的文本
|
||||
full_text = "".join(response_message)
|
||||
remaining_text = full_text[processed_chars:]
|
||||
if remaining_text:
|
||||
segment_text = get_string_no_punctuation_or_emoji(remaining_text)
|
||||
if segment_text:
|
||||
text_index += 1
|
||||
self.recode_first_last_text(segment_text, text_index)
|
||||
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)
|
||||
|
||||
if text_index == 0:
|
||||
msg_type = MsgType.START_TTS_REQUEST
|
||||
else:
|
||||
msg_type = MsgType.TTS_TEXT_REQUEST
|
||||
self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=msg_type, content=content))
|
||||
text_index += 1
|
||||
msg_type = MsgType.STOP_TTS_REQUEST
|
||||
self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=msg_type, content=""))
|
||||
self.llm_finish_task = True
|
||||
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
|
||||
self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False))
|
||||
@@ -372,6 +324,8 @@ class ConnectionHandler:
|
||||
function_id = None
|
||||
function_arguments = ""
|
||||
content_arguments = ""
|
||||
uuid_str = str(uuid.uuid4()).replace("-", "")
|
||||
msg_type = None
|
||||
for response in llm_responses:
|
||||
content, tools_call = response
|
||||
if content is not None and len(content) > 0:
|
||||
@@ -398,39 +352,16 @@ class ConnectionHandler:
|
||||
|
||||
end_time = time.time()
|
||||
self.logger.bind(tag=TAG).debug(f"大模型返回时间: {end_time - start_time} 秒, 生成token={content}")
|
||||
|
||||
# 处理文本分段和TTS逻辑
|
||||
# 合并当前全部文本并处理未分割部分
|
||||
full_text = "".join(response_message)
|
||||
current_text = full_text[processed_chars:] # 从未处理的位置开始
|
||||
|
||||
# 查找最后一个有效标点
|
||||
punctuations = ("。", "?", "!", ";", ":", ".", "?", "!", ";", ":", " ")
|
||||
last_punct_pos = -1
|
||||
for punct in punctuations:
|
||||
pos = current_text.rfind(punct)
|
||||
if pos > last_punct_pos:
|
||||
last_punct_pos = pos
|
||||
|
||||
# 找到分割点则处理
|
||||
if last_punct_pos != -1:
|
||||
segment_text_raw = current_text[:last_punct_pos + 1]
|
||||
segment_text = get_string_no_punctuation_or_emoji(segment_text_raw)
|
||||
if segment_text:
|
||||
text_index += 1
|
||||
self.recode_first_last_text(segment_text, text_index)
|
||||
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) # 更新已处理字符位置
|
||||
if text_index == 0:
|
||||
self.tts.tts_text_queue.put(
|
||||
TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.START_TTS_REQUEST, content=''))
|
||||
self.start_tts_request_flag = True
|
||||
self.tts.tts_text_queue.put(
|
||||
TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.TTS_TEXT_REQUEST, content=content))
|
||||
text_index += 1
|
||||
if self.start_tts_request_flag:
|
||||
self.start_tts_request_flag = False
|
||||
self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.STOP_TTS_REQUEST, content=''))
|
||||
|
||||
# 处理function call
|
||||
if tool_call_flag:
|
||||
@@ -464,26 +395,6 @@ class ConnectionHandler:
|
||||
result = self.func_handler.handle_llm_function_call(self, function_call_data)
|
||||
self._handle_function_result(result, function_call_data, text_index + 1)
|
||||
|
||||
# 处理最后剩余的文本
|
||||
full_text = "".join(response_message)
|
||||
remaining_text = full_text[processed_chars:]
|
||||
if remaining_text:
|
||||
segment_text = get_string_no_punctuation_or_emoji(remaining_text)
|
||||
if segment_text:
|
||||
text_index += 1
|
||||
self.recode_first_last_text(segment_text, text_index)
|
||||
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)
|
||||
|
||||
# 存储对话内容
|
||||
if len(response_message) > 0:
|
||||
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
|
||||
@@ -497,20 +408,9 @@ class ConnectionHandler:
|
||||
if result.action == Action.RESPONSE: # 直接回复前端
|
||||
text = result.response
|
||||
self.recode_first_last_text(text, text_index)
|
||||
if self.tts_stream:
|
||||
stream_queue = queue.Queue()
|
||||
self.executor.submit(self.speak_and_play_stream, text, stream_queue, text_index)
|
||||
self.tts_queue_stream.put({
|
||||
"text": text,
|
||||
"chunk_queque": stream_queue,
|
||||
"text_index": text_index
|
||||
})
|
||||
else:
|
||||
future = self.executor.submit(self.speak_and_play, text, text_index)
|
||||
self.tts_queue.put(future)
|
||||
asyncio.run_coroutine_threadsafe(self.tts.tts_one_sentence(text), loop=self.loop)
|
||||
self.dialogue.put(Message(role="assistant", content=text))
|
||||
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
|
||||
|
||||
text = result.result
|
||||
if text is not None and len(text) > 0:
|
||||
function_id = function_call_data["id"]
|
||||
@@ -528,14 +428,12 @@ class ConnectionHandler:
|
||||
elif result.action == Action.NOTFOUND:
|
||||
text = result.result
|
||||
self.recode_first_last_text(text, text_index)
|
||||
future = self.executor.submit(self.speak_and_play, text, text_index)
|
||||
self.tts_queue.put(future)
|
||||
asyncio.run_coroutine_threadsafe(self.tts.tts_one_sentence(text), loop=self.loop)
|
||||
self.dialogue.put(Message(role="assistant", content=text))
|
||||
else:
|
||||
text = result.result
|
||||
self.recode_first_last_text(text, text_index)
|
||||
future = self.executor.submit(self.speak_and_play, text, text_index)
|
||||
self.tts_queue.put(future)
|
||||
asyncio.run_coroutine_threadsafe(self.tts.tts_one_sentence(text), loop=self.loop)
|
||||
self.dialogue.put(Message(role="assistant", content=text))
|
||||
|
||||
def _tts_priority_thread(self):
|
||||
@@ -615,22 +513,10 @@ class ConnectionHandler:
|
||||
while not self.stop_event.is_set():
|
||||
text = None
|
||||
try:
|
||||
if self.tts_stream:
|
||||
data, text, text_index = self.audio_play_queue.get()
|
||||
if isinstance(data, list):
|
||||
future = asyncio.run_coroutine_threadsafe(sendAudioMessage(self, data, text, text_index),
|
||||
self.loop)
|
||||
future.result()
|
||||
else:
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
sendAudioMessageStream(self, data, 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()
|
||||
ttsMessageDTO = self.tts.tts_audio_queue.get()
|
||||
future = asyncio.run_coroutine_threadsafe(sendAudioMessage(self, ttsMessageDTO),
|
||||
self.loop)
|
||||
future.result()
|
||||
except Exception as e:
|
||||
self.logger.bind(tag=TAG).error(f"audio_play_priority priority_thread: {text} {e}")
|
||||
|
||||
@@ -668,6 +554,15 @@ class ConnectionHandler:
|
||||
self.tts_first_text_index = text_index
|
||||
self.tts_last_text_index = text_index
|
||||
|
||||
async def init_and_reset_tts(self):
|
||||
self.stop_event.set()
|
||||
# 释放之前的tts语音监听:重置监听队列
|
||||
await self.tts.reset()
|
||||
# 音频播放 消化线程
|
||||
self.stop_event.clear()
|
||||
audio_play_priority = threading.Thread(target=self._audio_play_priority_thread, daemon=True)
|
||||
audio_play_priority.start()
|
||||
|
||||
async def close(self):
|
||||
"""资源清理方法"""
|
||||
|
||||
@@ -676,6 +571,7 @@ class ConnectionHandler:
|
||||
self.executor.shutdown(wait=False)
|
||||
if self.websocket:
|
||||
await self.websocket.close()
|
||||
await self.tts.close()
|
||||
self.logger.bind(tag=TAG).info("连接资源已释放")
|
||||
|
||||
def reset_vad_states(self):
|
||||
|
||||
@@ -4,94 +4,31 @@ from config.logger import setup_logging
|
||||
import json
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType, MsgType
|
||||
from core.utils.util import remove_punctuation_and_length, get_string_no_punctuation_or_emoji
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
|
||||
|
||||
async def sendAudioMessageStream(conn, audios_queue, text, text_index=0, llm_finish_task=False):
|
||||
async def sendAudioMessage(conn, ttsMessageDTO: TTSMessageDTO):
|
||||
u_id = None
|
||||
# 发送句子开始消息
|
||||
if text_index == conn.tts_first_text_index:
|
||||
logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
||||
await send_tts_message(conn, "sentence_start", text)
|
||||
|
||||
# 初始化流控参数
|
||||
frame_duration = 60 # 毫秒
|
||||
start_time = time.time() # 使用高精度计时器
|
||||
# 初始化流控参数
|
||||
frame_duration = 60 # 毫秒
|
||||
start_time_chunk = time.perf_counter() # 使用高精度计时器
|
||||
play_position = 0 # 已播放的时长(毫秒)
|
||||
while True:
|
||||
try:
|
||||
start_get_queue = time.time()
|
||||
# 尝试获取数据,如果没有数据,则等待一小段时间再试
|
||||
audio_data_chunke = None
|
||||
try:
|
||||
audio_data_chunke = audios_queue.get(timeout=5) # 设置超时为1秒
|
||||
except Exception as e:
|
||||
# 如果超时,继续等待
|
||||
logger.bind(tag=TAG).error(f"获取队列超时~{e}")
|
||||
|
||||
audio_opus_datas = audio_data_chunke.get('data') if audio_data_chunke else None
|
||||
duration = audio_data_chunke.get('duration') if audio_data_chunke else 0
|
||||
|
||||
if audio_data_chunke:
|
||||
start_time = time.time()
|
||||
# 检查是否超过 5 秒没有数据
|
||||
if time.time() - start_time > 15:
|
||||
logger.bind(tag=TAG).error("超过15秒没有数据,退出。")
|
||||
break
|
||||
|
||||
if audio_data_chunke and audio_data_chunke.get("end", True):
|
||||
break
|
||||
|
||||
if audio_opus_datas:
|
||||
for opus_packet in audio_opus_datas:
|
||||
if conn.client_abort:
|
||||
return
|
||||
logger.bind(tag=TAG).info(f'发送数据长度:{len(opus_packet)}')
|
||||
await conn.websocket.send(opus_packet)
|
||||
play_position += frame_duration # 更新播放位置
|
||||
start_time = time.time() # 更新获取数据的时间
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"发生错误: {e}")
|
||||
traceback.print_exc() # 打印错误堆栈
|
||||
await send_tts_message(conn, "sentence_end", text)
|
||||
|
||||
print(f'{text_index}-{conn.tts_last_text_index}')
|
||||
expected_time = start_time_chunk + (play_position / 1000)
|
||||
current_time = time.perf_counter()
|
||||
# 等待直到预期时间
|
||||
delay = expected_time - current_time
|
||||
if delay > 0:
|
||||
await asyncio.sleep(delay)
|
||||
# 发送结束消息(如果是最后一个文本)
|
||||
logger.bind(tag=TAG).info(f"{conn.llm_finish_task},{text_index},{conn.tts_last_text_index}")
|
||||
if conn.llm_finish_task and text_index == conn.tts_last_text_index:
|
||||
await send_tts_message(conn, 'stop', None)
|
||||
if conn.close_after_chat or "拜拜" in text or "再见" in text:
|
||||
await conn.close()
|
||||
|
||||
|
||||
|
||||
async def sendAudioMessage(conn, audios, text, text_index=0):
|
||||
# 发送句子开始消息
|
||||
if text_index == conn.tts_first_text_index:
|
||||
logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
||||
await send_tts_message(conn, "sentence_start", text)
|
||||
if SentenceType.SENTENCE_START == ttsMessageDTO.sentence_type:
|
||||
logger.bind(tag=TAG).info(f"发送第一段语音: {ttsMessageDTO.tts_finish_text}")
|
||||
await send_tts_message(conn, "sentence_start", ttsMessageDTO.tts_finish_text)
|
||||
|
||||
# 流控参数优化
|
||||
original_frame_duration = 60 # 原始帧时长(毫秒)
|
||||
adjusted_frame_duration = int(original_frame_duration * 0.8) # 缩短20%
|
||||
total_frames = len(audios) # 获取总帧数
|
||||
total_frames = len(ttsMessageDTO.content) # 获取总帧数
|
||||
compensation = total_frames * (original_frame_duration - adjusted_frame_duration) / 1000 # 补偿时间(秒)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
play_position = 0 # 已播放时长(毫秒)
|
||||
|
||||
for opus_packet in audios:
|
||||
for opus_packet in ttsMessageDTO.content:
|
||||
if conn.client_abort:
|
||||
return
|
||||
|
||||
@@ -110,15 +47,17 @@ async def sendAudioMessage(conn, audios, text, text_index=0):
|
||||
# 补偿因加速损失的时长
|
||||
if compensation > 0:
|
||||
await asyncio.sleep(compensation)
|
||||
|
||||
await send_tts_message(conn, "sentence_end", text)
|
||||
if SentenceType.SENTENCE_END == ttsMessageDTO.sentence_type:
|
||||
logger.bind(tag=TAG).info(f"发送最后一段语音: {ttsMessageDTO.tts_finish_text}")
|
||||
await send_tts_message(conn, "sentence_end", ttsMessageDTO.tts_finish_text)
|
||||
|
||||
# 发送结束消息(如果是最后一个文本)
|
||||
if conn.llm_finish_task and text_index == conn.tts_last_text_index:
|
||||
if conn.llm_finish_task and MsgType.STOP_TTS_RESPONSE == ttsMessageDTO.msg_type:
|
||||
await send_tts_message(conn, 'stop', None)
|
||||
if conn.close_after_chat:
|
||||
await conn.close()
|
||||
|
||||
|
||||
async def send_tts_message(conn, state, text=None):
|
||||
"""发送 TTS 状态消息"""
|
||||
message = {
|
||||
|
||||
@@ -28,6 +28,8 @@ async def handleTextMessage(conn, message):
|
||||
if msg_json["state"] == "start":
|
||||
conn.client_have_voice = True
|
||||
conn.client_voice_stop = False
|
||||
# 打断,开启了行的对话,如果之前有tts存在,销毁掉重新建立tts
|
||||
await conn.init_and_reset_tts()
|
||||
elif msg_json["state"] == "stop":
|
||||
conn.client_have_voice = True
|
||||
conn.client_voice_stop = True
|
||||
|
||||
@@ -1,20 +1,231 @@
|
||||
import asyncio
|
||||
import gc
|
||||
import io
|
||||
import threading
|
||||
import traceback
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
from config.logger import setup_logging
|
||||
import os
|
||||
import numpy as np
|
||||
import opuslib_next
|
||||
from pydub import AudioSegment
|
||||
from abc import ABC, abstractmethod
|
||||
from core.utils import textUtils
|
||||
import queue
|
||||
|
||||
from core.providers.tts.dto.dto import MsgType, TTSMessageDTO, SentenceType
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
|
||||
|
||||
class TTSProviderBase(ABC):
|
||||
def __init__(self, config, delete_audio_file):
|
||||
self.config = config
|
||||
self.delete_audio_file = delete_audio_file
|
||||
self.output_file = config.get("output_dir")
|
||||
self.tts_text_queue = queue.Queue()
|
||||
self.tts_audio_queue = queue.Queue()
|
||||
self.enable_two_way = False
|
||||
self.stop_event = threading.Event()
|
||||
|
||||
self.tts_text_buff = []
|
||||
self.punctuations = ("。", "?", "!", ";", ":", ".", "?", "!", ";", ":", " ", ",", ",")
|
||||
self.tts_request = False
|
||||
self.processed_chars = 0
|
||||
self.stream = False
|
||||
self.last_to_opus_raw = b''
|
||||
|
||||
# 启动tts_text_queue监听线程
|
||||
# 线程任务相关
|
||||
self.loop = asyncio.get_event_loop()
|
||||
self.process_tasks_loop = asyncio.get_event_loop()
|
||||
self.max_workers = self.config.get("TTS_SET", {}).get("MAX_WORKERS", 3)
|
||||
self.active_tasks = set() # 追踪当前运行的任务
|
||||
self.executor = ThreadPoolExecutor(max_workers=self.max_workers)
|
||||
|
||||
async def open_audio_channels(self):
|
||||
pass
|
||||
|
||||
async def reset(self):
|
||||
try:
|
||||
logger.bind(tag=TAG).info("说明开始了新的对话,重建tts监听")
|
||||
await self.stop_listen_resource()
|
||||
self.tts_text_queue = queue.Queue()
|
||||
self.tts_audio_queue = queue.Queue()
|
||||
# 启动tts_text_queue监听线程
|
||||
self.stop_event.clear()
|
||||
tts_priority = threading.Thread(target=self._tts_text_priority_thread, daemon=True)
|
||||
tts_priority.start()
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"Failed to process TTS text: {e}")
|
||||
traceback.print_exc()
|
||||
|
||||
async def stop_listen_resource(self):
|
||||
"""资源清理方法"""
|
||||
self.stop_event.set()
|
||||
self.tts_text_queue = None
|
||||
self.tts_audio_queue = None
|
||||
gc.collect() # 强制执行垃圾回收
|
||||
|
||||
async def close(self):
|
||||
pass
|
||||
|
||||
def _get_segment_text(self):
|
||||
# 合并当前全部文本并处理未分割部分
|
||||
full_text = "".join(self.tts_text_buff)
|
||||
current_text = full_text[self.processed_chars:] # 从未处理的位置开始
|
||||
last_punct_pos = -1
|
||||
for punct in self.punctuations:
|
||||
pos = current_text.rfind(punct)
|
||||
if (pos != -1 and last_punct_pos == -1) or (pos != -1 and pos < last_punct_pos):
|
||||
last_punct_pos = pos
|
||||
if last_punct_pos != -1:
|
||||
segment_text_raw = current_text[:last_punct_pos + 1]
|
||||
segment_text = textUtils.get_string_no_punctuation_or_emoji(segment_text_raw)
|
||||
self.processed_chars += len(segment_text_raw) # 更新已处理字符位置
|
||||
return segment_text
|
||||
else:
|
||||
return None
|
||||
|
||||
async def process_generator(self, generator):
|
||||
async for tts_data in generator:
|
||||
self.tts_audio_queue.put(tts_data)
|
||||
|
||||
def _tts_text_priority_thread(self):
|
||||
logger.bind(tag=TAG).info("开始监听tts文本")
|
||||
if self.enable_two_way:
|
||||
self._enable_two_way_tts()
|
||||
else:
|
||||
self._no_enable_two_way_tts()
|
||||
|
||||
async def start_session(self, session_id):
|
||||
pass
|
||||
|
||||
async def finish_session(self, session_id):
|
||||
pass
|
||||
|
||||
async def tts_one_sentence(self,text):
|
||||
uuid_str = str(uuid.uuid4()).replace("-", "")
|
||||
self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.START_TTS_REQUEST, content=''))
|
||||
self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.TTS_TEXT_REQUEST, content=text))
|
||||
self.tts.tts_text_queue.put(TTSMessageDTO(u_id=uuid_str, msg_type=MsgType.STOP_TTS_REQUEST, content=text))
|
||||
|
||||
def _enable_two_way_tts(self):
|
||||
while not self.stop_event.is_set():
|
||||
try:
|
||||
ttsMessageDTO = self.tts_text_queue.get()
|
||||
msg_type = ttsMessageDTO.msg_type
|
||||
if msg_type == MsgType.START_TTS_REQUEST:
|
||||
# 开始传输tts文本
|
||||
self.tts_request = True
|
||||
self.u_id = ttsMessageDTO.u_id
|
||||
# 开启session
|
||||
future = asyncio.run_coroutine_threadsafe(self.start_session(ttsMessageDTO.u_id), loop=self.loop)
|
||||
future.result()
|
||||
# await self.start_session(ttsMessageDTO.u_id)
|
||||
elif self.tts_request and msg_type == MsgType.TTS_TEXT_REQUEST:
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
self.text_to_speak(u_id=ttsMessageDTO.u_id, text=ttsMessageDTO.content), loop=self.loop)
|
||||
future.result()
|
||||
elif msg_type == MsgType.STOP_TTS_REQUEST:
|
||||
self.tts_request = False
|
||||
future = asyncio.run_coroutine_threadsafe(self.finish_session(ttsMessageDTO.u_id), loop=self.loop)
|
||||
future.result()
|
||||
|
||||
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"Failed to process TTS text: {e}")
|
||||
# 报错了。要关闭说话
|
||||
self.tts_audio_queue.put(
|
||||
TTSMessageDTO(
|
||||
u_id=self.u_id, msg_type=MsgType.STOP_TTS_RESPONSE, content=[],
|
||||
tts_finish_text='',
|
||||
sentence_type=None
|
||||
)
|
||||
)
|
||||
traceback.print_exc()
|
||||
|
||||
def _no_enable_two_way_tts(self):
|
||||
# 为这个线程创建一个新的事件循环
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
while not self.stop_event.is_set():
|
||||
try:
|
||||
ttsMessageDTO = self.tts_text_queue.get()
|
||||
msg_type = ttsMessageDTO.msg_type
|
||||
if not self.enable_two_way:
|
||||
if msg_type == MsgType.START_TTS_REQUEST:
|
||||
# 开始传输tts文本
|
||||
self.tts_request = True
|
||||
self.processed_chars = 0
|
||||
self.tts_text_buff = []
|
||||
elif self.tts_request and msg_type == MsgType.TTS_TEXT_REQUEST:
|
||||
self.tts_text_buff.append(ttsMessageDTO.content)
|
||||
elif msg_type == MsgType.STOP_TTS_REQUEST:
|
||||
# 结束传输tts文本,处理最尾巴的数据
|
||||
self.tts_request = False
|
||||
segment_text = self._get_segment_text()
|
||||
if segment_text:
|
||||
# 修改部分:创建协程对象
|
||||
# 修改部分:创建协程对象
|
||||
tts_generator = self.text_to_speak(ttsMessageDTO.u_id, segment_text,
|
||||
True if msg_type == MsgType.STOP_TTS_REQUEST else False,
|
||||
True if msg_type == MsgType.START_TTS_REQUEST else False)
|
||||
future = asyncio.run_coroutine_threadsafe(self.process_generator(tts_generator), self.loop)
|
||||
self.active_tasks.add(future)
|
||||
if self.active_tasks:
|
||||
async def wrap_future(future):
|
||||
return await asyncio.wrap_future(future)
|
||||
|
||||
wrapped_tasks = [wrap_future(task) for task in self.active_tasks]
|
||||
done, _ = loop.run_until_complete(asyncio.wait(wrapped_tasks))
|
||||
self.active_tasks -= done
|
||||
|
||||
# 发送合成结束
|
||||
self.tts_audio_queue.put(TTSMessageDTO(u_id=ttsMessageDTO.u_id,
|
||||
msg_type=MsgType.STOP_TTS_RESPONSE,
|
||||
content=[],
|
||||
tts_finish_text='',
|
||||
sentence_type=SentenceType.SENTENCE_END))
|
||||
|
||||
segment_text = self._get_segment_text()
|
||||
if segment_text:
|
||||
# 确保这里得到的是协程对象
|
||||
tts_generator = self.text_to_speak(
|
||||
ttsMessageDTO.u_id,
|
||||
segment_text,
|
||||
msg_type == MsgType.STOP_TTS_REQUEST,
|
||||
msg_type == MsgType.START_TTS_REQUEST
|
||||
)
|
||||
# 提交协程到事件循环
|
||||
tts_generator_future = asyncio.run_coroutine_threadsafe(
|
||||
self.process_generator(tts_generator),
|
||||
loop
|
||||
)
|
||||
self.active_tasks.add(tts_generator_future)
|
||||
if len(self.active_tasks) >= self.max_workers:
|
||||
# 等待所有任务完成
|
||||
try:
|
||||
async def wrap_future(future):
|
||||
return await asyncio.wrap_future(future)
|
||||
|
||||
wrapped_tasks = [wrap_future(task) for task in self.active_tasks]
|
||||
done, _ = loop.run_until_complete(asyncio.wait(wrapped_tasks))
|
||||
self.active_tasks -= done
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"Failed to process TTS text: {e}")
|
||||
traceback.print_exc()
|
||||
else:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"Failed to process TTS text: {e}")
|
||||
traceback.print_exc()
|
||||
|
||||
@abstractmethod
|
||||
def generate_filename(self):
|
||||
@@ -38,7 +249,7 @@ class TTSProviderBase(ABC):
|
||||
logger.bind(tag=TAG).info(f"Failed to generate TTS file: {e}")
|
||||
return None
|
||||
|
||||
def to_tts_stream(self, text, queue: queue.Queue, text_index=0):
|
||||
def to_tts_stream(self, u_id, text, queue: queue.Queue, text_index=0):
|
||||
try:
|
||||
asyncio.run(self.text_to_speak_stream(text, queue, text_index))
|
||||
except Exception as e:
|
||||
@@ -46,7 +257,7 @@ class TTSProviderBase(ABC):
|
||||
return None
|
||||
|
||||
@abstractmethod
|
||||
async def text_to_speak(self, text, output_file):
|
||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||
pass
|
||||
|
||||
async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0):
|
||||
@@ -95,7 +306,17 @@ class TTSProviderBase(ABC):
|
||||
|
||||
return opus_datas, duration
|
||||
|
||||
def wav_to_opus_data_audio_raw(self, raw_data):
|
||||
def get_audio_from_tts(self, data_bytes, src_rate, to_rate=16000):
|
||||
tts_speech = torch.from_numpy(np.array(np.frombuffer(data_bytes, dtype=np.int16))).unsqueeze(dim=0)
|
||||
with io.BytesIO() as bf:
|
||||
torchaudio.save(bf, tts_speech, src_rate, format="wav")
|
||||
audio = AudioSegment.from_file(bf, format="wav")
|
||||
audio = audio.set_channels(1).set_frame_rate(to_rate)
|
||||
return audio
|
||||
|
||||
def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False):
|
||||
raw_data = self.last_to_opus_raw + raw_data_var
|
||||
self.last_to_opus_raw = b''
|
||||
# 初始化Opus编码器
|
||||
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
||||
|
||||
@@ -110,8 +331,13 @@ class TTSProviderBase(ABC):
|
||||
chunk = raw_data[i:i + frame_size * 2]
|
||||
|
||||
# 如果最后一帧不足,补零
|
||||
if len(chunk) < frame_size * 2:
|
||||
# logger.bind(tag=TAG).info("开始补0")
|
||||
# 缓存记录一下
|
||||
if len(chunk) < frame_size * 2 and not is_end:
|
||||
logger.bind(tag=TAG).info("如果最后一帧不足,缓存记录一下")
|
||||
self.last_to_opus_raw = chunk
|
||||
break
|
||||
if len(chunk) < frame_size * 2 and is_end:
|
||||
logger.bind(tag=TAG).info("是最后一句了,补零")
|
||||
chunk += b'\x00' * (frame_size * 2 - len(chunk))
|
||||
|
||||
# 转换为numpy数组处理
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
from enum import Enum
|
||||
from typing import Union
|
||||
|
||||
|
||||
class MsgType(Enum):
|
||||
# 请求类型
|
||||
START_TTS_REQUEST = "START_TTS_REQUEST"
|
||||
TTS_TEXT_REQUEST = "TTS_TEXT_REQUEST"
|
||||
STOP_TTS_REQUEST = "STOP_TTS_REQUEST"
|
||||
|
||||
# 返回类型
|
||||
START_TTS_RESPONSE = "START_TTS_RESPONSE"
|
||||
TTS_TEXT_RESPONSE = "TTS_TEXT_RESPONSE"
|
||||
STOP_TTS_RESPONSE = "STOP_TTS_RESPONSE"
|
||||
|
||||
|
||||
class SentenceType(Enum):
|
||||
# 句子开始
|
||||
SENTENCE_START = "SENTENCE_START"
|
||||
# 句子结束
|
||||
SENTENCE_END = "SENTENCE_END"
|
||||
|
||||
|
||||
class TTSMessageDTO:
|
||||
def __init__(self, u_id: str, msg_type: MsgType, content: Union[str, bytes], tts_finish_text=None,
|
||||
sentence_type: SentenceType = None, duration=0):
|
||||
if not isinstance(msg_type, MsgType):
|
||||
raise ValueError("msg_type must be an instance of MsgType Enum")
|
||||
if not isinstance(content, (str, list, bytes)):
|
||||
raise ValueError("content must be of type str or bytes")
|
||||
|
||||
# 唯一id,每个合成到合成结束,使用同一个id
|
||||
self.u_id = u_id
|
||||
self.msg_type = msg_type
|
||||
self.sentence_type = sentence_type
|
||||
self.content = content
|
||||
self.tts_finish_text = tts_finish_text
|
||||
self.duration = duration
|
||||
|
||||
def __repr__(self):
|
||||
content_preview = self.content if isinstance(self.content, str) else "<binary data>"
|
||||
return f"MessageDTO(msg_type={self.msg_type}, content={content_preview})"
|
||||
@@ -1,9 +1,17 @@
|
||||
import io
|
||||
import os
|
||||
import uuid
|
||||
import edge_tts
|
||||
from datetime import datetime
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
|
||||
from pydub import AudioSegment
|
||||
|
||||
from config.logger import setup_logging
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
|
||||
class TTSProvider(TTSProviderBase):
|
||||
def __init__(self, config, delete_audio_file):
|
||||
@@ -13,6 +21,25 @@ class TTSProvider(TTSProviderBase):
|
||||
def generate_filename(self, extension=".mp3"):
|
||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
||||
|
||||
async def text_to_speak(self, text, output_file):
|
||||
communicate = edge_tts.Communicate(text, voice=self.voice) # Use your preferred voice
|
||||
await communicate.save(output_file)
|
||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||
try:
|
||||
communicate = edge_tts.Communicate(text, voice=self.voice) # Use your preferred voice
|
||||
tmp_file = self.generate_filename()
|
||||
await communicate.save(tmp_file)
|
||||
|
||||
# 使用 pydub 读取临时文件
|
||||
audio = AudioSegment.from_file(tmp_file, format="mp3")
|
||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,sentence_type=SentenceType.SENTENCE_START)
|
||||
# 用完后删除临时文件
|
||||
try:
|
||||
os.remove(tmp_file)
|
||||
except FileNotFoundError:
|
||||
# 若文件不存在,忽略该异常
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"TTSProvider text_to_speak error: {e}")
|
||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[], tts_finish_text=text,sentence_type=SentenceType.SENTENCE_START)
|
||||
|
||||
|
||||
|
||||
@@ -17,6 +17,8 @@ from pydub import AudioSegment
|
||||
from typing_extensions import Annotated
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
||||
from core.utils.util import check_model_key
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
from config.logger import setup_logging
|
||||
@@ -118,55 +120,6 @@ class TTSProvider(TTSProviderBase):
|
||||
def generate_filename(self, extension=".wav"):
|
||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
||||
|
||||
async def text_to_speak(self, text, output_file):
|
||||
# Prepare reference data
|
||||
byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio]
|
||||
ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text]
|
||||
|
||||
data = {
|
||||
"text": text,
|
||||
"references": [
|
||||
ServeReferenceAudio(
|
||||
audio=audio if audio else b"", text=text
|
||||
)
|
||||
for text, audio in zip(ref_texts, byte_audios)
|
||||
],
|
||||
"reference_id": self.reference_id,
|
||||
"normalize": self.normalize,
|
||||
"format": self.format,
|
||||
"max_new_tokens": self.max_new_tokens,
|
||||
"chunk_length": self.chunk_length,
|
||||
"top_p": self.top_p,
|
||||
"repetition_penalty": self.repetition_penalty,
|
||||
"temperature": self.temperature,
|
||||
"streaming": self.streaming,
|
||||
"use_memory_cache": self.use_memory_cache,
|
||||
"seed": self.seed,
|
||||
}
|
||||
|
||||
pydantic_data = ServeTTSRequest(**data)
|
||||
|
||||
response = requests.post(
|
||||
self.api_url,
|
||||
data=ormsgpack.packb(pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC),
|
||||
headers={
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/msgpack",
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
audio_content = response.content
|
||||
|
||||
with open(output_file, "wb") as audio_file:
|
||||
audio_file.write(audio_content)
|
||||
|
||||
|
||||
|
||||
else:
|
||||
print(f"Request failed with status code {response.status_code}")
|
||||
print(response.json())
|
||||
|
||||
def _get_audio_from_tts(self, data_bytes):
|
||||
tts_speech = torch.from_numpy(np.array(np.frombuffer(data_bytes, dtype=np.int16))).unsqueeze(dim=0)
|
||||
with io.BytesIO() as bf:
|
||||
@@ -175,7 +128,7 @@ class TTSProvider(TTSProviderBase):
|
||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||
return audio
|
||||
|
||||
async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0):
|
||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||
try:
|
||||
data = {
|
||||
"text": text,
|
||||
@@ -219,6 +172,7 @@ class TTSProvider(TTSProviderBase):
|
||||
},
|
||||
) as response:
|
||||
if response.status_code == 200:
|
||||
index = 0
|
||||
for chunk in response.iter_content():
|
||||
# 拼接当前块和上一块数据
|
||||
chunk_total += chunk
|
||||
@@ -228,42 +182,29 @@ class TTSProvider(TTSProviderBase):
|
||||
audio_raw = audio_raw + audio.raw_data
|
||||
# 长度凑够2贞开始发送,60ms*4=240ms
|
||||
if len(audio_raw) >= 7680:
|
||||
duration = 60 * len(audio_raw) // 1920
|
||||
if (len(audio_raw) % 1920) > 0:
|
||||
duration += 60
|
||||
duration = duration / 1000.0
|
||||
# logger.bind(tag=TAG).info(f'发送数据长度:{len(audio_raw)}')
|
||||
opus_datas = self.wav_to_opus_data_audio_raw(audio_raw)
|
||||
queue.put({
|
||||
"data": opus_datas,
|
||||
"duration": duration,
|
||||
"end": False,
|
||||
"text_index": text_index
|
||||
})
|
||||
if index == 0:
|
||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||
content=opus_datas,
|
||||
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_START)
|
||||
else:
|
||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||
content=opus_datas,
|
||||
tts_finish_text=text, sentence_type=None)
|
||||
audio_raw = b''
|
||||
chunk_total = b''
|
||||
if len(chunk_total) > 0:
|
||||
audio = self._get_audio_from_tts(chunk_total)
|
||||
audio_raw = audio_raw + audio.raw_data
|
||||
duration = 60 * len(audio_raw) // 1920
|
||||
if (len(audio_raw) % 1920) > 0:
|
||||
duration += 60
|
||||
duration = duration / 1000.0
|
||||
# 把 audio 转成 opus
|
||||
# logger.bind(tag=TAG).info(f'发送数据长度:{len(audio_raw)}')
|
||||
opus_datas = self.wav_to_opus_data_audio_raw(audio_raw)
|
||||
queue.put({
|
||||
"data": opus_datas,
|
||||
"duration": duration,
|
||||
"end": False
|
||||
})
|
||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas,
|
||||
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END)
|
||||
else:
|
||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[],
|
||||
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END)
|
||||
|
||||
else:
|
||||
print('请求失败:', response.status_code, response.text)
|
||||
queue.put({
|
||||
"data": None,
|
||||
"end": True
|
||||
})
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error("tts发生错误")
|
||||
traceback.print_exc()
|
||||
|
||||
@@ -0,0 +1,390 @@
|
||||
import asyncio
|
||||
import io
|
||||
import os
|
||||
import subprocess
|
||||
import threading
|
||||
import traceback
|
||||
import uuid
|
||||
import json
|
||||
import base64
|
||||
import requests
|
||||
from datetime import datetime
|
||||
from mutagen.oggopus import OggOpus
|
||||
|
||||
import websockets
|
||||
|
||||
from config.logger import setup_logging
|
||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
||||
from core.utils.util import check_model_key
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
|
||||
PROTOCOL_VERSION = 0b0001
|
||||
DEFAULT_HEADER_SIZE = 0b0001
|
||||
|
||||
# Message Type:
|
||||
FULL_CLIENT_REQUEST = 0b0001
|
||||
AUDIO_ONLY_RESPONSE = 0b1011
|
||||
FULL_SERVER_RESPONSE = 0b1001
|
||||
ERROR_INFORMATION = 0b1111
|
||||
|
||||
# Message Type Specific Flags
|
||||
MsgTypeFlagNoSeq = 0b0000 # Non-terminal packet with no sequence
|
||||
MsgTypeFlagPositiveSeq = 0b1 # Non-terminal packet with sequence > 0
|
||||
MsgTypeFlagLastNoSeq = 0b10 # last packet with no sequence
|
||||
MsgTypeFlagNegativeSeq = 0b11 # Payload contains event number (int32)
|
||||
MsgTypeFlagWithEvent = 0b100
|
||||
# Message Serialization
|
||||
NO_SERIALIZATION = 0b0000
|
||||
JSON = 0b0001
|
||||
# Message Compression
|
||||
COMPRESSION_NO = 0b0000
|
||||
COMPRESSION_GZIP = 0b0001
|
||||
|
||||
EVENT_NONE = 0
|
||||
EVENT_Start_Connection = 1
|
||||
|
||||
EVENT_FinishConnection = 2
|
||||
|
||||
EVENT_ConnectionStarted = 50 # 成功建连
|
||||
|
||||
EVENT_ConnectionFailed = 51 # 建连失败(可能是无法通过权限认证)
|
||||
|
||||
EVENT_ConnectionFinished = 52 # 连接结束
|
||||
|
||||
# 上行Session事件
|
||||
EVENT_StartSession = 100
|
||||
|
||||
EVENT_FinishSession = 102
|
||||
# 下行Session事件
|
||||
EVENT_SessionStarted = 150
|
||||
EVENT_SessionFinished = 152
|
||||
|
||||
EVENT_SessionFailed = 153
|
||||
|
||||
# 上行通用事件
|
||||
EVENT_TaskRequest = 200
|
||||
|
||||
# 下行TTS事件
|
||||
EVENT_TTSSentenceStart = 350
|
||||
|
||||
EVENT_TTSSentenceEnd = 351
|
||||
|
||||
EVENT_TTSResponse = 352
|
||||
|
||||
|
||||
class Header:
|
||||
def __init__(self,
|
||||
protocol_version=PROTOCOL_VERSION,
|
||||
header_size=DEFAULT_HEADER_SIZE,
|
||||
message_type: int = 0,
|
||||
message_type_specific_flags: int = 0,
|
||||
serial_method: int = NO_SERIALIZATION,
|
||||
compression_type: int = COMPRESSION_NO,
|
||||
reserved_data=0):
|
||||
self.header_size = header_size
|
||||
self.protocol_version = protocol_version
|
||||
self.message_type = message_type
|
||||
self.message_type_specific_flags = message_type_specific_flags
|
||||
self.serial_method = serial_method
|
||||
self.compression_type = compression_type
|
||||
self.reserved_data = reserved_data
|
||||
|
||||
def as_bytes(self) -> bytes:
|
||||
return bytes([
|
||||
(self.protocol_version << 4) | self.header_size,
|
||||
(self.message_type << 4) | self.message_type_specific_flags,
|
||||
(self.serial_method << 4) | self.compression_type,
|
||||
self.reserved_data
|
||||
])
|
||||
|
||||
|
||||
class Optional:
|
||||
def __init__(self, event: int = EVENT_NONE, sessionId: str = None, sequence: int = None):
|
||||
self.event = event
|
||||
self.sessionId = sessionId
|
||||
self.errorCode: int = 0
|
||||
self.connectionId: str | None = None
|
||||
self.response_meta_json: str | None = None
|
||||
self.sequence = sequence
|
||||
|
||||
# 转成 byte 序列
|
||||
def as_bytes(self) -> bytes:
|
||||
option_bytes = bytearray()
|
||||
if self.event != EVENT_NONE:
|
||||
option_bytes.extend(self.event.to_bytes(4, "big", signed=True))
|
||||
if self.sessionId is not None:
|
||||
session_id_bytes = str.encode(self.sessionId)
|
||||
size = len(session_id_bytes).to_bytes(4, "big", signed=True)
|
||||
option_bytes.extend(size)
|
||||
option_bytes.extend(session_id_bytes)
|
||||
if self.sequence is not None:
|
||||
option_bytes.extend(self.sequence.to_bytes(4, "big", signed=True))
|
||||
return option_bytes
|
||||
|
||||
|
||||
class Response:
|
||||
def __init__(self, header: Header, optional: Optional):
|
||||
self.optional = optional
|
||||
self.header = header
|
||||
self.payload: bytes | None = None
|
||||
|
||||
def __str__(self):
|
||||
return super().__str__()
|
||||
|
||||
|
||||
class TTSProvider(TTSProviderBase):
|
||||
def __init__(self, config, delete_audio_file):
|
||||
super().__init__(config, delete_audio_file)
|
||||
self.appId = config.get("appid")
|
||||
self.access_token = config.get("access_token")
|
||||
self.cluster = config.get("cluster")
|
||||
self.voice = config.get("voice")
|
||||
self.ws_url = config.get("ws_url")
|
||||
self.authorization = config.get("authorization")
|
||||
self.speaker = config.get("speaker")
|
||||
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
||||
self.stop_event_response = threading.Event()
|
||||
self.enable_two_way = True
|
||||
self.start_connection_flag = False
|
||||
self.tts_text = ""
|
||||
|
||||
async def open_audio_channels(self):
|
||||
self.loop_tts = asyncio.new_event_loop()
|
||||
ws_header = {
|
||||
"X-Api-App-Key": self.appId,
|
||||
"X-Api-Access-Key": self.access_token,
|
||||
"X-Api-Resource-Id": 'volc.service_type.10029',
|
||||
"X-Api-Connect-Id": uuid.uuid4(),
|
||||
}
|
||||
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()
|
||||
|
||||
def generate_filename(self, extension=".wav"):
|
||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
||||
|
||||
async def send_event(self, header: bytes, optional: bytes | None = None,
|
||||
payload: bytes = None):
|
||||
full_client_request = bytearray(header)
|
||||
if optional is not None:
|
||||
full_client_request.extend(optional)
|
||||
if payload is not None:
|
||||
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)
|
||||
|
||||
async def send_text(self, speaker: str, text: str, session_id):
|
||||
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=speaker)
|
||||
return await self.send_event(header, optional, payload)
|
||||
|
||||
# 读取 res 数组某段 字符串内容
|
||||
def read_res_content(self, res: bytes, offset: int):
|
||||
content_size = int.from_bytes(res[offset: offset + 4], "big", signed=True)
|
||||
offset += 4
|
||||
content = str(res[offset: offset + content_size])
|
||||
offset += content_size
|
||||
return content, offset
|
||||
|
||||
# 读取 payload
|
||||
def read_res_payload(self, res: bytes, offset: int):
|
||||
payload_size = int.from_bytes(res[offset: offset + 4], "big", signed=True)
|
||||
offset += 4
|
||||
payload = res[offset: offset + payload_size]
|
||||
offset += payload_size
|
||||
return payload, offset
|
||||
|
||||
def parser_response(self, res) -> Response:
|
||||
if isinstance(res, str):
|
||||
raise RuntimeError(res)
|
||||
response = Response(Header(), Optional())
|
||||
# 解析结果
|
||||
# header
|
||||
header = response.header
|
||||
num = 0b00001111
|
||||
header.protocol_version = res[0] >> 4 & num
|
||||
header.header_size = res[0] & 0x0f
|
||||
header.message_type = (res[1] >> 4) & num
|
||||
header.message_type_specific_flags = res[1] & 0x0f
|
||||
header.serialization_method = res[2] >> num
|
||||
header.message_compression = res[2] & 0x0f
|
||||
header.reserved = res[3]
|
||||
#
|
||||
offset = 4
|
||||
optional = response.optional
|
||||
if header.message_type == FULL_SERVER_RESPONSE or AUDIO_ONLY_RESPONSE:
|
||||
# read event
|
||||
if header.message_type_specific_flags == MsgTypeFlagWithEvent:
|
||||
optional.event = int.from_bytes(res[offset:8], "big", signed=True)
|
||||
offset += 4
|
||||
if optional.event == EVENT_NONE:
|
||||
return response
|
||||
# read connectionId
|
||||
elif optional.event == EVENT_ConnectionStarted:
|
||||
optional.connectionId, offset = self.read_res_content(res, offset)
|
||||
elif optional.event == EVENT_ConnectionFailed:
|
||||
optional.response_meta_json, offset = self.read_res_content(res, offset)
|
||||
elif (optional.event == EVENT_SessionStarted
|
||||
or optional.event == EVENT_SessionFailed
|
||||
or optional.event == EVENT_SessionFinished):
|
||||
optional.sessionId, offset = self.read_res_content(res, offset)
|
||||
optional.response_meta_json, offset = self.read_res_content(res, offset)
|
||||
else:
|
||||
optional.sessionId, offset = self.read_res_content(res, offset)
|
||||
response.payload, offset = self.read_res_payload(res, offset)
|
||||
|
||||
elif header.message_type == ERROR_INFORMATION:
|
||||
optional.errorCode = int.from_bytes(res[offset:offset + 4], "big", signed=True)
|
||||
offset += 4
|
||||
response.payload, offset = self.read_res_payload(res, offset)
|
||||
return response
|
||||
|
||||
async def start_connection(self):
|
||||
header = Header(message_type=FULL_CLIENT_REQUEST, message_type_specific_flags=MsgTypeFlagWithEvent).as_bytes()
|
||||
optional = Optional(event=EVENT_Start_Connection).as_bytes()
|
||||
payload = str.encode("{}")
|
||||
return await self.send_event(header, optional, payload)
|
||||
|
||||
def print_response(self, res, tag_msg: str):
|
||||
logger.bind(tag=TAG).info(f'===>{tag_msg} header:{res.header.__dict__}')
|
||||
logger.bind(tag=TAG).info(f'===>{tag_msg} optional:{res.optional.__dict__}')
|
||||
|
||||
def get_payload_bytes(self, uid='1234', event=EVENT_NONE, text='', speaker='', audio_format='pcm',
|
||||
audio_sample_rate=16000):
|
||||
return str.encode(json.dumps(
|
||||
{
|
||||
"user": {"uid": uid},
|
||||
"event": event,
|
||||
"namespace": "BidirectionalTTS",
|
||||
"req_params": {
|
||||
"text": text,
|
||||
"speaker": speaker,
|
||||
"audio_params": {
|
||||
"format": audio_format,
|
||||
"sample_rate": audio_sample_rate
|
||||
}
|
||||
}
|
||||
}
|
||||
))
|
||||
|
||||
async def finish_connection(self):
|
||||
header = Header(message_type=FULL_CLIENT_REQUEST,
|
||||
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||
serial_method=JSON
|
||||
).as_bytes()
|
||||
optional = Optional(event=EVENT_FinishConnection).as_bytes()
|
||||
payload = str.encode('{}')
|
||||
await self.send_event(header, optional, payload)
|
||||
return
|
||||
|
||||
async def start_session(self, session_id):
|
||||
self.stop_event_response.clear()
|
||||
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.speaker)
|
||||
await self.send_event(header, optional, payload)
|
||||
|
||||
async def finish_session(self, session_id):
|
||||
self.stop_event_response.set()
|
||||
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(header, optional, payload)
|
||||
return
|
||||
|
||||
async def reset(self):
|
||||
# 关闭之前的对话
|
||||
if self.start_connection_flag:
|
||||
await self.finish_connection()
|
||||
self.start_connection_flag = False
|
||||
await self.start_connection()
|
||||
self.start_connection_flag = True
|
||||
await super().reset()
|
||||
|
||||
async def close(self):
|
||||
"""资源清理方法"""
|
||||
await self.ws.close()
|
||||
|
||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||
# 发送文本
|
||||
await self.send_text(self.speaker, text, u_id)
|
||||
return
|
||||
|
||||
def _start_monitor_tts_response_thread(self):
|
||||
# 初始化链接
|
||||
asyncio.run_coroutine_threadsafe(self._start_monitor_tts_response(), loop=self.loop)
|
||||
|
||||
async def _start_monitor_tts_response(self):
|
||||
chunk_total = b''
|
||||
while True:
|
||||
try:
|
||||
msg = await self.ws.recv() # 确保 `recv()` 运行在同一个 event loop
|
||||
res = self.parser_response(msg)
|
||||
self.print_response(res, 'send_text res:')
|
||||
|
||||
if res.optional.event == EVENT_TTSResponse and res.header.message_type == AUDIO_ONLY_RESPONSE:
|
||||
logger.bind(tag=TAG).info(f'推送数据到队列里面~~')
|
||||
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
||||
self.tts_audio_queue.put(
|
||||
TTSMessageDTO(
|
||||
u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas,
|
||||
tts_finish_text="", sentence_type=None, duration=0
|
||||
)
|
||||
)
|
||||
elif res.optional.event == EVENT_TTSSentenceStart:
|
||||
json_data = json.loads(res.payload.decode('utf-8'))
|
||||
self.tts_text = json_data.get("text", "")
|
||||
logger.bind(tag=TAG).info(f'句子开始~~{self.tts_text}')
|
||||
self.tts_audio_queue.put(
|
||||
TTSMessageDTO(
|
||||
u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[],
|
||||
tts_finish_text=self.tts_text,
|
||||
sentence_type=SentenceType.SENTENCE_START
|
||||
)
|
||||
)
|
||||
elif res.optional.event == EVENT_TTSSentenceEnd:
|
||||
logger.bind(tag=TAG).info(f'句子结束~~{self.tts_text}')
|
||||
self.tts_audio_queue.put(
|
||||
TTSMessageDTO(
|
||||
u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[],
|
||||
tts_finish_text=self.tts_text,
|
||||
sentence_type=SentenceType.SENTENCE_END
|
||||
)
|
||||
)
|
||||
elif res.optional.event == EVENT_SessionFinished:
|
||||
logger.bind(tag=TAG).info(f'会话结束~~,最后一句补零')
|
||||
opus_datas = self.wav_to_opus_data_audio_raw(b'', is_end=True)
|
||||
self.tts_audio_queue.put(
|
||||
TTSMessageDTO(
|
||||
u_id=self.u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas,
|
||||
tts_finish_text="", sentence_type=None, duration=0
|
||||
)
|
||||
)
|
||||
self.tts_audio_queue.put(
|
||||
TTSMessageDTO(
|
||||
u_id=self.u_id, msg_type=MsgType.STOP_TTS_RESPONSE, content=[],
|
||||
tts_finish_text=self.tts_text,
|
||||
sentence_type=SentenceType.SENTENCE_END
|
||||
)
|
||||
)
|
||||
else:
|
||||
continue
|
||||
except websockets.ConnectionClosed:
|
||||
break # 连接关闭时退出监听
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"Error in _start_monitor_tts_response: {e}")
|
||||
traceback.print_exc()
|
||||
continue
|
||||
@@ -0,0 +1,34 @@
|
||||
def get_string_no_punctuation_or_emoji(s):
|
||||
"""去除字符串首尾的空格、标点符号和表情符号"""
|
||||
chars = list(s)
|
||||
# 处理开头的字符
|
||||
start = 0
|
||||
while start < len(chars) and is_punctuation_or_emoji(chars[start]):
|
||||
start += 1
|
||||
# 处理结尾的字符
|
||||
end = len(chars) - 1
|
||||
while end >= start and is_punctuation_or_emoji(chars[end]):
|
||||
end -= 1
|
||||
return ''.join(chars[start:end + 1])
|
||||
|
||||
def is_punctuation_or_emoji(char):
|
||||
"""检查字符是否为空格、指定标点或表情符号"""
|
||||
# 定义需要去除的中英文标点(包括全角/半角)
|
||||
punctuation_set = {
|
||||
',', ',', # 中文逗号 + 英文逗号
|
||||
'。', '.', # 中文句号 + 英文句号
|
||||
'!', '!', # 中文感叹号 + 英文感叹号
|
||||
'-', '-', # 英文连字符 + 中文全角横线
|
||||
'、' # 中文顿号
|
||||
}
|
||||
if char.isspace() or char in punctuation_set:
|
||||
return True
|
||||
# 检查表情符号(保留原有逻辑)
|
||||
code_point = ord(char)
|
||||
emoji_ranges = [
|
||||
(0x1F600, 0x1F64F), (0x1F300, 0x1F5FF),
|
||||
(0x1F680, 0x1F6FF), (0x1F900, 0x1F9FF),
|
||||
(0x1FA70, 0x1FAFF), (0x2600, 0x26FF),
|
||||
(0x2700, 0x27BF)
|
||||
]
|
||||
return any(start <= code_point <= end for start, end in emoji_ranges)
|
||||
@@ -12,7 +12,7 @@ class WebSocketServer:
|
||||
def __init__(self, config: dict):
|
||||
self.config = config
|
||||
self.logger = setup_logging()
|
||||
self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = self._create_processing_instances()
|
||||
self._vad, self._asr, self._llm, self._memory, self.intent = self._create_processing_instances()
|
||||
self.active_connections = set() # 添加全局连接记录
|
||||
|
||||
def _create_processing_instances(self):
|
||||
@@ -41,14 +41,6 @@ class WebSocketServer:
|
||||
self.config["LLM"][self.config["selected_module"]["LLM"]]['type'],
|
||||
self.config["LLM"][self.config["selected_module"]["LLM"]],
|
||||
),
|
||||
tts.create_instance(
|
||||
self.config["selected_module"]["TTS"]
|
||||
if not 'type' in self.config["TTS"][self.config["selected_module"]["TTS"]]
|
||||
else
|
||||
self.config["TTS"][self.config["selected_module"]["TTS"]]["type"],
|
||||
self.config["TTS"][self.config["selected_module"]["TTS"]],
|
||||
self.config["delete_audio"]
|
||||
),
|
||||
memory.create_instance(memory_cls_name, memory_cfg),
|
||||
intent.create_instance(
|
||||
self.config["selected_module"]["Intent"]
|
||||
@@ -78,7 +70,17 @@ class WebSocketServer:
|
||||
async def _handle_connection(self, websocket):
|
||||
"""处理新连接,每次创建独立的ConnectionHandler"""
|
||||
# 创建ConnectionHandler时传入当前server实例
|
||||
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._memory, self.intent)
|
||||
# tts 变成链接的时候创建,避免并非问题
|
||||
f_tts = tts.create_instance(
|
||||
self.config["selected_module"]["TTS"]
|
||||
if not 'type' in self.config["TTS"][self.config["selected_module"]["TTS"]]
|
||||
else
|
||||
self.config["TTS"][self.config["selected_module"]["TTS"]]["type"],
|
||||
self.config["TTS"][self.config["selected_module"]["TTS"]],
|
||||
self.config["delete_audio"]
|
||||
)
|
||||
await f_tts.open_audio_channels()
|
||||
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, f_tts, self._memory, self.intent)
|
||||
self.active_connections.add(handler)
|
||||
try:
|
||||
await handler.handle_connection(websocket)
|
||||
|
||||
@@ -7,6 +7,8 @@ import asyncio
|
||||
import difflib
|
||||
import traceback
|
||||
from pathlib import Path
|
||||
|
||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType
|
||||
from core.utils import p3
|
||||
from core.handle.sendAudioHandle import send_stt_message
|
||||
from plugins_func.register import register_function,ToolType, ActionResponse, Action
|
||||
@@ -189,7 +191,12 @@ async def play_local_music(conn, specific_file=None):
|
||||
opus_packets, duration = p3.decode_opus_from_file(music_path)
|
||||
else:
|
||||
opus_packets, duration = conn.tts.audio_to_opus_data(music_path)
|
||||
conn.audio_play_queue.put((opus_packets, selected_music, 0))
|
||||
conn.tts.tts_audio_queue.put(
|
||||
TTSMessageDTO(
|
||||
u_id="", msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_packets,
|
||||
tts_finish_text="", sentence_type=None, duration=0
|
||||
)
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}")
|
||||
|
||||
@@ -22,3 +22,4 @@ mem0ai==0.1.62
|
||||
bs4==0.0.2
|
||||
modelscope==1.23.2
|
||||
sherpa_onnx==1.11.0
|
||||
mutagen==1.47.0
|
||||
|
||||
Reference in New Issue
Block a user