mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 17:33:56 +08:00
@@ -61,10 +61,6 @@ delete_audio: true
|
|||||||
close_connection_no_voice_time: 120
|
close_connection_no_voice_time: 120
|
||||||
# TTS请求超时时间(秒)
|
# TTS请求超时时间(秒)
|
||||||
tts_timeout: 10
|
tts_timeout: 10
|
||||||
# 开启唤醒词加速
|
|
||||||
enable_wakeup_words_response_cache: true
|
|
||||||
# 开场是否回复唤醒词
|
|
||||||
enable_greeting: true
|
|
||||||
# 说完话是否开启提示音
|
# 说完话是否开启提示音
|
||||||
enable_stop_tts_notify: false
|
enable_stop_tts_notify: false
|
||||||
# 说完话是否开启提示音,音效地址
|
# 说完话是否开启提示音,音效地址
|
||||||
|
|||||||
@@ -664,16 +664,14 @@ class ConnectionHandler:
|
|||||||
# 更新系统prompt至上下文
|
# 更新系统prompt至上下文
|
||||||
self.dialogue.update_system_message(self.prompt)
|
self.dialogue.update_system_message(self.prompt)
|
||||||
|
|
||||||
def chat(self, query, tool_call=False, depth=0):
|
def chat(self, query, depth=0):
|
||||||
self.logger.bind(tag=TAG).info(f"大模型收到用户消息: {query}")
|
self.logger.bind(tag=TAG).info(f"大模型收到用户消息: {query}")
|
||||||
self.llm_finish_task = False
|
self.llm_finish_task = False
|
||||||
|
|
||||||
if not tool_call:
|
|
||||||
self.dialogue.put(Message(role="user", content=query))
|
|
||||||
|
|
||||||
# 为最顶层时新建会话ID和发送FIRST请求
|
# 为最顶层时新建会话ID和发送FIRST请求
|
||||||
if depth == 0:
|
if depth == 0:
|
||||||
self.sentence_id = str(uuid.uuid4().hex)
|
self.sentence_id = str(uuid.uuid4().hex)
|
||||||
|
self.dialogue.put(Message(role="user", content=query))
|
||||||
self.tts.tts_text_queue.put(
|
self.tts.tts_text_queue.put(
|
||||||
TTSMessageDTO(
|
TTSMessageDTO(
|
||||||
sentence_id=self.sentence_id,
|
sentence_id=self.sentence_id,
|
||||||
@@ -878,7 +876,7 @@ class ConnectionHandler:
|
|||||||
content=text,
|
content=text,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
self.chat(text, tool_call=True, depth=depth + 1)
|
self.chat(text, depth=depth + 1)
|
||||||
elif result.action == Action.NOTFOUND or result.action == Action.ERROR:
|
elif result.action == Action.NOTFOUND or result.action == Action.ERROR:
|
||||||
text = result.response if result.response else result.result
|
text = result.response if result.response else result.result
|
||||||
self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text)
|
self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text)
|
||||||
|
|||||||
@@ -1,32 +1,13 @@
|
|||||||
import time
|
|
||||||
import json
|
import json
|
||||||
import random
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from core.utils.dialogue import Message
|
|
||||||
from core.utils.util import audio_to_data
|
|
||||||
from core.handle.sendAudioHandle import sendAudioMessage, send_stt_message
|
|
||||||
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
|
|
||||||
from core.providers.tts.dto.dto import ContentType, SentenceType
|
|
||||||
from core.providers.tools.device_mcp import (
|
from core.providers.tools.device_mcp import (
|
||||||
MCPClient,
|
MCPClient,
|
||||||
send_mcp_initialize_message,
|
send_mcp_initialize_message,
|
||||||
send_mcp_tools_list_request,
|
send_mcp_tools_list_request,
|
||||||
)
|
)
|
||||||
from core.utils.wakeup_word import WakeupWordsConfig
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
WAKEUP_CONFIG = {
|
|
||||||
"refresh_time": 5,
|
|
||||||
"words": ["你好", "你好啊", "嘿,你好", "嗨"],
|
|
||||||
}
|
|
||||||
|
|
||||||
# 创建全局的唤醒词配置管理器
|
|
||||||
wakeup_words_config = WakeupWordsConfig()
|
|
||||||
|
|
||||||
# 用于防止并发调用wakeupWordsResponse的锁
|
|
||||||
_wakeup_response_lock = asyncio.Lock()
|
|
||||||
|
|
||||||
|
|
||||||
async def handleHelloMessage(conn, msg_json):
|
async def handleHelloMessage(conn, msg_json):
|
||||||
"""处理hello消息"""
|
"""处理hello消息"""
|
||||||
@@ -49,93 +30,3 @@ async def handleHelloMessage(conn, msg_json):
|
|||||||
asyncio.create_task(send_mcp_tools_list_request(conn))
|
asyncio.create_task(send_mcp_tools_list_request(conn))
|
||||||
|
|
||||||
await conn.websocket.send(json.dumps(conn.welcome_msg))
|
await conn.websocket.send(json.dumps(conn.welcome_msg))
|
||||||
|
|
||||||
|
|
||||||
async def checkWakeupWords(conn, text):
|
|
||||||
enable_wakeup_words_response_cache = conn.config[
|
|
||||||
"enable_wakeup_words_response_cache"
|
|
||||||
]
|
|
||||||
|
|
||||||
if not enable_wakeup_words_response_cache or not conn.tts:
|
|
||||||
return False
|
|
||||||
|
|
||||||
_, filtered_text = remove_punctuation_and_length(text)
|
|
||||||
if filtered_text not in conn.config.get("wakeup_words"):
|
|
||||||
return False
|
|
||||||
|
|
||||||
conn.just_woken_up = True
|
|
||||||
await send_stt_message(conn, text)
|
|
||||||
|
|
||||||
# 获取当前音色
|
|
||||||
voice = getattr(conn.tts, "voice", "default")
|
|
||||||
if not voice:
|
|
||||||
voice = "default"
|
|
||||||
|
|
||||||
# 获取唤醒词回复配置
|
|
||||||
response = wakeup_words_config.get_wakeup_response(voice)
|
|
||||||
if not response or not response.get("file_path"):
|
|
||||||
response = {
|
|
||||||
"voice": "default",
|
|
||||||
"file_path": "config/assets/wakeup_words.wav",
|
|
||||||
"time": 0,
|
|
||||||
"text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦",
|
|
||||||
}
|
|
||||||
|
|
||||||
# 播放唤醒词回复
|
|
||||||
conn.client_abort = False
|
|
||||||
opus_packets, _ = audio_to_data(response.get("file_path"))
|
|
||||||
|
|
||||||
conn.logger.bind(tag=TAG).info(f"播放唤醒词回复: {response.get('text')}")
|
|
||||||
await sendAudioMessage(conn, SentenceType.FIRST, opus_packets, response.get("text"))
|
|
||||||
await sendAudioMessage(conn, SentenceType.LAST, [], None)
|
|
||||||
|
|
||||||
# 补充对话
|
|
||||||
conn.dialogue.put(Message(role="assistant", content=response.get("text")))
|
|
||||||
|
|
||||||
# 检查是否需要更新唤醒词回复
|
|
||||||
if time.time() - response.get("time", 0) > WAKEUP_CONFIG["refresh_time"]:
|
|
||||||
if not _wakeup_response_lock.locked():
|
|
||||||
asyncio.create_task(wakeupWordsResponse(conn))
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
async def wakeupWordsResponse(conn):
|
|
||||||
if not conn.tts or not conn.llm or not conn.llm.response_no_stream:
|
|
||||||
return
|
|
||||||
|
|
||||||
try:
|
|
||||||
# 尝试获取锁,如果获取不到就返回
|
|
||||||
if not await _wakeup_response_lock.acquire():
|
|
||||||
return
|
|
||||||
|
|
||||||
# 生成唤醒词回复
|
|
||||||
wakeup_word = random.choice(WAKEUP_CONFIG["words"])
|
|
||||||
question = (
|
|
||||||
"此刻用户正在和你说```"
|
|
||||||
+ wakeup_word
|
|
||||||
+ "```。\n请你根据以上用户的内容进行20-30字回复。要符合系统设置的角色情感和态度,不要像机器人一样说话。\n"
|
|
||||||
+ "请勿对这条内容本身进行任何解释和回应,请勿返回表情符号,仅返回对用户的内容的回复。"
|
|
||||||
)
|
|
||||||
|
|
||||||
result = conn.llm.response_no_stream(conn.config["prompt"], question)
|
|
||||||
if not result or len(result) == 0:
|
|
||||||
return
|
|
||||||
|
|
||||||
# 生成TTS音频
|
|
||||||
tts_result = await asyncio.to_thread(conn.tts.to_tts, result)
|
|
||||||
if not tts_result:
|
|
||||||
return
|
|
||||||
|
|
||||||
# 获取当前音色
|
|
||||||
voice = getattr(conn.tts, "voice", "default")
|
|
||||||
|
|
||||||
wav_bytes = opus_datas_to_wav_bytes(tts_result, sample_rate=16000)
|
|
||||||
file_path = wakeup_words_config.generate_file_path(voice)
|
|
||||||
with open(file_path, "wb") as f:
|
|
||||||
f.write(wav_bytes)
|
|
||||||
# 更新配置
|
|
||||||
wakeup_words_config.update_wakeup_response(voice, file_path, result)
|
|
||||||
finally:
|
|
||||||
# 确保在任何情况下都释放锁
|
|
||||||
if _wakeup_response_lock.locked():
|
|
||||||
_wakeup_response_lock.release()
|
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ import json
|
|||||||
import asyncio
|
import asyncio
|
||||||
import uuid
|
import uuid
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
from core.handle.helloHandle import checkWakeupWords
|
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
from core.providers.tts.dto.dto import ContentType
|
from core.providers.tts.dto.dto import ContentType
|
||||||
from core.utils.dialogue import Message
|
from core.utils.dialogue import Message
|
||||||
@@ -27,9 +26,6 @@ async def handle_user_intent(conn, text):
|
|||||||
filtered_text = remove_punctuation_and_length(text)[1]
|
filtered_text = remove_punctuation_and_length(text)[1]
|
||||||
if await check_direct_exit(conn, filtered_text):
|
if await check_direct_exit(conn, filtered_text):
|
||||||
return True
|
return True
|
||||||
# 检查是否是唤醒词
|
|
||||||
if await checkWakeupWords(conn, filtered_text):
|
|
||||||
return True
|
|
||||||
|
|
||||||
if conn.intent_type == "function_call":
|
if conn.intent_type == "function_call":
|
||||||
# 使用支持function calling的聊天方法,不再进行意图分析
|
# 使用支持function calling的聊天方法,不再进行意图分析
|
||||||
|
|||||||
@@ -1,12 +1,12 @@
|
|||||||
|
import time
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
from core.handle.intentHandler import handle_user_intent
|
from core.handle.intentHandler import handle_user_intent
|
||||||
from core.utils.output_counter import check_device_output_limit
|
from core.utils.output_counter import check_device_output_limit
|
||||||
from core.handle.abortHandle import handleAbortMessage
|
from core.handle.abortHandle import handleAbortMessage
|
||||||
import time
|
|
||||||
import asyncio
|
|
||||||
import json
|
|
||||||
from core.handle.sendAudioHandle import SentenceType
|
from core.handle.sendAudioHandle import SentenceType
|
||||||
from core.utils.util import audio_to_data
|
from core.utils.util import audio_to_data_stream
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -121,8 +121,9 @@ async def max_out_size(conn):
|
|||||||
text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
|
text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
file_path = "config/assets/max_output_size.wav"
|
file_path = "config/assets/max_output_size.wav"
|
||||||
opus_packets, _ = audio_to_data(file_path)
|
conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
|
play_audio_frames(conn, file_path)
|
||||||
|
conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
conn.close_after_chat = True
|
conn.close_after_chat = True
|
||||||
|
|
||||||
|
|
||||||
@@ -140,16 +141,15 @@ async def check_bind_device(conn):
|
|||||||
|
|
||||||
# 播放提示音
|
# 播放提示音
|
||||||
music_path = "config/assets/bind_code.wav"
|
music_path = "config/assets/bind_code.wav"
|
||||||
opus_packets, _ = audio_to_data(music_path)
|
conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text))
|
play_audio_frames(conn, music_path)
|
||||||
|
|
||||||
# 逐个播放数字
|
# 逐个播放数字
|
||||||
for i in range(6): # 确保只播放6位数字
|
for i in range(6): # 确保只播放6位数字
|
||||||
try:
|
try:
|
||||||
digit = conn.bind_code[i]
|
digit = conn.bind_code[i]
|
||||||
num_path = f"config/assets/bind_code/{digit}.wav"
|
num_path = f"config/assets/bind_code/{digit}.wav"
|
||||||
num_packets, _ = audio_to_data(num_path)
|
play_audio_frames(conn, num_path)
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, num_packets, None))
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
|
conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
|
||||||
continue
|
continue
|
||||||
@@ -158,5 +158,18 @@ async def check_bind_device(conn):
|
|||||||
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
|
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
music_path = "config/assets/bind_not_found.wav"
|
music_path = "config/assets/bind_not_found.wav"
|
||||||
opus_packets, _ = audio_to_data(music_path)
|
conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
|
play_audio_frames(conn, music_path)
|
||||||
|
conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
|
|
||||||
|
|
||||||
|
def play_audio_frames(conn, file_path):
|
||||||
|
"""播放音频文件并处理发送帧数据"""
|
||||||
|
def handle_audio_frame(frame_data):
|
||||||
|
conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, frame_data, None))
|
||||||
|
|
||||||
|
audio_to_data_stream(
|
||||||
|
file_path,
|
||||||
|
is_opus=True,
|
||||||
|
callback=handle_audio_frame
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,6 +1,4 @@
|
|||||||
import json
|
import json
|
||||||
import asyncio
|
|
||||||
import time
|
|
||||||
from core.providers.tts.dto.dto import SentenceType
|
from core.providers.tts.dto.dto import SentenceType
|
||||||
from core.utils import textUtils
|
from core.utils import textUtils
|
||||||
|
|
||||||
@@ -8,18 +6,18 @@ TAG = __name__
|
|||||||
|
|
||||||
|
|
||||||
async def sendAudioMessage(conn, sentenceType, audios, text):
|
async def sendAudioMessage(conn, sentenceType, audios, text):
|
||||||
# 发送句子开始消息
|
|
||||||
conn.logger.bind(tag=TAG).info(f"发送音频消息: {sentenceType}, {text}")
|
|
||||||
|
|
||||||
pre_buffer = False
|
|
||||||
if conn.tts.tts_audio_first_sentence:
|
if conn.tts.tts_audio_first_sentence:
|
||||||
conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
||||||
conn.tts.tts_audio_first_sentence = False
|
conn.tts.tts_audio_first_sentence = False
|
||||||
pre_buffer = True
|
await send_tts_message(conn, "start", None)
|
||||||
|
|
||||||
await send_tts_message(conn, "sentence_start", text)
|
if sentenceType == SentenceType.FIRST:
|
||||||
|
await send_tts_message(conn, "sentence_start", text)
|
||||||
|
|
||||||
await sendAudio(conn, audios, pre_buffer)
|
await sendAudio(conn, audios)
|
||||||
|
# 发送句子开始消息
|
||||||
|
if sentenceType is not SentenceType.MIDDLE:
|
||||||
|
conn.logger.bind(tag=TAG).info(f"发送音频消息: {sentenceType}, {text}")
|
||||||
|
|
||||||
# 发送结束消息(如果是最后一个文本)
|
# 发送结束消息(如果是最后一个文本)
|
||||||
if conn.llm_finish_task and sentenceType == SentenceType.LAST:
|
if conn.llm_finish_task and sentenceType == SentenceType.LAST:
|
||||||
@@ -30,45 +28,18 @@ async def sendAudioMessage(conn, sentenceType, audios, text):
|
|||||||
|
|
||||||
|
|
||||||
# 播放音频
|
# 播放音频
|
||||||
async def sendAudio(conn, audios, pre_buffer=True):
|
async def sendAudio(conn, audios):
|
||||||
if audios is None or len(audios) == 0:
|
if audios is None:
|
||||||
return
|
return
|
||||||
# 流控参数优化
|
# 如果audios不是opus数组,则不需要进行遍历,可以直接发送;这里需要进行流控管理,防止发送过快引发客户端溢出
|
||||||
frame_duration = 60 # 帧时长(毫秒),匹配 Opus 编码
|
if isinstance(audios, bytes):
|
||||||
start_time = time.perf_counter()
|
await conn.websocket.send(audios)
|
||||||
play_position = 0
|
|
||||||
|
|
||||||
# 仅当第一句话时执行预缓冲
|
|
||||||
if pre_buffer:
|
|
||||||
pre_buffer_frames = min(3, len(audios))
|
|
||||||
for i in range(pre_buffer_frames):
|
|
||||||
await conn.websocket.send(audios[i])
|
|
||||||
remaining_audios = audios[pre_buffer_frames:]
|
|
||||||
else:
|
|
||||||
remaining_audios = audios
|
|
||||||
|
|
||||||
# 播放剩余音频帧
|
|
||||||
for opus_packet in remaining_audios:
|
|
||||||
if conn.client_abort:
|
|
||||||
break
|
|
||||||
|
|
||||||
# 重置没有声音的状态
|
|
||||||
conn.last_activity_time = time.time() * 1000
|
|
||||||
|
|
||||||
# 计算预期发送时间
|
|
||||||
expected_time = start_time + (play_position / 1000)
|
|
||||||
current_time = time.perf_counter()
|
|
||||||
delay = expected_time - current_time
|
|
||||||
if delay > 0:
|
|
||||||
await asyncio.sleep(delay)
|
|
||||||
|
|
||||||
await conn.websocket.send(opus_packet)
|
|
||||||
|
|
||||||
play_position += frame_duration
|
|
||||||
|
|
||||||
|
|
||||||
async def send_tts_message(conn, state, text=None):
|
async def send_tts_message(conn, state, text=None):
|
||||||
"""发送 TTS 状态消息"""
|
"""发送 TTS 状态消息"""
|
||||||
|
if text is None and state == "sentence_start":
|
||||||
|
return
|
||||||
message = {"type": "tts", "state": state, "session_id": conn.session_id}
|
message = {"type": "tts", "state": state, "session_id": conn.session_id}
|
||||||
if text is not None:
|
if text is not None:
|
||||||
message["text"] = textUtils.check_emoji(text)
|
message["text"] = textUtils.check_emoji(text)
|
||||||
@@ -91,13 +62,12 @@ async def send_tts_message(conn, state, text=None):
|
|||||||
|
|
||||||
|
|
||||||
async def send_stt_message(conn, text):
|
async def send_stt_message(conn, text):
|
||||||
|
"""发送 STT 状态消息"""
|
||||||
end_prompt_str = conn.config.get("end_prompt", {}).get("prompt")
|
end_prompt_str = conn.config.get("end_prompt", {}).get("prompt")
|
||||||
if end_prompt_str and end_prompt_str == text:
|
if end_prompt_str and end_prompt_str == text:
|
||||||
await send_tts_message(conn, "start")
|
await send_tts_message(conn, "start")
|
||||||
return
|
return
|
||||||
|
|
||||||
"""发送 STT 状态消息"""
|
|
||||||
|
|
||||||
# 解析JSON格式,提取实际的用户说话内容
|
# 解析JSON格式,提取实际的用户说话内容
|
||||||
display_text = text
|
display_text = text
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ from core.handle.helloHandle import handleHelloMessage
|
|||||||
from core.providers.tools.device_mcp import handle_mcp_message
|
from core.providers.tools.device_mcp import handle_mcp_message
|
||||||
from core.utils.util import remove_punctuation_and_length, filter_sensitive_info
|
from core.utils.util import remove_punctuation_and_length, filter_sensitive_info
|
||||||
from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
|
from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
|
||||||
from core.handle.sendAudioHandle import send_stt_message, send_tts_message
|
|
||||||
from core.providers.tools.device_iot import handleIotDescriptors, handleIotStatus
|
from core.providers.tools.device_iot import handleIotDescriptors, handleIotStatus
|
||||||
from core.handle.reportHandle import enqueue_asr_report
|
from core.handle.reportHandle import enqueue_asr_report
|
||||||
import asyncio
|
import asyncio
|
||||||
@@ -51,23 +50,9 @@ async def handleTextMessage(conn, message):
|
|||||||
filtered_len, filtered_text = remove_punctuation_and_length(
|
filtered_len, filtered_text = remove_punctuation_and_length(
|
||||||
original_text
|
original_text
|
||||||
)
|
)
|
||||||
|
|
||||||
# 识别是否是唤醒词
|
# 识别是否是唤醒词
|
||||||
is_wakeup_words = filtered_text in conn.config.get("wakeup_words")
|
is_wakeup_words = filtered_text in conn.config.get("wakeup_words")
|
||||||
# 是否开启唤醒词回复
|
if not is_wakeup_words:
|
||||||
enable_greeting = conn.config.get("enable_greeting", True)
|
|
||||||
|
|
||||||
if is_wakeup_words and not enable_greeting:
|
|
||||||
# 如果是唤醒词,且关闭了唤醒词回复,就不用回答
|
|
||||||
await send_stt_message(conn, original_text)
|
|
||||||
await send_tts_message(conn, "stop", None)
|
|
||||||
conn.client_is_speaking = False
|
|
||||||
elif is_wakeup_words:
|
|
||||||
conn.just_woken_up = True
|
|
||||||
# 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
|
|
||||||
enqueue_asr_report(conn, "嘿,你好呀", [])
|
|
||||||
await startToChat(conn, "嘿,你好呀")
|
|
||||||
else:
|
|
||||||
# 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
|
# 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
|
||||||
enqueue_asr_report(conn, original_text, [])
|
enqueue_asr_report(conn, original_text, [])
|
||||||
# 否则需要LLM对文字内容进行答复
|
# 否则需要LLM对文字内容进行答复
|
||||||
|
|||||||
@@ -216,6 +216,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
if message.sentence_type == SentenceType.FIRST:
|
if message.sentence_type == SentenceType.FIRST:
|
||||||
self.conn.client_abort = False
|
self.conn.client_abort = False
|
||||||
|
self.reset_flow_controller()
|
||||||
|
|
||||||
if self.conn.client_abort:
|
if self.conn.client_abort:
|
||||||
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
|
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
|
||||||
@@ -268,11 +269,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
||||||
@@ -422,9 +419,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
async def _start_monitor_tts_response(self):
|
async def _start_monitor_tts_response(self):
|
||||||
"""监听TTS响应"""
|
"""监听TTS响应"""
|
||||||
opus_datas_cache = []
|
|
||||||
is_first_sentence = True
|
|
||||||
first_sentence_segment_count = 0 # 添加计数器
|
|
||||||
try:
|
try:
|
||||||
session_finished = False # 标记会话是否正常结束
|
session_finished = False # 标记会话是否正常结束
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
@@ -445,28 +439,16 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(SentenceType.FIRST, [], None)
|
(SentenceType.FIRST, [], None)
|
||||||
)
|
)
|
||||||
elif event_name == "SentenceBegin":
|
|
||||||
opus_datas_cache = []
|
|
||||||
elif event_name == "SentenceEnd":
|
elif event_name == "SentenceEnd":
|
||||||
if (
|
# 发送缓存的数据
|
||||||
not is_first_sentence
|
if self.conn.tts_MessageText:
|
||||||
or first_sentence_segment_count > 10
|
logger.bind(tag=TAG).info(
|
||||||
):
|
f"句子语音生成成功: {self.conn.tts_MessageText}"
|
||||||
# 发送缓存的数据
|
)
|
||||||
if self.conn.tts_MessageText:
|
self.tts_audio_queue.put(
|
||||||
logger.bind(tag=TAG).info(
|
(SentenceType.FIRST, [], self.conn.tts_MessageText)
|
||||||
f"句子语音生成成功: {self.conn.tts_MessageText}"
|
)
|
||||||
)
|
self.conn.tts_MessageText = None
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, self.conn.tts_MessageText)
|
|
||||||
)
|
|
||||||
self.conn.tts_MessageText = None
|
|
||||||
else:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
# 第一句话结束后,将标志设置为False
|
|
||||||
is_first_sentence = False
|
|
||||||
elif event_name == "SynthesisCompleted":
|
elif event_name == "SynthesisCompleted":
|
||||||
logger.bind(tag=TAG).debug(f"会话结束~~")
|
logger.bind(tag=TAG).debug(f"会话结束~~")
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -477,22 +459,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
# 二进制消息(音频数据)
|
# 二进制消息(音频数据)
|
||||||
elif isinstance(msg, (bytes, bytearray)):
|
elif isinstance(msg, (bytes, bytearray)):
|
||||||
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
||||||
opus_datas = self.opus_encoder.encode_pcm_to_opus(msg, False)
|
self.opus_encoder.encode_pcm_to_opus_stream(msg, False, self.handle_opus)
|
||||||
logger.bind(tag=TAG).debug(
|
|
||||||
f"推送数据到队列里面帧数~~{len(opus_datas)}"
|
|
||||||
)
|
|
||||||
if is_first_sentence:
|
|
||||||
first_sentence_segment_count += 1
|
|
||||||
if first_sentence_segment_count <= 6:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas, None)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
opus_datas_cache.extend(opus_datas)
|
|
||||||
else:
|
|
||||||
# 后续句子缓存
|
|
||||||
opus_datas_cache.extend(opus_datas)
|
|
||||||
|
|
||||||
except websockets.ConnectionClosed:
|
except websockets.ConnectionClosed:
|
||||||
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
|
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
|
||||||
break
|
break
|
||||||
@@ -512,142 +479,3 @@ class TTSProvider(TTSProviderBase):
|
|||||||
finally:
|
finally:
|
||||||
self._monitor_task = None
|
self._monitor_task = None
|
||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
|
||||||
"""非流式TTS处理,用于测试及保存音频文件的场景"""
|
|
||||||
try:
|
|
||||||
# 创建新的事件循环
|
|
||||||
loop = asyncio.new_event_loop()
|
|
||||||
asyncio.set_event_loop(loop)
|
|
||||||
|
|
||||||
# 生成会话ID
|
|
||||||
session_id = uuid.uuid4().hex
|
|
||||||
# 存储音频数据
|
|
||||||
audio_data = []
|
|
||||||
|
|
||||||
async def _generate_audio():
|
|
||||||
# 刷新Token(如果需要)
|
|
||||||
if self._is_token_expired():
|
|
||||||
self._refresh_token()
|
|
||||||
|
|
||||||
# 建立WebSocket连接
|
|
||||||
ws = await websockets.connect(
|
|
||||||
self.ws_url,
|
|
||||||
additional_headers={"X-NLS-Token": self.token},
|
|
||||||
ping_interval=30,
|
|
||||||
ping_timeout=10,
|
|
||||||
close_timeout=10,
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
# 发送StartSynthesis请求
|
|
||||||
start_message_id = str(uuid.uuid4().hex)
|
|
||||||
start_request = {
|
|
||||||
"header": {
|
|
||||||
"message_id": start_message_id,
|
|
||||||
"task_id": session_id,
|
|
||||||
"namespace": "FlowingSpeechSynthesizer",
|
|
||||||
"name": "StartSynthesis",
|
|
||||||
"appkey": self.appkey,
|
|
||||||
},
|
|
||||||
"payload": {
|
|
||||||
"voice": self.voice,
|
|
||||||
"format": self.format,
|
|
||||||
"sample_rate": self.sample_rate,
|
|
||||||
"volume": self.volume,
|
|
||||||
"speech_rate": self.speech_rate,
|
|
||||||
"pitch_rate": self.pitch_rate,
|
|
||||||
"enable_subtitle": True,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
await ws.send(json.dumps(start_request))
|
|
||||||
|
|
||||||
# 等待SynthesisStarted响应
|
|
||||||
synthesis_started = False
|
|
||||||
while not synthesis_started:
|
|
||||||
msg = await ws.recv()
|
|
||||||
if isinstance(msg, str):
|
|
||||||
data = json.loads(msg)
|
|
||||||
header = data.get("header", {})
|
|
||||||
if header.get("name") == "SynthesisStarted":
|
|
||||||
synthesis_started = True
|
|
||||||
logger.bind(tag=TAG).debug("TTS合成已启动")
|
|
||||||
elif header.get("name") == "TaskFailed":
|
|
||||||
error_info = data.get("payload", {}).get(
|
|
||||||
"error_info", {}
|
|
||||||
)
|
|
||||||
error_code = error_info.get("error_code")
|
|
||||||
error_message = error_info.get(
|
|
||||||
"error_message", "未知错误"
|
|
||||||
)
|
|
||||||
raise Exception(
|
|
||||||
f"启动合成失败: {error_code} - {error_message}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 发送文本合成请求
|
|
||||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
|
||||||
run_message_id = str(uuid.uuid4().hex)
|
|
||||||
run_request = {
|
|
||||||
"header": {
|
|
||||||
"message_id": run_message_id,
|
|
||||||
"task_id": session_id,
|
|
||||||
"namespace": "FlowingSpeechSynthesizer",
|
|
||||||
"name": "RunSynthesis",
|
|
||||||
"appkey": self.appkey,
|
|
||||||
},
|
|
||||||
"payload": {"text": filtered_text},
|
|
||||||
}
|
|
||||||
await ws.send(json.dumps(run_request))
|
|
||||||
|
|
||||||
# 发送停止合成请求
|
|
||||||
stop_message_id = str(uuid.uuid4().hex)
|
|
||||||
stop_request = {
|
|
||||||
"header": {
|
|
||||||
"message_id": stop_message_id,
|
|
||||||
"task_id": session_id,
|
|
||||||
"namespace": "FlowingSpeechSynthesizer",
|
|
||||||
"name": "StopSynthesis",
|
|
||||||
"appkey": self.appkey,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
await ws.send(json.dumps(stop_request))
|
|
||||||
|
|
||||||
# 接收音频数据
|
|
||||||
synthesis_completed = False
|
|
||||||
while not synthesis_completed:
|
|
||||||
msg = await ws.recv()
|
|
||||||
if isinstance(msg, (bytes, bytearray)):
|
|
||||||
# 编码为Opus并收集
|
|
||||||
opus_frames = self.opus_encoder.encode_pcm_to_opus(
|
|
||||||
msg, False
|
|
||||||
)
|
|
||||||
audio_data.extend(opus_frames)
|
|
||||||
elif isinstance(msg, str):
|
|
||||||
data = json.loads(msg)
|
|
||||||
header = data.get("header", {})
|
|
||||||
event_name = header.get("name")
|
|
||||||
if event_name == "SynthesisCompleted":
|
|
||||||
synthesis_completed = True
|
|
||||||
logger.bind(tag=TAG).debug("TTS合成完成")
|
|
||||||
elif event_name == "TaskFailed":
|
|
||||||
error_info = data.get("payload", {}).get(
|
|
||||||
"error_info", {}
|
|
||||||
)
|
|
||||||
error_code = error_info.get("error_code")
|
|
||||||
error_message = error_info.get(
|
|
||||||
"error_message", "未知错误"
|
|
||||||
)
|
|
||||||
raise Exception(
|
|
||||||
f"合成失败: {error_code} - {error_message}"
|
|
||||||
)
|
|
||||||
finally:
|
|
||||||
try:
|
|
||||||
await ws.close()
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
|
|
||||||
loop.run_until_complete(_generate_audio())
|
|
||||||
loop.close()
|
|
||||||
|
|
||||||
return audio_data
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
|
|
||||||
return []
|
|
||||||
|
|||||||
@@ -4,12 +4,15 @@ import queue
|
|||||||
import uuid
|
import uuid
|
||||||
import asyncio
|
import asyncio
|
||||||
import threading
|
import threading
|
||||||
|
from typing import Callable, Any
|
||||||
from core.utils import p3
|
from core.utils import p3
|
||||||
|
import time
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from core.utils import textUtils
|
from core.utils import textUtils
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.util import audio_to_data, audio_bytes_to_data
|
from core.utils.audio_flow_control import FlowControlConfig
|
||||||
|
from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.utils.output_counter import add_device_output
|
from core.utils.output_counter import add_device_output
|
||||||
from core.handle.reportHandle import enqueue_tts_report
|
from core.handle.reportHandle import enqueue_tts_report
|
||||||
@@ -50,11 +53,9 @@ class TTSProviderBase(ABC):
|
|||||||
";",
|
";",
|
||||||
";",
|
";",
|
||||||
":",
|
":",
|
||||||
"~",
|
|
||||||
)
|
)
|
||||||
self.first_sentence_punctuations = (
|
self.first_sentence_punctuations = (
|
||||||
",",
|
",",
|
||||||
"~",
|
|
||||||
"~",
|
"~",
|
||||||
"、",
|
"、",
|
||||||
",",
|
",",
|
||||||
@@ -70,6 +71,7 @@ class TTSProviderBase(ABC):
|
|||||||
self.tts_stop_request = False
|
self.tts_stop_request = False
|
||||||
self.processed_chars = 0
|
self.processed_chars = 0
|
||||||
self.is_first_sentence = True
|
self.is_first_sentence = True
|
||||||
|
self.flow_controller = FlowControlConfig.create_flow_controller()
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(
|
return os.path.join(
|
||||||
@@ -77,7 +79,20 @@ class TTSProviderBase(ABC):
|
|||||||
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
)
|
)
|
||||||
|
|
||||||
def to_tts(self, text):
|
def handle_opus(self, opus_data: bytes):
|
||||||
|
logger.bind(tag=TAG).debug(
|
||||||
|
f"推送数据到队列里面帧数~~ {len(opus_data)}"
|
||||||
|
)
|
||||||
|
self.tts_audio_queue.put(
|
||||||
|
(SentenceType.MIDDLE, opus_data, None)
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle_audio_file(self, file_audio: bytes, text):
|
||||||
|
self.before_stop_play_files.append(
|
||||||
|
(file_audio, text)
|
||||||
|
)
|
||||||
|
|
||||||
|
def to_tts_stream(self, text, opus_handler: Callable[[bytes], None] = None) -> None:
|
||||||
text = MarkdownCleaner.clean_markdown(text)
|
text = MarkdownCleaner.clean_markdown(text)
|
||||||
max_repeat_time = 5
|
max_repeat_time = 5
|
||||||
if self.delete_audio_file:
|
if self.delete_audio_file:
|
||||||
@@ -86,10 +101,13 @@ class TTSProviderBase(ABC):
|
|||||||
try:
|
try:
|
||||||
audio_bytes = asyncio.run(self.text_to_speak(text, None))
|
audio_bytes = asyncio.run(self.text_to_speak(text, None))
|
||||||
if audio_bytes:
|
if audio_bytes:
|
||||||
audio_datas, _ = audio_bytes_to_data(
|
self.tts_audio_queue.put(
|
||||||
audio_bytes, file_type=self.audio_file_type, is_opus=True
|
(SentenceType.FIRST, None, text)
|
||||||
)
|
)
|
||||||
return audio_datas
|
audio_bytes_to_data_stream(
|
||||||
|
audio_bytes, file_type=self.audio_file_type, is_opus=True, callback=opus_handler
|
||||||
|
)
|
||||||
|
break
|
||||||
else:
|
else:
|
||||||
max_repeat_time -= 1
|
max_repeat_time -= 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -129,8 +147,10 @@ class TTSProviderBase(ABC):
|
|||||||
logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(
|
||||||
f"语音生成失败: {text},请检查网络或服务是否正常"
|
f"语音生成失败: {text},请检查网络或服务是否正常"
|
||||||
)
|
)
|
||||||
|
self.tts_audio_queue.put(
|
||||||
return tmp_file
|
(SentenceType.FIRST, None, text)
|
||||||
|
)
|
||||||
|
self._process_audio_file_stream(tmp_file, callback=opus_handler)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
|
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
|
||||||
return None
|
return None
|
||||||
@@ -139,13 +159,13 @@ class TTSProviderBase(ABC):
|
|||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def audio_to_pcm_data(self, audio_file_path):
|
def audio_to_pcm_data_stream(self, audio_file_path, callback: Callable[[Any], Any] = None):
|
||||||
"""音频文件转换为PCM编码"""
|
"""音频文件转换为PCM编码"""
|
||||||
return audio_to_data(audio_file_path, is_opus=False)
|
return audio_to_data_stream(audio_file_path, is_opus=False, callback=callback)
|
||||||
|
|
||||||
def audio_to_opus_data(self, audio_file_path):
|
def audio_to_opus_data_stream(self, audio_file_path, callback: Callable[[Any], Any] = None):
|
||||||
"""音频文件转换为Opus编码"""
|
"""音频文件转换为Opus编码"""
|
||||||
return audio_to_data(audio_file_path, is_opus=True)
|
return audio_to_data_stream(audio_file_path, is_opus=True, callback=callback)
|
||||||
|
|
||||||
def tts_one_sentence(
|
def tts_one_sentence(
|
||||||
self,
|
self,
|
||||||
@@ -208,34 +228,19 @@ class TTSProviderBase(ABC):
|
|||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.is_first_sentence = True
|
self.is_first_sentence = True
|
||||||
self.tts_audio_first_sentence = True
|
self.tts_audio_first_sentence = True
|
||||||
|
self.reset_flow_controller()
|
||||||
elif ContentType.TEXT == message.content_type:
|
elif ContentType.TEXT == message.content_type:
|
||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
segment_text = self._get_segment_text()
|
segment_text = self._get_segment_text()
|
||||||
if segment_text:
|
if segment_text:
|
||||||
if self.delete_audio_file:
|
self.to_tts_stream(segment_text, opus_handler=self.handle_opus)
|
||||||
audio_datas = self.to_tts(segment_text)
|
|
||||||
if audio_datas:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(message.sentence_type, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
tts_file = self.to_tts(segment_text)
|
|
||||||
if tts_file:
|
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(message.sentence_type, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
elif ContentType.FILE == message.content_type:
|
elif ContentType.FILE == message.content_type:
|
||||||
self._process_remaining_text()
|
self._process_remaining_text_stream(opus_handler=self.handle_opus)
|
||||||
tts_file = message.content_file
|
tts_file = message.content_file
|
||||||
if tts_file and os.path.exists(tts_file):
|
if tts_file and os.path.exists(tts_file):
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
self._process_audio_file_stream(tts_file, callback=self.handle_opus)
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(message.sentence_type, audio_datas, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
self._process_remaining_text()
|
self._process_remaining_text_stream(opus_handler=self.handle_opus)
|
||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(message.sentence_type, [], message.content_detail)
|
(message.sentence_type, [], message.content_detail)
|
||||||
)
|
)
|
||||||
@@ -249,30 +254,112 @@ class TTSProviderBase(ABC):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
def _audio_play_priority_thread(self):
|
def _audio_play_priority_thread(self):
|
||||||
|
# 需要上报的文本和音频列表
|
||||||
|
enqueue_text = None
|
||||||
|
enqueue_audio = None
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
text = None
|
text = None
|
||||||
try:
|
try:
|
||||||
try:
|
try:
|
||||||
sentence_type, audio_datas, text = self.tts_audio_queue.get(
|
sentence_type, audio_datas, text = self.tts_audio_queue.get(timeout=0.1)
|
||||||
timeout=1
|
|
||||||
)
|
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
if self.conn.stop_event.is_set():
|
if self.conn.stop_event.is_set():
|
||||||
break
|
break
|
||||||
continue
|
continue
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
|
||||||
sendAudioMessage(self.conn, sentence_type, audio_datas, text),
|
if self.conn.client_abort:
|
||||||
self.conn.loop,
|
logger.bind(tag=TAG).debug("收到打断信号,跳过当前音频数据")
|
||||||
)
|
# 打断时丢弃未上报的音频数据
|
||||||
future.result()
|
enqueue_text, enqueue_audio = None, []
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 收到下一个文本开始或会话结束时进行上报
|
||||||
|
if sentence_type is not SentenceType.MIDDLE:
|
||||||
|
# 上报TTS数据
|
||||||
|
if enqueue_text is not None and enqueue_audio is not None:
|
||||||
|
enqueue_tts_report(self.conn, enqueue_text, enqueue_audio)
|
||||||
|
enqueue_audio = []
|
||||||
|
enqueue_text = text
|
||||||
|
|
||||||
|
# 计算音频数据的帧数
|
||||||
|
if isinstance(audio_datas, bytes):
|
||||||
|
frame_count = 1 # 单个字节流作为一帧
|
||||||
|
enqueue_audio.append(audio_datas)
|
||||||
|
else:
|
||||||
|
frame_count = 0
|
||||||
|
|
||||||
|
# 记录输出和报告
|
||||||
if self.conn.max_output_size > 0 and text:
|
if self.conn.max_output_size > 0 and text:
|
||||||
add_device_output(self.conn.headers.get("device-id"), len(text))
|
add_device_output(self.conn.headers.get("device-id"), len(text))
|
||||||
enqueue_tts_report(self.conn, text, audio_datas)
|
|
||||||
|
# 流控检查
|
||||||
|
if frame_count > 0:
|
||||||
|
max_wait_time = FlowControlConfig.DEFAULT_MAX_WAIT_TIME
|
||||||
|
wait_start_time = time.time()
|
||||||
|
retry_interval = FlowControlConfig.DEFAULT_RETRY_INTERVAL
|
||||||
|
|
||||||
|
while not self.flow_controller.can_send_frames(frame_count):
|
||||||
|
# 检查是否超时或需要停止
|
||||||
|
if (time.time() - wait_start_time > max_wait_time or
|
||||||
|
self.conn.stop_event.is_set() or
|
||||||
|
self.conn.client_abort):
|
||||||
|
logger.bind(tag=TAG).debug("流控等待超时或收到停止信号,跳过音频发送")
|
||||||
|
break
|
||||||
|
# 短暂等待后重试
|
||||||
|
time.sleep(retry_interval)
|
||||||
|
else:
|
||||||
|
# 可以发送,记录发送的帧数
|
||||||
|
self.flow_controller.record_sent_frames(frame_count)
|
||||||
|
|
||||||
|
# 发送音频
|
||||||
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
|
self._send_audio_with_flow_control(sentence_type, audio_datas, text),
|
||||||
|
self.conn.loop,
|
||||||
|
)
|
||||||
|
future.result()
|
||||||
|
|
||||||
|
# 输出流控状态(调试用)
|
||||||
|
# status = self.flow_controller.get_status()
|
||||||
|
# logger.bind(tag=TAG).debug(
|
||||||
|
# f"流控状态: 缓冲区使用率={status['buffer_usage_percent']:.1f}%, "
|
||||||
|
# f"可用令牌={status['available_tokens']}..."
|
||||||
|
# f"发送帧数={status['sent_frames']}..."
|
||||||
|
# f"消费帧数={status['consumed_frames']}..."
|
||||||
|
# f"代播放帧数={status['sent_frames'] - status['consumed_frames']}..."
|
||||||
|
# )
|
||||||
|
else:
|
||||||
|
# 没有音频数据,直接发送
|
||||||
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
|
self._send_audio_with_flow_control(sentence_type, audio_datas, text),
|
||||||
|
self.conn.loop,
|
||||||
|
)
|
||||||
|
future.result()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(
|
||||||
f"audio_play_priority priority_thread: {text} {e}"
|
f"audio_play_priority_thread: {text} {e}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def _send_audio_with_flow_control(self, sentence_type, audio_datas, text):
|
||||||
|
"""
|
||||||
|
带流控的音频发送方法 模拟设备消费音频帧的过程
|
||||||
|
实际应用中应该根据设备反馈来更新消费情况
|
||||||
|
"""
|
||||||
|
await sendAudioMessage(self.conn, sentence_type, audio_datas, text)
|
||||||
|
|
||||||
|
# 模拟设备消费(实际应用中应该从设备获取反馈)防止音字不同步
|
||||||
|
if isinstance(audio_datas, bytes):
|
||||||
|
# 模拟设备播放延迟(60ms per frame), 实际情况可以低一点(50ms),增加使用体验
|
||||||
|
await asyncio.sleep(0.055)
|
||||||
|
self.flow_controller.update_device_consumption(1)
|
||||||
|
|
||||||
|
# 在类中添加流控制器重置方法
|
||||||
|
def reset_flow_controller(self):
|
||||||
|
"""重置流控制器状态,通常在新会话开始时调用"""
|
||||||
|
if hasattr(self, 'flow_controller'):
|
||||||
|
self.flow_controller.reset()
|
||||||
|
logger.bind(tag=TAG).info("流控制器状态已重置")
|
||||||
|
|
||||||
async def start_session(self, session_id):
|
async def start_session(self, session_id):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@@ -323,22 +410,19 @@ class TTSProviderBase(ABC):
|
|||||||
else:
|
else:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _process_audio_file(self, tts_file):
|
def _process_audio_file_stream(self, tts_file, callback: Callable[[Any], Any]) -> None:
|
||||||
"""处理音频文件并转换为指定格式
|
"""处理音频文件并转换为指定格式
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
tts_file: 音频文件路径
|
tts_file: 音频文件路径
|
||||||
content_detail: 内容详情
|
callback: 文件处理函数
|
||||||
|
|
||||||
Returns:
|
|
||||||
tuple: (sentence_type, audio_datas, content_detail)
|
|
||||||
"""
|
"""
|
||||||
if tts_file.endswith(".p3"):
|
if tts_file.endswith(".p3"):
|
||||||
audio_datas, _ = p3.decode_opus_from_file(tts_file)
|
p3.decode_opus_from_file_stream(tts_file, callback=callback)
|
||||||
elif self.conn.audio_format == "pcm":
|
elif self.conn.audio_format == "pcm":
|
||||||
audio_datas, _ = self.audio_to_pcm_data(tts_file)
|
self.audio_to_pcm_data_stream(tts_file, callback=callback)
|
||||||
else:
|
else:
|
||||||
audio_datas, _ = self.audio_to_opus_data(tts_file)
|
self.audio_to_opus_data_stream(tts_file, callback=callback)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.delete_audio_file
|
self.delete_audio_file
|
||||||
@@ -347,7 +431,6 @@ class TTSProviderBase(ABC):
|
|||||||
and tts_file.startswith(self.output_file)
|
and tts_file.startswith(self.output_file)
|
||||||
):
|
):
|
||||||
os.remove(tts_file)
|
os.remove(tts_file)
|
||||||
return audio_datas
|
|
||||||
|
|
||||||
def _process_before_stop_play_files(self):
|
def _process_before_stop_play_files(self):
|
||||||
for audio_datas, text in self.before_stop_play_files:
|
for audio_datas, text in self.before_stop_play_files:
|
||||||
@@ -355,7 +438,7 @@ class TTSProviderBase(ABC):
|
|||||||
self.before_stop_play_files.clear()
|
self.before_stop_play_files.clear()
|
||||||
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
|
|
||||||
def _process_remaining_text(self):
|
def _process_remaining_text_stream(self, opus_handler: Callable[[bytes], None] = None):
|
||||||
"""处理剩余的文本并生成语音
|
"""处理剩余的文本并生成语音
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -366,18 +449,7 @@ class TTSProviderBase(ABC):
|
|||||||
if remaining_text:
|
if remaining_text:
|
||||||
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
||||||
if segment_text:
|
if segment_text:
|
||||||
if self.delete_audio_file:
|
self.to_tts_stream(segment_text, opus_handler=opus_handler)
|
||||||
audio_datas = self.to_tts(segment_text)
|
|
||||||
if audio_datas:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
tts_file = self.to_tts(segment_text)
|
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
self.processed_chars += len(full_text)
|
self.processed_chars += len(full_text)
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import json
|
|||||||
import queue
|
import queue
|
||||||
import asyncio
|
import asyncio
|
||||||
import traceback
|
import traceback
|
||||||
|
from typing import Callable, Any
|
||||||
import websockets
|
import websockets
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
@@ -212,6 +213,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
if message.sentence_type == SentenceType.FIRST:
|
if message.sentence_type == SentenceType.FIRST:
|
||||||
self.conn.client_abort = False
|
self.conn.client_abort = False
|
||||||
|
self.reset_flow_controller()
|
||||||
|
|
||||||
if self.conn.client_abort:
|
if self.conn.client_abort:
|
||||||
try:
|
try:
|
||||||
@@ -266,11 +268,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
||||||
@@ -428,9 +426,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
async def _start_monitor_tts_response(self):
|
async def _start_monitor_tts_response(self):
|
||||||
"""监听TTS响应"""
|
"""监听TTS响应"""
|
||||||
opus_datas_cache = []
|
|
||||||
is_first_sentence = True
|
|
||||||
first_sentence_segment_count = 0 # 添加计数器
|
|
||||||
try:
|
try:
|
||||||
session_finished = False # 标记会话是否正常结束
|
session_finished = False # 标记会话是否正常结束
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
@@ -451,37 +446,14 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(SentenceType.FIRST, [], self.tts_text)
|
(SentenceType.FIRST, [], self.tts_text)
|
||||||
)
|
)
|
||||||
opus_datas_cache = []
|
|
||||||
first_sentence_segment_count = 0 # 重置计数器
|
|
||||||
elif (
|
elif (
|
||||||
res.optional.event == EVENT_TTSResponse
|
res.optional.event == EVENT_TTSResponse
|
||||||
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
||||||
):
|
):
|
||||||
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
self.wav_to_opus_data_audio_raw_stream(res.payload, callback=self.handle_opus)
|
||||||
logger.bind(tag=TAG).debug(
|
|
||||||
f"推送数据到队列里面帧数~~{len(opus_datas)}"
|
|
||||||
)
|
|
||||||
if is_first_sentence:
|
|
||||||
first_sentence_segment_count += 1
|
|
||||||
if first_sentence_segment_count <= 6:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas, None)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
opus_datas_cache.extend(opus_datas)
|
|
||||||
else:
|
|
||||||
# 后续句子缓存
|
|
||||||
opus_datas_cache.extend(opus_datas)
|
|
||||||
elif res.optional.event == EVENT_TTSSentenceEnd:
|
elif res.optional.event == EVENT_TTSSentenceEnd:
|
||||||
logger.bind(tag=TAG).info(f"句子语音生成成功:{self.tts_text}")
|
logger.bind(tag=TAG).info(f"句子语音生成成功:{self.tts_text}")
|
||||||
if not is_first_sentence or first_sentence_segment_count > 10:
|
|
||||||
# 发送缓存的数据
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
# 第一句话结束后,将标志设置为False
|
|
||||||
is_first_sentence = False
|
|
||||||
elif res.optional.event == EVENT_SessionFinished:
|
elif res.optional.event == EVENT_SessionFinished:
|
||||||
logger.bind(tag=TAG).debug(f"会话结束~~")
|
logger.bind(tag=TAG).debug(f"会话结束~~")
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -655,110 +627,5 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False):
|
def wav_to_opus_data_audio_raw_stream(self, raw_data_var, is_end=False, callback: Callable[[Any], Any]=None):
|
||||||
opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end)
|
return self.opus_encoder.encode_pcm_to_opus_stream(raw_data_var, is_end, callback=callback)
|
||||||
return opus_datas
|
|
||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
|
||||||
"""非流式生成音频数据,用于生成音频及测试场景
|
|
||||||
|
|
||||||
Args:
|
|
||||||
text: 要转换的文本
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
list: 音频数据列表
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
# 创建事件循环
|
|
||||||
loop = asyncio.new_event_loop()
|
|
||||||
asyncio.set_event_loop(loop)
|
|
||||||
|
|
||||||
# 生成会话ID
|
|
||||||
session_id = uuid.uuid4().__str__().replace("-", "")
|
|
||||||
|
|
||||||
# 存储音频数据
|
|
||||||
audio_data = []
|
|
||||||
|
|
||||||
async def _generate_audio():
|
|
||||||
# 创建新的WebSocket连接
|
|
||||||
ws_header = {
|
|
||||||
"X-Api-App-Key": self.appId,
|
|
||||||
"X-Api-Access-Key": self.access_token,
|
|
||||||
"X-Api-Resource-Id": self.resource_id,
|
|
||||||
"X-Api-Connect-Id": uuid.uuid4(),
|
|
||||||
}
|
|
||||||
ws = await websockets.connect(
|
|
||||||
self.ws_url, additional_headers=ws_header, max_size=1000000000
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
# 启动会话
|
|
||||||
header = Header(
|
|
||||||
message_type=FULL_CLIENT_REQUEST,
|
|
||||||
message_type_specific_flags=MsgTypeFlagWithEvent,
|
|
||||||
serial_method=JSON,
|
|
||||||
).as_bytes()
|
|
||||||
optional = Optional(
|
|
||||||
event=EVENT_StartSession, sessionId=session_id
|
|
||||||
).as_bytes()
|
|
||||||
payload = self.get_payload_bytes(
|
|
||||||
event=EVENT_StartSession, speaker=self.voice
|
|
||||||
)
|
|
||||||
await self.send_event(ws, header, optional, payload)
|
|
||||||
|
|
||||||
# 发送文本
|
|
||||||
header = Header(
|
|
||||||
message_type=FULL_CLIENT_REQUEST,
|
|
||||||
message_type_specific_flags=MsgTypeFlagWithEvent,
|
|
||||||
serial_method=JSON,
|
|
||||||
).as_bytes()
|
|
||||||
optional = Optional(
|
|
||||||
event=EVENT_TaskRequest, sessionId=session_id
|
|
||||||
).as_bytes()
|
|
||||||
payload = self.get_payload_bytes(
|
|
||||||
event=EVENT_TaskRequest, text=text, speaker=self.voice
|
|
||||||
)
|
|
||||||
await self.send_event(ws, header, optional, payload)
|
|
||||||
|
|
||||||
# 发送结束会话请求
|
|
||||||
header = Header(
|
|
||||||
message_type=FULL_CLIENT_REQUEST,
|
|
||||||
message_type_specific_flags=MsgTypeFlagWithEvent,
|
|
||||||
serial_method=JSON,
|
|
||||||
).as_bytes()
|
|
||||||
optional = Optional(
|
|
||||||
event=EVENT_FinishSession, sessionId=session_id
|
|
||||||
).as_bytes()
|
|
||||||
payload = str.encode("{}")
|
|
||||||
await self.send_event(ws, header, optional, payload)
|
|
||||||
|
|
||||||
# 接收音频数据
|
|
||||||
while True:
|
|
||||||
msg = await ws.recv()
|
|
||||||
res = self.parser_response(msg)
|
|
||||||
|
|
||||||
if (
|
|
||||||
res.optional.event == EVENT_TTSResponse
|
|
||||||
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
|
||||||
):
|
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
|
||||||
audio_data.extend(opus_datas)
|
|
||||||
elif res.optional.event == EVENT_SessionFinished:
|
|
||||||
break
|
|
||||||
|
|
||||||
finally:
|
|
||||||
# 清理资源
|
|
||||||
try:
|
|
||||||
await ws.close()
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
|
|
||||||
# 运行异步任务
|
|
||||||
loop.run_until_complete(_generate_audio())
|
|
||||||
loop.close()
|
|
||||||
|
|
||||||
return audio_data
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
|
|
||||||
return []
|
|
||||||
|
|||||||
@@ -3,8 +3,6 @@ import queue
|
|||||||
import asyncio
|
import asyncio
|
||||||
import traceback
|
import traceback
|
||||||
import aiohttp
|
import aiohttp
|
||||||
import requests
|
|
||||||
import time
|
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
@@ -27,15 +25,13 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url = config.get("api_url", "http://8.138.114.124:11996/tts")
|
self.api_url = config.get("api_url", "http://8.138.114.124:11996/tts")
|
||||||
self.audio_format = "pcm"
|
self.audio_format = "pcm"
|
||||||
self.before_stop_play_files = []
|
self.before_stop_play_files = []
|
||||||
self.segment_count = 0
|
|
||||||
|
|
||||||
# 创建Opus编码器 需注意接口返回的采样率为24000
|
# 创建Opus编码器 需注意接口返回的采样率为24000
|
||||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
sample_rate=24000, channels=1, frame_size_ms=60
|
sample_rate=24000, channels=1, frame_size_ms=60
|
||||||
)
|
)
|
||||||
|
|
||||||
# 文本缓冲区和PCM缓冲区
|
# PCM缓冲区
|
||||||
self.text_buffer = ""
|
|
||||||
self.pcm_buffer = bytearray()
|
self.pcm_buffer = bytearray()
|
||||||
|
|
||||||
def tts_text_priority_thread(self):
|
def tts_text_priority_thread(self):
|
||||||
@@ -48,8 +44,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_stop_request = False
|
self.tts_stop_request = False
|
||||||
self.processed_chars = 0
|
self.processed_chars = 0
|
||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.segment_count = 0
|
|
||||||
self.before_stop_play_files.clear()
|
self.before_stop_play_files.clear()
|
||||||
|
self.reset_flow_controller()
|
||||||
elif ContentType.TEXT == message.content_type:
|
elif ContentType.TEXT == message.content_type:
|
||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
segment_text = self._get_segment_text()
|
segment_text = self._get_segment_text()
|
||||||
@@ -62,14 +58,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
# 处理剩余的文本
|
# 处理剩余的文本
|
||||||
self._process_remaining_text(True)
|
self._process_remaining_text_stream(True)
|
||||||
|
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
continue
|
continue
|
||||||
@@ -78,7 +71,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _process_remaining_text(self, is_last=False):
|
def _process_remaining_text_stream(self, is_last=False):
|
||||||
"""处理剩余的文本并生成语音
|
"""处理剩余的文本并生成语音
|
||||||
Returns:
|
Returns:
|
||||||
bool: 是否成功处理了文本
|
bool: 是否成功处理了文本
|
||||||
@@ -143,8 +136,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
return
|
return
|
||||||
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
opus_datas_cache = []
|
|
||||||
|
|
||||||
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
|
|
||||||
# 处理音频流数据
|
# 处理音频流数据
|
||||||
@@ -158,41 +149,22 @@ class TTSProvider(TTSProviderBase):
|
|||||||
while len(self.pcm_buffer) >= frame_bytes:
|
while len(self.pcm_buffer) >= frame_bytes:
|
||||||
frame = bytes(self.pcm_buffer[:frame_bytes])
|
frame = bytes(self.pcm_buffer[:frame_bytes])
|
||||||
del self.pcm_buffer[:frame_bytes]
|
del self.pcm_buffer[:frame_bytes]
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
|
||||||
frame, end_of_stream=False
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
|
frame,
|
||||||
|
end_of_stream=False,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
|
|
||||||
# flush 剩余不足一帧的数据
|
# flush 剩余不足一帧的数据
|
||||||
if self.pcm_buffer:
|
if self.pcm_buffer:
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
bytes(self.pcm_buffer), end_of_stream=True
|
bytes(self.pcm_buffer),
|
||||||
|
end_of_stream=True,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
# 直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
# 后续片段缓存
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
|
|
||||||
# 如果不是前10个片段,发送缓存的数据
|
|
||||||
if self.segment_count >= 10 and opus_datas_cache:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
|
|
||||||
# 如果是最后一段,输出音频获取完毕
|
# 如果是最后一段,输出音频获取完毕
|
||||||
if is_last:
|
if is_last:
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -206,59 +178,3 @@ class TTSProvider(TTSProviderBase):
|
|||||||
await super().close()
|
await super().close()
|
||||||
if hasattr(self, "opus_encoder"):
|
if hasattr(self, "opus_encoder"):
|
||||||
self.opus_encoder.close()
|
self.opus_encoder.close()
|
||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
|
||||||
"""非流式TTS处理,用于测试及保存音频文件的场景
|
|
||||||
|
|
||||||
Args:
|
|
||||||
text: 要转换的文本
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
list: 返回opus编码后的音频数据列表
|
|
||||||
"""
|
|
||||||
start_time = time.time()
|
|
||||||
text = MarkdownCleaner.clean_markdown(text)
|
|
||||||
|
|
||||||
payload = {"text": text, "character": self.character}
|
|
||||||
|
|
||||||
try:
|
|
||||||
with requests.post(self.api_url, json=payload, timeout=5) as response:
|
|
||||||
if response.status_code != 200:
|
|
||||||
logger.bind(tag=TAG).error(
|
|
||||||
f"TTS请求失败: {response.status_code}, {response.text}"
|
|
||||||
)
|
|
||||||
return []
|
|
||||||
|
|
||||||
logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒")
|
|
||||||
|
|
||||||
# 使用opus编码器处理PCM数据
|
|
||||||
opus_datas = []
|
|
||||||
pcm_data = response.content
|
|
||||||
|
|
||||||
# 计算每帧的字节数
|
|
||||||
frame_bytes = int(
|
|
||||||
self.opus_encoder.sample_rate
|
|
||||||
* self.opus_encoder.channels
|
|
||||||
* self.opus_encoder.frame_size_ms
|
|
||||||
/ 1000
|
|
||||||
* 2
|
|
||||||
)
|
|
||||||
|
|
||||||
# 分帧处理PCM数据
|
|
||||||
for i in range(0, len(pcm_data), frame_bytes):
|
|
||||||
frame = pcm_data[i : i + frame_bytes]
|
|
||||||
if len(frame) < frame_bytes:
|
|
||||||
# 最后一帧可能不足,用0填充
|
|
||||||
frame = frame + b"\x00" * (frame_bytes - len(frame))
|
|
||||||
|
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
|
||||||
frame, end_of_stream=(i + frame_bytes >= len(pcm_data))
|
|
||||||
)
|
|
||||||
if opus:
|
|
||||||
opus_datas.extend(opus)
|
|
||||||
|
|
||||||
return opus_datas
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
|
||||||
return []
|
|
||||||
|
|||||||
@@ -3,8 +3,6 @@ import queue
|
|||||||
import asyncio
|
import asyncio
|
||||||
import traceback
|
import traceback
|
||||||
import aiohttp
|
import aiohttp
|
||||||
import requests
|
|
||||||
import time
|
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
@@ -24,23 +22,15 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url = config.get("api_url")
|
self.api_url = config.get("api_url")
|
||||||
self.audio_format = "pcm"
|
self.audio_format = "pcm"
|
||||||
self.before_stop_play_files = []
|
self.before_stop_play_files = []
|
||||||
self.segment_count = 0 # 添加片段计数器
|
|
||||||
|
|
||||||
# 创建Opus编码器
|
# 创建Opus编码器
|
||||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
sample_rate=16000, channels=1, frame_size_ms=60
|
sample_rate=16000, channels=1, frame_size_ms=60
|
||||||
)
|
)
|
||||||
|
|
||||||
# 添加文本缓冲区
|
|
||||||
self.text_buffer = ""
|
|
||||||
|
|
||||||
# PCM缓冲区
|
# PCM缓冲区
|
||||||
self.pcm_buffer = bytearray()
|
self.pcm_buffer = bytearray()
|
||||||
|
|
||||||
###################################################################################
|
|
||||||
# linkerai单流式TTS重写父类的方法--开始
|
|
||||||
###################################################################################
|
|
||||||
|
|
||||||
def tts_text_priority_thread(self):
|
def tts_text_priority_thread(self):
|
||||||
"""流式文本处理线程"""
|
"""流式文本处理线程"""
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
@@ -51,8 +41,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_stop_request = False
|
self.tts_stop_request = False
|
||||||
self.processed_chars = 0
|
self.processed_chars = 0
|
||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.segment_count = 0
|
|
||||||
self.before_stop_play_files.clear()
|
self.before_stop_play_files.clear()
|
||||||
|
self.reset_flow_controller()
|
||||||
elif ContentType.TEXT == message.content_type:
|
elif ContentType.TEXT == message.content_type:
|
||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
segment_text = self._get_segment_text()
|
segment_text = self._get_segment_text()
|
||||||
@@ -65,14 +55,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
# 处理剩余的文本
|
# 处理剩余的文本
|
||||||
self._process_remaining_text(True)
|
self._process_remaining_text_stream(True)
|
||||||
|
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
continue
|
continue
|
||||||
@@ -81,7 +67,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _process_remaining_text(self, is_last=False):
|
def _process_remaining_text_stream(self, is_last=False):
|
||||||
"""处理剩余的文本并生成语音
|
"""处理剩余的文本并生成语音
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -176,8 +162,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
return
|
return
|
||||||
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
opus_datas_cache = []
|
|
||||||
|
|
||||||
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
|
|
||||||
# 兼容 iter_chunked / iter_chunks / iter_any
|
# 兼容 iter_chunked / iter_chunks / iter_any
|
||||||
@@ -194,41 +178,21 @@ class TTSProvider(TTSProviderBase):
|
|||||||
frame = bytes(self.pcm_buffer[:frame_bytes])
|
frame = bytes(self.pcm_buffer[:frame_bytes])
|
||||||
del self.pcm_buffer[:frame_bytes]
|
del self.pcm_buffer[:frame_bytes]
|
||||||
|
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
frame, end_of_stream=False
|
frame,
|
||||||
|
end_of_stream=False,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
|
|
||||||
# flush 剩余不足一帧的数据
|
# flush 剩余不足一帧的数据
|
||||||
if self.pcm_buffer:
|
if self.pcm_buffer:
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
bytes(self.pcm_buffer), end_of_stream=True
|
bytes(self.pcm_buffer),
|
||||||
|
end_of_stream=True,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
# 直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
# 后续片段缓存
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
|
|
||||||
# 如果不是前10个片段,发送缓存的数据
|
|
||||||
if self.segment_count >= 10 and opus_datas_cache:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
|
|
||||||
# 如果是最后一段,输出音频获取完毕
|
# 如果是最后一段,输出音频获取完毕
|
||||||
if is_last:
|
if is_last:
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -236,73 +200,3 @@ class TTSProvider(TTSProviderBase):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
||||||
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
|
||||||
"""非流式TTS处理,用于测试及保存音频文件的场景
|
|
||||||
|
|
||||||
Args:
|
|
||||||
text: 要转换的文本
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
list: 返回opus编码后的音频数据列表
|
|
||||||
"""
|
|
||||||
start_time = time.time()
|
|
||||||
text = MarkdownCleaner.clean_markdown(text)
|
|
||||||
|
|
||||||
params = {
|
|
||||||
"tts_text": text,
|
|
||||||
"spk_id": self.voice,
|
|
||||||
"frame_duration": 60,
|
|
||||||
"stream": False,
|
|
||||||
"target_sr": 16000,
|
|
||||||
"audio_format": self.audio_format,
|
|
||||||
"instruct_text": "请生成一段自然流畅的语音",
|
|
||||||
}
|
|
||||||
headers = {
|
|
||||||
"Authorization": f"Bearer {self.access_token}",
|
|
||||||
"Content-Type": "application/json",
|
|
||||||
}
|
|
||||||
|
|
||||||
try:
|
|
||||||
with requests.get(
|
|
||||||
self.api_url, params=params, headers=headers, timeout=5
|
|
||||||
) as response:
|
|
||||||
if response.status_code != 200:
|
|
||||||
logger.bind(tag=TAG).error(
|
|
||||||
f"TTS请求失败: {response.status_code}, {response.text}"
|
|
||||||
)
|
|
||||||
return []
|
|
||||||
|
|
||||||
logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒")
|
|
||||||
|
|
||||||
# 使用opus编码器处理PCM数据
|
|
||||||
opus_datas = []
|
|
||||||
pcm_data = response.content
|
|
||||||
|
|
||||||
# 计算每帧的字节数
|
|
||||||
frame_bytes = int(
|
|
||||||
self.opus_encoder.sample_rate
|
|
||||||
* self.opus_encoder.channels
|
|
||||||
* self.opus_encoder.frame_size_ms
|
|
||||||
/ 1000
|
|
||||||
* 2
|
|
||||||
)
|
|
||||||
|
|
||||||
# 分帧处理PCM数据
|
|
||||||
for i in range(0, len(pcm_data), frame_bytes):
|
|
||||||
frame = pcm_data[i : i + frame_bytes]
|
|
||||||
if len(frame) < frame_bytes:
|
|
||||||
# 最后一帧可能不足,用0填充
|
|
||||||
frame = frame + b"\x00" * (frame_bytes - len(frame))
|
|
||||||
|
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
|
||||||
frame, end_of_stream=(i + frame_bytes >= len(pcm_data))
|
|
||||||
)
|
|
||||||
if opus:
|
|
||||||
opus_datas.extend(opus)
|
|
||||||
|
|
||||||
return opus_datas
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
|
||||||
return []
|
|
||||||
|
|||||||
@@ -0,0 +1,186 @@
|
|||||||
|
"""
|
||||||
|
音频流控模块
|
||||||
|
包含令牌桶算法和音频流控制器的实现
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import time
|
||||||
|
import threading
|
||||||
|
from collections import deque
|
||||||
|
from typing import Optional, Dict, Any
|
||||||
|
|
||||||
|
|
||||||
|
class TokenBucket:
|
||||||
|
"""令牌桶实现,用于限流控制"""
|
||||||
|
|
||||||
|
def __init__(self, capacity: int, refill_rate: float, initial_tokens: Optional[int] = None):
|
||||||
|
"""
|
||||||
|
初始化令牌桶
|
||||||
|
|
||||||
|
Args:
|
||||||
|
capacity: 桶容量(最大令牌数)
|
||||||
|
refill_rate: 令牌补充速率(每秒补充的令牌数)
|
||||||
|
initial_tokens: 初始令牌数,默认为桶容量
|
||||||
|
"""
|
||||||
|
self.capacity = capacity
|
||||||
|
self.refill_rate = refill_rate
|
||||||
|
self.tokens = initial_tokens if initial_tokens is not None else capacity
|
||||||
|
self.last_refill_time = time.time()
|
||||||
|
self.lock = threading.Lock()
|
||||||
|
|
||||||
|
def get_tokens(self, requested_tokens: int = 1) -> bool:
|
||||||
|
"""
|
||||||
|
获取指定数量的令牌
|
||||||
|
|
||||||
|
Args:
|
||||||
|
requested_tokens: 请求的令牌数量
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: 是否成功获取到令牌
|
||||||
|
"""
|
||||||
|
with self.lock:
|
||||||
|
self._refill_tokens()
|
||||||
|
|
||||||
|
if self.tokens >= requested_tokens:
|
||||||
|
self.tokens -= requested_tokens
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def get_available_tokens(self) -> int:
|
||||||
|
"""获取当前可用令牌数"""
|
||||||
|
with self.lock:
|
||||||
|
self._refill_tokens()
|
||||||
|
return int(self.tokens)
|
||||||
|
|
||||||
|
def _refill_tokens(self):
|
||||||
|
"""内部方法:补充令牌"""
|
||||||
|
current_time = time.time()
|
||||||
|
time_passed = current_time - self.last_refill_time
|
||||||
|
tokens_to_add = time_passed * self.refill_rate
|
||||||
|
|
||||||
|
self.tokens = min(self.capacity, self.tokens + tokens_to_add)
|
||||||
|
self.last_refill_time = current_time
|
||||||
|
|
||||||
|
|
||||||
|
class AudioFlowController:
|
||||||
|
"""音频流控制器,基于令牌桶算法控制音频数据发送"""
|
||||||
|
|
||||||
|
def __init__(self, max_device_buffer: int = 3000, refill_rate: float = 20):
|
||||||
|
"""
|
||||||
|
初始化音频流控制器
|
||||||
|
|
||||||
|
Args:
|
||||||
|
max_device_buffer: 设备端最大缓冲区大小(Opus帧数)
|
||||||
|
refill_rate: 令牌补充速率(每秒允许发送的帧数)
|
||||||
|
"""
|
||||||
|
self.max_device_buffer = max_device_buffer
|
||||||
|
self.token_bucket = TokenBucket(
|
||||||
|
capacity=max_device_buffer,
|
||||||
|
refill_rate=refill_rate,
|
||||||
|
initial_tokens=max_device_buffer // 2 # 初始令牌为容量的一半
|
||||||
|
)
|
||||||
|
self.sent_frames_count = 0 # 已发送帧数计数
|
||||||
|
self.device_consumed_frames = 0 # 设备端已消费帧数
|
||||||
|
self.pending_queue = deque() # 等待发送的数据队列
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def can_send_frames(self, frame_count: int) -> bool:
|
||||||
|
"""
|
||||||
|
检查是否可以发送指定数量的帧
|
||||||
|
|
||||||
|
Args:
|
||||||
|
frame_count: 要发送的帧数
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: 是否可以发送
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
# 检查设备端缓冲区是否会溢出
|
||||||
|
estimated_device_buffer = self.sent_frames_count - self.device_consumed_frames
|
||||||
|
if estimated_device_buffer + frame_count > self.max_device_buffer:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 检查令牌桶是否有足够令牌
|
||||||
|
return self.token_bucket.get_tokens(frame_count)
|
||||||
|
|
||||||
|
def update_device_consumption(self, consumed_frames: int):
|
||||||
|
"""
|
||||||
|
更新设备端消费的帧数
|
||||||
|
|
||||||
|
Args:
|
||||||
|
consumed_frames: 设备端消费的帧数
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
self.device_consumed_frames += consumed_frames
|
||||||
|
|
||||||
|
def record_sent_frames(self, frame_count: int):
|
||||||
|
"""
|
||||||
|
记录已发送的帧数
|
||||||
|
|
||||||
|
Args:
|
||||||
|
frame_count: 发送的帧数
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
self.sent_frames_count += frame_count
|
||||||
|
|
||||||
|
def get_status(self) -> Dict[str, Any]:
|
||||||
|
"""获取流控状态信息"""
|
||||||
|
with self._lock:
|
||||||
|
estimated_buffer = self.sent_frames_count - self.device_consumed_frames
|
||||||
|
return {
|
||||||
|
"sent_frames": self.sent_frames_count,
|
||||||
|
"consumed_frames": self.device_consumed_frames,
|
||||||
|
"estimated_device_buffer": estimated_buffer,
|
||||||
|
"available_tokens": self.token_bucket.get_available_tokens(),
|
||||||
|
"pending_queue_size": len(self.pending_queue),
|
||||||
|
"buffer_usage_percent": (estimated_buffer / self.max_device_buffer) * 100
|
||||||
|
}
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
"""重置流控状态"""
|
||||||
|
with self._lock:
|
||||||
|
self.sent_frames_count = 0
|
||||||
|
self.device_consumed_frames = 0
|
||||||
|
self.pending_queue.clear()
|
||||||
|
# 重新初始化令牌桶
|
||||||
|
self.token_bucket = TokenBucket(
|
||||||
|
capacity=self.max_device_buffer,
|
||||||
|
refill_rate=self.token_bucket.refill_rate,
|
||||||
|
initial_tokens=self.max_device_buffer // 2
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# 流控配置常量
|
||||||
|
class FlowControlConfig:
|
||||||
|
"""流控配置常量"""
|
||||||
|
# Opus 编码参数
|
||||||
|
OPUS_FRAME_DURATION_MS = 60 # Opus帧时长(毫秒)
|
||||||
|
OPUS_FRAMES_PER_SECOND = 1000 / OPUS_FRAME_DURATION_MS # 每秒帧数
|
||||||
|
|
||||||
|
# 默认流控参数
|
||||||
|
DEFAULT_MAX_DEVICE_BUFFER = 40 # 设备端最大缓冲帧数
|
||||||
|
DEFAULT_REFILL_RATE = OPUS_FRAMES_PER_SECOND # 默认令牌补充速率(帧/秒)
|
||||||
|
DEFAULT_MAX_WAIT_TIME = 5.0 # 流控最大等待时间(秒)
|
||||||
|
DEFAULT_RETRY_INTERVAL = 0.06 # 流控重试间隔(秒)
|
||||||
|
|
||||||
|
# 预缓冲参数
|
||||||
|
PRE_BUFFER_FRAMES = 3 # 预缓冲帧数
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create_flow_controller(cls, max_buffer: Optional[int] = None,
|
||||||
|
refill_rate: Optional[float] = None) -> AudioFlowController:
|
||||||
|
"""
|
||||||
|
创建流控制器的工厂方法
|
||||||
|
|
||||||
|
Args:
|
||||||
|
max_buffer: 最大缓冲区大小,使用默认值如果为None
|
||||||
|
refill_rate: 令牌补充速率,使用默认值如果为None
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
AudioFlowController: 配置好的流控制器实例
|
||||||
|
"""
|
||||||
|
return AudioFlowController(
|
||||||
|
max_device_buffer=max_buffer or cls.DEFAULT_MAX_DEVICE_BUFFER,
|
||||||
|
refill_rate=refill_rate or cls.DEFAULT_REFILL_RATE
|
||||||
|
)
|
||||||
@@ -5,9 +5,8 @@ Opus编码工具类
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import traceback
|
import traceback
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from typing import List, Optional
|
from typing import Optional, Callable, Any
|
||||||
from opuslib_next import Encoder
|
from opuslib_next import Encoder
|
||||||
from opuslib_next import constants
|
from opuslib_next import constants
|
||||||
|
|
||||||
@@ -56,13 +55,14 @@ class OpusEncoderUtils:
|
|||||||
self.encoder.reset_state()
|
self.encoder.reset_state()
|
||||||
self.buffer = np.array([], dtype=np.int16)
|
self.buffer = np.array([], dtype=np.int16)
|
||||||
|
|
||||||
def encode_pcm_to_opus(self, pcm_data: bytes, end_of_stream: bool) -> List[bytes]:
|
def encode_pcm_to_opus_stream(self, pcm_data: bytes, end_of_stream: bool, callback: Callable[[Any], Any]):
|
||||||
"""
|
"""
|
||||||
将PCM数据编码为Opus格式
|
将PCM数据编码为Opus格式,以流式方式进行处理
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pcm_data: PCM字节数据
|
pcm_data: PCM字节数据
|
||||||
end_of_stream: 是否为流的结束
|
end_of_stream: 是否为流的结束,
|
||||||
|
callback: opus处理方法
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Opus数据包列表
|
Opus数据包列表
|
||||||
@@ -76,7 +76,6 @@ class OpusEncoderUtils:
|
|||||||
# 将新数据追加到缓冲区
|
# 将新数据追加到缓冲区
|
||||||
self.buffer = np.append(self.buffer, new_samples)
|
self.buffer = np.append(self.buffer, new_samples)
|
||||||
|
|
||||||
opus_packets = []
|
|
||||||
offset = 0
|
offset = 0
|
||||||
|
|
||||||
# 处理所有完整帧
|
# 处理所有完整帧
|
||||||
@@ -84,7 +83,7 @@ class OpusEncoderUtils:
|
|||||||
frame = self.buffer[offset : offset + self.total_frame_size]
|
frame = self.buffer[offset : offset + self.total_frame_size]
|
||||||
output = self._encode(frame)
|
output = self._encode(frame)
|
||||||
if output:
|
if output:
|
||||||
opus_packets.append(output)
|
callback(output)
|
||||||
offset += self.total_frame_size
|
offset += self.total_frame_size
|
||||||
|
|
||||||
# 保留未处理的样本
|
# 保留未处理的样本
|
||||||
@@ -98,11 +97,9 @@ class OpusEncoderUtils:
|
|||||||
|
|
||||||
output = self._encode(last_frame)
|
output = self._encode(last_frame)
|
||||||
if output:
|
if output:
|
||||||
opus_packets.append(output)
|
callback(output)
|
||||||
self.buffer = np.array([], dtype=np.int16)
|
self.buffer = np.array([], dtype=np.int16)
|
||||||
|
|
||||||
return opus_packets
|
|
||||||
|
|
||||||
def _encode(self, frame: np.ndarray) -> Optional[bytes]:
|
def _encode(self, frame: np.ndarray) -> Optional[bytes]:
|
||||||
"""编码一帧音频数据"""
|
"""编码一帧音频数据"""
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -1,15 +1,12 @@
|
|||||||
|
import io
|
||||||
import struct
|
import struct
|
||||||
|
from typing import Callable, Any
|
||||||
|
|
||||||
def decode_opus_from_file(input_file):
|
|
||||||
"""
|
|
||||||
从p3文件中解码 Opus 数据,并返回一个 Opus 数据包的列表以及总时长。
|
|
||||||
"""
|
|
||||||
opus_datas = []
|
|
||||||
total_frames = 0
|
|
||||||
sample_rate = 16000 # 文件采样率
|
|
||||||
frame_duration_ms = 60 # 帧时长
|
|
||||||
frame_size = int(sample_rate * frame_duration_ms / 1000)
|
|
||||||
|
|
||||||
|
def decode_opus_from_file_stream(input_file, callback: Callable[[Any], Any]):
|
||||||
|
"""
|
||||||
|
从p3文件中解码 Opus 数据,由 callback 处理 Opus 数据包。
|
||||||
|
"""
|
||||||
with open(input_file, 'rb') as f:
|
with open(input_file, 'rb') as f:
|
||||||
while True:
|
while True:
|
||||||
# 读取头部(4字节):[1字节类型,1字节保留,2字节长度]
|
# 读取头部(4字节):[1字节类型,1字节保留,2字节长度]
|
||||||
@@ -25,23 +22,13 @@ def decode_opus_from_file(input_file):
|
|||||||
if len(opus_data) != data_len:
|
if len(opus_data) != data_len:
|
||||||
raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the file.")
|
raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the file.")
|
||||||
|
|
||||||
opus_datas.append(opus_data)
|
callback(opus_data)
|
||||||
total_frames += 1
|
|
||||||
|
|
||||||
# 计算总时长
|
|
||||||
total_duration = (total_frames * frame_duration_ms) / 1000.0
|
|
||||||
return opus_datas, total_duration
|
|
||||||
|
|
||||||
def decode_opus_from_bytes(input_bytes):
|
def decode_opus_from_bytes_stream(input_bytes, callback: Callable[[Any], Any]):
|
||||||
"""
|
"""
|
||||||
从p3二进制数据中解码 Opus 数据,并返回一个 Opus 数据包的列表以及总时长。
|
从p3二进制数据中解码 Opus 数据,由 callback 处理 Opus 数据包。
|
||||||
"""
|
"""
|
||||||
import io
|
|
||||||
opus_datas = []
|
|
||||||
total_frames = 0
|
|
||||||
sample_rate = 16000 # 文件采样率
|
|
||||||
frame_duration_ms = 60 # 帧时长
|
|
||||||
frame_size = int(sample_rate * frame_duration_ms / 1000)
|
|
||||||
|
|
||||||
f = io.BytesIO(input_bytes)
|
f = io.BytesIO(input_bytes)
|
||||||
while True:
|
while True:
|
||||||
@@ -52,8 +39,4 @@ def decode_opus_from_bytes(input_bytes):
|
|||||||
opus_data = f.read(data_len)
|
opus_data = f.read(data_len)
|
||||||
if len(opus_data) != data_len:
|
if len(opus_data) != data_len:
|
||||||
raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the bytes.")
|
raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the bytes.")
|
||||||
opus_datas.append(opus_data)
|
callback(opus_data)
|
||||||
total_frames += 1
|
|
||||||
|
|
||||||
total_duration = (total_frames * frame_duration_ms) / 1000.0
|
|
||||||
return opus_datas, total_duration
|
|
||||||
|
|||||||
@@ -3,8 +3,8 @@ import socket
|
|||||||
import subprocess
|
import subprocess
|
||||||
import re
|
import re
|
||||||
import os
|
import os
|
||||||
import wave
|
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
|
from typing import Callable, Any
|
||||||
from core.utils import p3
|
from core.utils import p3
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import requests
|
import requests
|
||||||
@@ -211,7 +211,7 @@ def extract_json_from_string(input_string):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def audio_to_data(audio_file_path, is_opus=True):
|
def audio_to_data_stream(audio_file_path, is_opus=True, callback: Callable[[Any], Any]=None) -> None:
|
||||||
# 获取文件后缀名
|
# 获取文件后缀名
|
||||||
file_type = os.path.splitext(audio_file_path)[1]
|
file_type = os.path.splitext(audio_file_path)[1]
|
||||||
if file_type:
|
if file_type:
|
||||||
@@ -224,33 +224,29 @@ def audio_to_data(audio_file_path, is_opus=True):
|
|||||||
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
|
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
||||||
|
|
||||||
# 音频时长(秒)
|
|
||||||
duration = len(audio) / 1000.0
|
|
||||||
|
|
||||||
# 获取原始PCM数据(16位小端)
|
# 获取原始PCM数据(16位小端)
|
||||||
raw_data = audio.raw_data
|
raw_data = audio.raw_data
|
||||||
return pcm_to_data(raw_data, is_opus), duration
|
pcm_to_data_stream(raw_data, is_opus, callback)
|
||||||
|
|
||||||
|
|
||||||
def audio_bytes_to_data(audio_bytes, file_type, is_opus=True):
|
def audio_bytes_to_data_stream(audio_bytes, file_type, is_opus, callback: Callable[[Any], Any]) -> None:
|
||||||
"""
|
"""
|
||||||
直接用音频二进制数据转为opus/pcm数据,支持wav、mp3、p3
|
直接用音频二进制数据转为opus/pcm数据,支持wav、mp3、p3
|
||||||
"""
|
"""
|
||||||
if file_type == "p3":
|
if file_type == "p3":
|
||||||
# 直接用p3解码
|
# 直接用p3解码
|
||||||
return p3.decode_opus_from_bytes(audio_bytes)
|
return p3.decode_opus_from_bytes_stream(audio_bytes, callback)
|
||||||
else:
|
else:
|
||||||
# 其他格式用pydub
|
# 其他格式用pydub
|
||||||
audio = AudioSegment.from_file(
|
audio = AudioSegment.from_file(
|
||||||
BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"]
|
BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"]
|
||||||
)
|
)
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
||||||
duration = len(audio) / 1000.0
|
|
||||||
raw_data = audio.raw_data
|
raw_data = audio.raw_data
|
||||||
return pcm_to_data(raw_data, is_opus), duration
|
pcm_to_data_stream(raw_data, is_opus, callback)
|
||||||
|
|
||||||
|
|
||||||
def pcm_to_data(raw_data, is_opus=True):
|
def pcm_to_data_stream(raw_data, is_opus=True, callback: Callable[[Any], Any] = None):
|
||||||
# 初始化Opus编码器
|
# 初始化Opus编码器
|
||||||
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
||||||
|
|
||||||
@@ -258,7 +254,6 @@ def pcm_to_data(raw_data, is_opus=True):
|
|||||||
frame_duration = 60 # 60ms per frame
|
frame_duration = 60 # 60ms per frame
|
||||||
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
|
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
|
||||||
|
|
||||||
datas = []
|
|
||||||
# 按帧处理所有音频数据(包括最后一帧可能补零)
|
# 按帧处理所有音频数据(包括最后一帧可能补零)
|
||||||
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
|
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
|
||||||
# 获取当前帧的二进制数据
|
# 获取当前帧的二进制数据
|
||||||
@@ -273,39 +268,10 @@ def pcm_to_data(raw_data, is_opus=True):
|
|||||||
np_frame = np.frombuffer(chunk, dtype=np.int16)
|
np_frame = np.frombuffer(chunk, dtype=np.int16)
|
||||||
# 编码Opus数据
|
# 编码Opus数据
|
||||||
frame_data = encoder.encode(np_frame.tobytes(), frame_size)
|
frame_data = encoder.encode(np_frame.tobytes(), frame_size)
|
||||||
|
callback(frame_data)
|
||||||
else:
|
else:
|
||||||
frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
|
frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
|
||||||
|
callback(frame_data)
|
||||||
datas.append(frame_data)
|
|
||||||
|
|
||||||
return datas
|
|
||||||
|
|
||||||
|
|
||||||
def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1):
|
|
||||||
"""
|
|
||||||
将opus帧列表解码为wav字节流
|
|
||||||
"""
|
|
||||||
decoder = opuslib_next.Decoder(sample_rate, channels)
|
|
||||||
pcm_datas = []
|
|
||||||
|
|
||||||
frame_duration = 60 # ms
|
|
||||||
frame_size = int(sample_rate * frame_duration / 1000) # 960
|
|
||||||
|
|
||||||
for opus_frame in opus_datas:
|
|
||||||
# 解码为PCM(返回bytes,2字节/采样点)
|
|
||||||
pcm = decoder.decode(opus_frame, frame_size)
|
|
||||||
pcm_datas.append(pcm)
|
|
||||||
|
|
||||||
pcm_bytes = b"".join(pcm_datas)
|
|
||||||
|
|
||||||
# 写入wav字节流
|
|
||||||
wav_buffer = BytesIO()
|
|
||||||
with wave.open(wav_buffer, "wb") as wf:
|
|
||||||
wf.setnchannels(channels)
|
|
||||||
wf.setsampwidth(2) # 16bit
|
|
||||||
wf.setframerate(sample_rate)
|
|
||||||
wf.writeframes(pcm_bytes)
|
|
||||||
return wav_buffer.getvalue()
|
|
||||||
|
|
||||||
|
|
||||||
def check_vad_update(before_config, new_config):
|
def check_vad_update(before_config, new_config):
|
||||||
|
|||||||
@@ -1,140 +0,0 @@
|
|||||||
import os
|
|
||||||
import re
|
|
||||||
import yaml
|
|
||||||
import time
|
|
||||||
import hashlib
|
|
||||||
import portalocker
|
|
||||||
from typing import Dict
|
|
||||||
|
|
||||||
|
|
||||||
class FileLock:
|
|
||||||
def __init__(self, file, timeout=5):
|
|
||||||
self.file = file
|
|
||||||
self.timeout = timeout
|
|
||||||
self.start_time = None
|
|
||||||
|
|
||||||
def __enter__(self):
|
|
||||||
self.start_time = time.time()
|
|
||||||
while True:
|
|
||||||
try:
|
|
||||||
portalocker.lock(self.file, portalocker.LOCK_EX | portalocker.LOCK_NB)
|
|
||||||
return self.file
|
|
||||||
except portalocker.LockException:
|
|
||||||
if time.time() - self.start_time > self.timeout:
|
|
||||||
raise TimeoutError("获取文件锁超时")
|
|
||||||
time.sleep(0.1)
|
|
||||||
|
|
||||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
|
||||||
portalocker.unlock(self.file)
|
|
||||||
|
|
||||||
|
|
||||||
class WakeupWordsConfig:
|
|
||||||
def __init__(self):
|
|
||||||
self.config_file = "data/.wakeup_words.yaml"
|
|
||||||
self.assets_dir = "config/assets/wakeup_words"
|
|
||||||
self._ensure_directories()
|
|
||||||
self._config_cache = None
|
|
||||||
self._last_load_time = 0
|
|
||||||
self._cache_ttl = 1 # 缓存有效期(秒)
|
|
||||||
self._lock_timeout = 5 # 文件锁超时时间(秒)
|
|
||||||
|
|
||||||
def _ensure_directories(self):
|
|
||||||
"""确保必要的目录存在"""
|
|
||||||
os.makedirs(os.path.dirname(self.config_file), exist_ok=True)
|
|
||||||
os.makedirs(self.assets_dir, exist_ok=True)
|
|
||||||
|
|
||||||
def _load_config(self) -> Dict:
|
|
||||||
"""加载配置文件,使用缓存机制"""
|
|
||||||
current_time = time.time()
|
|
||||||
|
|
||||||
# 如果缓存有效,直接返回缓存
|
|
||||||
if (
|
|
||||||
self._config_cache is not None
|
|
||||||
and current_time - self._last_load_time < self._cache_ttl
|
|
||||||
):
|
|
||||||
return self._config_cache
|
|
||||||
|
|
||||||
try:
|
|
||||||
with open(self.config_file, "a+") as f:
|
|
||||||
with FileLock(f, timeout=self._lock_timeout):
|
|
||||||
f.seek(0)
|
|
||||||
content = f.read()
|
|
||||||
config = yaml.safe_load(content) if content else {}
|
|
||||||
self._config_cache = config
|
|
||||||
self._last_load_time = current_time
|
|
||||||
return config
|
|
||||||
except (TimeoutError, IOError) as e:
|
|
||||||
print(f"加载配置文件失败: {e}")
|
|
||||||
return {}
|
|
||||||
except Exception as e:
|
|
||||||
print(f"加载配置文件时发生未知错误: {e}")
|
|
||||||
return {}
|
|
||||||
|
|
||||||
def _save_config(self, config: Dict):
|
|
||||||
"""保存配置到文件,使用文件锁保护"""
|
|
||||||
try:
|
|
||||||
with open(self.config_file, "w") as f:
|
|
||||||
with FileLock(f, timeout=self._lock_timeout):
|
|
||||||
yaml.dump(config, f, allow_unicode=True)
|
|
||||||
self._config_cache = config
|
|
||||||
self._last_load_time = time.time()
|
|
||||||
except (TimeoutError, IOError) as e:
|
|
||||||
print(f"保存配置文件失败: {e}")
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
print(f"保存配置文件时发生未知错误: {e}")
|
|
||||||
raise
|
|
||||||
|
|
||||||
def get_wakeup_response(self, voice: str) -> Dict:
|
|
||||||
voice = hashlib.md5(voice.encode()).hexdigest()
|
|
||||||
"""获取唤醒词回复配置"""
|
|
||||||
config = self._load_config()
|
|
||||||
|
|
||||||
if not config or voice not in config:
|
|
||||||
return None
|
|
||||||
|
|
||||||
# 检查文件大小
|
|
||||||
file_path = config[voice]["file_path"]
|
|
||||||
if not os.path.exists(file_path) or os.stat(file_path).st_size < (15 * 1024):
|
|
||||||
return None
|
|
||||||
|
|
||||||
return config[voice]
|
|
||||||
|
|
||||||
def update_wakeup_response(self, voice: str, file_path: str, text: str):
|
|
||||||
"""更新唤醒词回复配置"""
|
|
||||||
try:
|
|
||||||
# 过滤表情符号
|
|
||||||
filtered_text = re.sub(r'[\U0001F600-\U0001F64F\U0001F900-\U0001F9FF]', '', text)
|
|
||||||
|
|
||||||
config = self._load_config()
|
|
||||||
voice_hash = hashlib.md5(voice.encode()).hexdigest()
|
|
||||||
config[voice_hash] = {
|
|
||||||
"voice": voice,
|
|
||||||
"file_path": file_path,
|
|
||||||
"time": time.time(),
|
|
||||||
"text": filtered_text,
|
|
||||||
}
|
|
||||||
self._save_config(config)
|
|
||||||
except Exception as e:
|
|
||||||
print(f"更新唤醒词回复配置失败: {e}")
|
|
||||||
raise
|
|
||||||
|
|
||||||
def generate_file_path(self, voice: str) -> str:
|
|
||||||
"""生成音频文件路径,使用voice的哈希值作为文件名"""
|
|
||||||
try:
|
|
||||||
# 生成voice的哈希值
|
|
||||||
voice_hash = hashlib.md5(voice.encode()).hexdigest()
|
|
||||||
file_path = os.path.join(self.assets_dir, f"{voice_hash}.wav")
|
|
||||||
|
|
||||||
# 如果文件已存在,先删除
|
|
||||||
if os.path.exists(file_path):
|
|
||||||
try:
|
|
||||||
os.remove(file_path)
|
|
||||||
except Exception as e:
|
|
||||||
print(f"删除已存在的音频文件失败: {e}")
|
|
||||||
raise
|
|
||||||
|
|
||||||
return file_path
|
|
||||||
except Exception as e:
|
|
||||||
print(f"生成音频文件路径失败: {e}")
|
|
||||||
raise
|
|
||||||
Reference in New Issue
Block a user