add:火山双向tts语音流式输入输出

This commit is contained in:
lizhongxiang
2025-03-29 14:17:09 +08:00
parent 1136ca4b24
commit 264487574b
15 changed files with 839 additions and 320 deletions
+3
View File
@@ -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
+9
View File
@@ -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
+45 -149
View File
@@ -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
+231 -5
View File
@@ -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})"
+31 -4
View File
@@ -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 -10
View File
@@ -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)}")
+1
View File
@@ -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