update: TTS复用链接,VAD双阈值判断 (#1742)

* update: TTS复用链接,VAD双阈值判断

* fix: 播放音乐时,引导词卡顿

* update:优化引导词

* update: 优化chat函数流程 优化huoshan处理 会话保持一致性

* fix: 等待时可能已经完成(设置为None),后续对None错误访问

* update:增加书名号

* fix: 打断状态未重置 监听未完成时服务端可能还在发送数据 此时复用链接会接收上个语音的残余 需要两者一同关闭

---------

Co-authored-by: hrz <1710360675@qq.com>
This commit is contained in:
欣南科技
2025-07-09 18:34:04 +08:00
committed by GitHub
co-authored by hrz
8 changed files with 137 additions and 128 deletions
+1
View File
@@ -326,6 +326,7 @@ VAD:
SileroVAD: SileroVAD:
type: silero type: silero
threshold: 0.5 threshold: 0.5
threshold_low: 0.3
model_dir: models/snakers4_silero-vad model_dir: models/snakers4_silero-vad
min_silence_duration_ms: 200 # 如果说话停顿比较长,可以把这个值设置大一些 min_silence_duration_ms: 200 # 如果说话停顿比较长,可以把这个值设置大一些
+27 -17
View File
@@ -17,6 +17,7 @@ from core.utils.util import (
filter_sensitive_info, filter_sensitive_info,
) )
from typing import Dict, Any from typing import Dict, Any
from collections import deque
from core.utils.modules_initialize import ( from core.utils.modules_initialize import (
initialize_modules, initialize_modules,
initialize_tts, initialize_tts,
@@ -112,6 +113,8 @@ class ConnectionHandler:
self.client_have_voice = False self.client_have_voice = False
self.last_activity_time = 0.0 # 统一的活动时间戳(毫秒) self.last_activity_time = 0.0 # 统一的活动时间戳(毫秒)
self.client_voice_stop = False self.client_voice_stop = False
self.client_voice_window = deque(maxlen=5)
self.last_is_voice = False
# asr相关变量 # asr相关变量
# 因为实际部署时可能会用到公共的本地ASR,不能把变量暴露给公共ASR # 因为实际部署时可能会用到公共的本地ASR,不能把变量暴露给公共ASR
@@ -608,13 +611,24 @@ 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): def chat(self, query, tool_call=False, 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: if not tool_call:
self.dialogue.put(Message(role="user", content=query)) self.dialogue.put(Message(role="user", content=query))
# 为最顶层时新建会话ID和发送FIRST请求
if depth == 0:
self.sentence_id = str(uuid.uuid4().hex)
self.tts.tts_text_queue.put(
TTSMessageDTO(
sentence_id=self.sentence_id,
sentence_type=SentenceType.FIRST,
content_type=ContentType.ACTION,
)
)
# Define intent functions # Define intent functions
functions = None functions = None
if self.intent_type == "function_call" and hasattr(self, "func_handler"): if self.intent_type == "function_call" and hasattr(self, "func_handler"):
@@ -630,8 +644,6 @@ class ConnectionHandler:
) )
memory_str = future.result() memory_str = future.result()
self.sentence_id = str(uuid.uuid4().hex)
if self.intent_type == "function_call" and functions is not None: if self.intent_type == "function_call" and functions is not None:
# 使用支持functions的streaming接口 # 使用支持functions的streaming接口
llm_responses = self.llm.response_with_functions( llm_responses = self.llm.response_with_functions(
@@ -654,7 +666,6 @@ class ConnectionHandler:
function_id = None function_id = None
function_arguments = "" function_arguments = ""
content_arguments = "" content_arguments = ""
text_index = 0
self.client_abort = False self.client_abort = False
for response in llm_responses: for response in llm_responses:
if self.client_abort: if self.client_abort:
@@ -684,14 +695,6 @@ class ConnectionHandler:
if content is not None and len(content) > 0: if content is not None and len(content) > 0:
if not tool_call_flag: if not tool_call_flag:
response_message.append(content) response_message.append(content)
if text_index == 0:
self.tts.tts_text_queue.put(
TTSMessageDTO(
sentence_id=self.sentence_id,
sentence_type=SentenceType.FIRST,
content_type=ContentType.ACTION,
)
)
self.tts.tts_text_queue.put( self.tts.tts_text_queue.put(
TTSMessageDTO( TTSMessageDTO(
sentence_id=self.sentence_id, sentence_id=self.sentence_id,
@@ -700,7 +703,6 @@ class ConnectionHandler:
content_detail=content, content_detail=content,
) )
) )
text_index += 1
# 处理function call # 处理function call
if tool_call_flag: if tool_call_flag:
bHasError = False bHasError = False
@@ -725,6 +727,11 @@ class ConnectionHandler:
f"function call error: {content_arguments}" f"function call error: {content_arguments}"
) )
if not bHasError: if not bHasError:
# 如需要大模型先处理一轮,添加相关处理后的日志情况
if len(response_message) > 0:
self.dialogue.put(
Message(role="assistant", content="".join(response_message))
)
response_message.clear() response_message.clear()
self.logger.bind(tag=TAG).debug( self.logger.bind(tag=TAG).debug(
f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}" f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}"
@@ -742,14 +749,14 @@ class ConnectionHandler:
), ),
self.loop, self.loop,
).result() ).result()
self._handle_function_result(result, function_call_data) self._handle_function_result(result, function_call_data, depth=depth)
# 存储对话内容 # 存储对话内容
if len(response_message) > 0: if len(response_message) > 0:
self.dialogue.put( self.dialogue.put(
Message(role="assistant", content="".join(response_message)) Message(role="assistant", content="".join(response_message))
) )
if text_index > 0: if depth == 0:
self.tts.tts_text_queue.put( self.tts.tts_text_queue.put(
TTSMessageDTO( TTSMessageDTO(
sentence_id=self.sentence_id, sentence_id=self.sentence_id,
@@ -764,7 +771,7 @@ class ConnectionHandler:
return True return True
def _handle_function_result(self, result, function_call_data): def _handle_function_result(self, result, function_call_data, depth):
if result.action == Action.RESPONSE: # 直接回复前端 if result.action == Action.RESPONSE: # 直接回复前端
text = result.response text = result.response
self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text) self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text)
@@ -801,7 +808,7 @@ class ConnectionHandler:
content=text, content=text,
) )
) )
self.chat(text, tool_call=True) self.chat(text, tool_call=True, 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)
@@ -914,6 +921,9 @@ class ConnectionHandler:
except Exception as ws_error: except Exception as ws_error:
self.logger.bind(tag=TAG).error(f"关闭WebSocket连接时出错: {ws_error}") self.logger.bind(tag=TAG).error(f"关闭WebSocket连接时出错: {ws_error}")
if self.tts:
await self.tts.close()
# 最后关闭线程池(避免阻塞) # 最后关闭线程池(避免阻塞)
if self.executor: if self.executor:
try: try:
@@ -6,9 +6,8 @@ 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
from core.providers.tools.device_mcp import call_mcp_tool
from plugins_func.register import Action, ActionResponse from plugins_func.register import Action, ActionResponse
from loguru import logger from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType
TAG = __name__ TAG = __name__
@@ -29,6 +28,8 @@ async def handle_user_intent(conn, text):
intent_result = await analyze_intent_with_llm(conn, text) intent_result = await analyze_intent_with_llm(conn, text)
if not intent_result: if not intent_result:
return False return False
# 会话开始时生成sentence_id
conn.sentence_id = str(uuid.uuid4().hex)
# 处理各种意图 # 处理各种意图
return await process_intent_result(conn, intent_result, text) return await process_intent_result(conn, intent_result, text)
@@ -153,5 +154,19 @@ async def process_intent_result(conn, intent_result, original_text):
def speak_txt(conn, text): def speak_txt(conn, text):
conn.tts.tts_text_queue.put(
TTSMessageDTO(
sentence_id=conn.sentence_id,
sentence_type=SentenceType.FIRST,
content_type=ContentType.ACTION,
)
)
conn.tts.tts_one_sentence(conn, ContentType.TEXT, content_detail=text) conn.tts.tts_one_sentence(conn, ContentType.TEXT, content_detail=text)
conn.tts.tts_text_queue.put(
TTSMessageDTO(
sentence_id=conn.sentence_id,
sentence_type=SentenceType.LAST,
content_type=ContentType.ACTION,
)
)
conn.dialogue.put(Message(role="assistant", content=text)) conn.dialogue.put(Message(role="assistant", content=text))
+4 -18
View File
@@ -161,13 +161,6 @@ class TTSProviderBase(ABC):
else: else:
sentence_id = str(uuid.uuid4()).replace("-", "") sentence_id = str(uuid.uuid4()).replace("-", "")
conn.sentence_id = sentence_id conn.sentence_id = sentence_id
self.tts_text_queue.put(
TTSMessageDTO(
sentence_id=sentence_id,
sentence_type=SentenceType.FIRST,
content_type=ContentType.ACTION,
)
)
# 对于单句的文本,进行分段处理 # 对于单句的文本,进行分段处理
segments = re.split(r"([。!?!?;\n])", content_detail) segments = re.split(r"([。!?!?;\n])", content_detail)
for seg in segments: for seg in segments:
@@ -180,13 +173,6 @@ class TTSProviderBase(ABC):
content_file=content_file, content_file=content_file,
) )
) )
self.tts_text_queue.put(
TTSMessageDTO(
sentence_id=sentence_id,
sentence_type=SentenceType.LAST,
content_type=ContentType.ACTION,
)
)
async def open_audio_channels(self, conn): async def open_audio_channels(self, conn):
self.conn = conn self.conn = conn
@@ -209,6 +195,8 @@ class TTSProviderBase(ABC):
while not self.conn.stop_event.is_set(): while not self.conn.stop_event.is_set():
try: try:
message = self.tts_text_queue.get(timeout=1) message = self.tts_text_queue.get(timeout=1)
if message.sentence_type == SentenceType.FIRST:
self.conn.client_abort = False
if self.conn.client_abort: if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程") logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
continue continue
@@ -362,10 +350,8 @@ class TTSProviderBase(ABC):
return audio_datas return audio_datas
def _process_before_stop_play_files(self): def _process_before_stop_play_files(self):
for tts_file, text in self.before_stop_play_files: for audio_datas, text in self.before_stop_play_files:
if tts_file and os.path.exists(tts_file): self.tts_audio_queue.put((SentenceType.MIDDLE, audio_datas, text))
audio_datas = self._process_audio_file(tts_file)
self.tts_audio_queue.put((SentenceType.MIDDLE, audio_datas, text))
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))
@@ -10,7 +10,6 @@ from config.logger import setup_logging
from core.utils import opus_encoder_utils from core.utils import opus_encoder_utils
from core.utils.util import check_model_key from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
from core.handle.abortHandle import handleAbortMessage
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
from asyncio import Task from asyncio import Task
@@ -175,6 +174,9 @@ class TTSProvider(TTSProviderBase):
async def _ensure_connection(self): async def _ensure_connection(self):
"""建立新的WebSocket连接""" """建立新的WebSocket连接"""
try: try:
if self.ws:
logger.bind(tag=TAG).info(f"使用已有链接...")
return self.ws
logger.bind(tag=TAG).info("开始建立新连接...") logger.bind(tag=TAG).info("开始建立新连接...")
ws_header = { ws_header = {
"X-Api-App-Key": self.appId, "X-Api-App-Key": self.appId,
@@ -200,6 +202,10 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).debug( logger.bind(tag=TAG).debug(
f"收到TTS任务|{message.sentence_type.name} {message.content_type.name} | 会话ID: {self.conn.sentence_id}" f"收到TTS任务|{message.sentence_type.name} {message.content_type.name} | 会话ID: {self.conn.sentence_id}"
) )
if message.sentence_type == SentenceType.FIRST:
self.conn.client_abort = False
if self.conn.client_abort: if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程") logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
continue continue
@@ -244,9 +250,12 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).info( logger.bind(tag=TAG).info(
f"添加音频文件到待播放列表: {message.content_file}" f"添加音频文件到待播放列表: {message.content_file}"
) )
self.before_stop_play_files.append( if message.content_file and os.path.exists(message.content_file):
(message.content_file, message.content_detail) # 先处理文件音频数据
) file_audio = self._process_audio_file(message.content_file)
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:
@@ -295,25 +304,15 @@ class TTSProvider(TTSProviderBase):
async def start_session(self, session_id): async def start_session(self, session_id):
logger.bind(tag=TAG).info(f"开始会话~~{session_id}") logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
try: try:
task = self._monitor_task # 会话开始时检测上个会话的监听状态
if ( if (
task is not None self._monitor_task is not None
and isinstance(task, Task) and isinstance(self._monitor_task, Task)
and not task.done() and not self._monitor_task.done()
): ):
logger.bind(tag=TAG).info("等待上一个监听任务结束...") logger.bind(tag=TAG).info("检测到未完成的上个会话,关闭监听任务和连接...")
if self.ws is not None: await self.close()
logger.bind(tag=TAG).info("强制关闭上一个WebSocket连接以唤醒监听任务...")
try:
await self.ws.close()
except Exception as e:
logger.bind(tag=TAG).warning(f"关闭上一个ws异常: {e}")
self.ws = None
try:
await asyncio.wait_for(task, timeout=8)
except Exception as e:
logger.bind(tag=TAG).warning(f"等待监听任务异常: {e}")
self._monitor_task = None
# 建立新连接 # 建立新连接
await self._ensure_connection() await self._ensure_connection()
@@ -336,19 +335,7 @@ class TTSProvider(TTSProviderBase):
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}") logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
# 确保清理资源 # 确保清理资源
if hasattr(self, "_monitor_task"): await self.close()
try:
self._monitor_task.cancel()
await self._monitor_task
except:
pass
self._monitor_task = None
if self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
raise raise
async def finish_session(self, session_id): async def finish_session(self, session_id):
@@ -368,7 +355,7 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).info("会话结束请求已发送") logger.bind(tag=TAG).info("会话结束请求已发送")
# 等待监听任务完成 # 等待监听任务完成
if hasattr(self, "_monitor_task"): if self._monitor_task:
try: try:
await self._monitor_task await self._monitor_task
except Exception as e: except Exception as e:
@@ -378,28 +365,25 @@ class TTSProvider(TTSProviderBase):
finally: finally:
self._monitor_task = None self._monitor_task = None
# 关闭连接
await self.close()
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"关闭会话失败: {str(e)}") logger.bind(tag=TAG).error(f"关闭会话失败: {str(e)}")
# 确保清理资源 # 确保清理资源
if hasattr(self, "_monitor_task"): await self.close()
try:
self._monitor_task.cancel()
await self._monitor_task
except:
pass
self._monitor_task = None
if self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
raise raise
async def close(self): async def close(self):
"""资源清理方法""" """资源清理方法"""
# 取消监听任务
if self._monitor_task:
try:
self._monitor_task.cancel()
await self._monitor_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"关闭时取消监听任务错误: {e}")
self._monitor_task = None
if self.ws: if self.ws:
try: try:
await self.ws.close() await self.ws.close()
@@ -413,6 +397,7 @@ class TTSProvider(TTSProviderBase):
is_first_sentence = True is_first_sentence = True
first_sentence_segment_count = 0 # 添加计数器 first_sentence_segment_count = 0 # 添加计数器
try: try:
session_finished = False # 标记会话是否正常结束
while not self.conn.stop_event.is_set(): while not self.conn.stop_event.is_set():
try: try:
# 确保 `recv()` 运行在同一个 event loop # 确保 `recv()` 运行在同一个 event loop
@@ -466,6 +451,7 @@ class TTSProvider(TTSProviderBase):
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()
session_finished = True
break break
except websockets.ConnectionClosed: except websockets.ConnectionClosed:
logger.bind(tag=TAG).warning("WebSocket连接已关闭") logger.bind(tag=TAG).warning("WebSocket连接已关闭")
@@ -476,15 +462,15 @@ class TTSProvider(TTSProviderBase):
) )
traceback.print_exc() traceback.print_exc()
break break
finally: # 仅在连接异常时才关闭
# 确保清理资源 if not session_finished and self.ws:
if self.ws:
try: try:
await self.ws.close() await self.ws.close()
except: except:
pass pass
self.ws = None self.ws = None
# 监听任务退出时清理引用 # 监听任务退出时清理引用
finally:
self._monitor_task = None self._monitor_task = None
async def send_event( async def send_event(
@@ -1,3 +1,4 @@
import os
import queue import queue
import asyncio import asyncio
import traceback import traceback
@@ -63,9 +64,12 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).info( logger.bind(tag=TAG).info(
f"添加音频文件到待播放列表: {message.content_file}" f"添加音频文件到待播放列表: {message.content_file}"
) )
self.before_stop_play_files.append( if message.content_file and os.path.exists(message.content_file):
(message.content_file, message.content_detail) # 先处理文件音频数据
) file_audio = self._process_audio_file(message.content_file)
self.before_stop_play_files.append(
(file_audio, message.content_detail)
)
if message.sentence_type == SentenceType.LAST: if message.sentence_type == SentenceType.LAST:
# 处理剩余的文本 # 处理剩余的文本
@@ -23,22 +23,24 @@ class VADProvider(VADProviderBase):
# 处理空字符串的情况 # 处理空字符串的情况
threshold = config.get("threshold", "0.5") threshold = config.get("threshold", "0.5")
threshold_low = config.get("threshold_low", "0.2")
min_silence_duration_ms = config.get("min_silence_duration_ms", "1000") min_silence_duration_ms = config.get("min_silence_duration_ms", "1000")
self.vad_threshold = float(threshold) if threshold else 0.5 self.vad_threshold = float(threshold) if threshold else 0.5
self.vad_threshold_low = float(threshold_low) if threshold_low else 0.2
self.silence_threshold_ms = ( self.silence_threshold_ms = (
int(min_silence_duration_ms) if min_silence_duration_ms else 1000 int(min_silence_duration_ms) if min_silence_duration_ms else 1000
) )
# 至少要多少帧才算有语音
self.frame_window_threshold = 3
def is_vad(self, conn, opus_packet): def is_vad(self, conn, opus_packet):
try: try:
pcm_frame = self.decoder.decode(opus_packet, 960) pcm_frame = self.decoder.decode(opus_packet, 960)
conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区 conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区
# 确保帧计数器存在
if not hasattr(conn, "client_voice_frame_count"):
conn.client_voice_frame_count = 0
# 处理缓冲区中的完整帧(每次处理512采样点) # 处理缓冲区中的完整帧(每次处理512采样点)
client_have_voice = False client_have_voice = False
while len(conn.client_audio_buffer) >= 512 * 2: while len(conn.client_audio_buffer) >= 512 * 2:
@@ -54,15 +56,21 @@ class VADProvider(VADProviderBase):
# 检测语音活动 # 检测语音活动
with torch.no_grad(): with torch.no_grad():
speech_prob = self.model(audio_tensor, 16000).item() speech_prob = self.model(audio_tensor, 16000).item()
is_voice = speech_prob >= self.vad_threshold
if is_voice: # 双阈值判断
conn.client_voice_frame_count += 1 if speech_prob >= self.vad_threshold:
is_voice = True
elif speech_prob <= self.vad_threshold_low:
is_voice = False
else: else:
conn.client_voice_frame_count = 0 is_voice = conn.last_is_voice
# 只有连续4帧检测到语音才认为有 # 声音没低于最低值则延续前一个状态,判断为有
client_have_voice = conn.client_voice_frame_count >= 4 conn.last_is_voice = is_voice
# 更新滑动窗口
conn.client_voice_window.append(is_voice)
client_have_voice = (conn.client_voice_window.count(True) >= self.frame_window_threshold)
# 如果之前有声音,但本次没有声音,且与上次有声音的时间差已经超过了静默阈值,则认为已经说完一句话 # 如果之前有声音,但本次没有声音,且与上次有声音的时间差已经超过了静默阈值,则认为已经说完一句话
if conn.client_have_voice and not client_have_voice: if conn.client_have_voice and not client_have_voice:
@@ -1,13 +1,10 @@
from config.logger import setup_logging
import os import os
import re import re
import time import time
import random import random
import asyncio
import difflib import difflib
import traceback import traceback
from pathlib import Path from pathlib import Path
from core.utils import p3
from core.handle.sendAudioHandle import send_stt_message from core.handle.sendAudioHandle import send_stt_message
from plugins_func.register import register_function, ToolType, ActionResponse, Action from plugins_func.register import register_function, ToolType, ActionResponse, Action
from core.utils.dialogue import Message from core.utils.dialogue import Message
@@ -51,8 +48,8 @@ def play_music(conn, song_name: str):
) )
# 提交异步任务 # 提交异步任务
future = asyncio.run_coroutine_threadsafe( task = conn.loop.create_task(
handle_music_command(conn, music_intent), conn.loop handle_music_command(conn, music_intent) # 封装异步逻辑
) )
# 非阻塞回调处理 # 非阻塞回调处理
@@ -63,7 +60,7 @@ def play_music(conn, song_name: str):
except Exception as e: except Exception as e:
conn.logger.bind(tag=TAG).error(f"播放失败: {e}") conn.logger.bind(tag=TAG).error(f"播放失败: {e}")
future.add_done_callback(handle_done) task.add_done_callback(handle_done)
return ActionResponse( return ActionResponse(
action=Action.NONE, result="指令已接收", response="正在为您播放音乐" action=Action.NONE, result="指令已接收", response="正在为您播放音乐"
@@ -178,13 +175,13 @@ def _get_random_play_prompt(song_name):
# 移除文件扩展名 # 移除文件扩展名
clean_name = os.path.splitext(song_name)[0] clean_name = os.path.splitext(song_name)[0]
prompts = [ prompts = [
f"正在为您播放,{clean_name}", f"正在为您播放,{clean_name}",
f"请欣赏歌曲,{clean_name}", f"请欣赏歌曲,{clean_name}",
f"即将为您播放,{clean_name}", f"即将为您播放,{clean_name}",
f"为您带来,{clean_name}", f"现在为您带来,{clean_name}",
f"让我们聆听,{clean_name}", f"让我们一起聆听,{clean_name}",
f"接下来请欣赏,{clean_name}", f"接下来请欣赏,{clean_name}",
f"为您献上,{clean_name}", f"此刻为您献上,{clean_name}",
] ]
# 直接使用random.choice,不设置seed # 直接使用random.choice,不设置seed
return random.choice(prompts) return random.choice(prompts)
@@ -218,13 +215,14 @@ async def play_local_music(conn, specific_file=None):
await send_stt_message(conn, text) await send_stt_message(conn, text)
conn.dialogue.put(Message(role="assistant", content=text)) conn.dialogue.put(Message(role="assistant", content=text))
conn.tts.tts_text_queue.put( if conn.intent_type == "intent_llm":
TTSMessageDTO( conn.tts.tts_text_queue.put(
sentence_id=conn.sentence_id, TTSMessageDTO(
sentence_type=SentenceType.FIRST, sentence_id=conn.sentence_id,
content_type=ContentType.ACTION, sentence_type=SentenceType.FIRST,
content_type=ContentType.ACTION,
)
) )
)
conn.tts.tts_text_queue.put( conn.tts.tts_text_queue.put(
TTSMessageDTO( TTSMessageDTO(
sentence_id=conn.sentence_id, sentence_id=conn.sentence_id,
@@ -241,13 +239,14 @@ async def play_local_music(conn, specific_file=None):
content_file=music_path, content_file=music_path,
) )
) )
conn.tts.tts_text_queue.put( if conn.intent_type == "intent_llm":
TTSMessageDTO( conn.tts.tts_text_queue.put(
sentence_id=conn.sentence_id, TTSMessageDTO(
sentence_type=SentenceType.LAST, sentence_id=conn.sentence_id,
content_type=ContentType.ACTION, sentence_type=SentenceType.LAST,
content_type=ContentType.ACTION,
)
) )
)
except Exception as e: except Exception as e:
conn.logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}") conn.logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}")