fix: 反复打断任务没有被清除,vad四帧语音识别

This commit is contained in:
Sakura-RanChen
2025-06-20 16:25:08 +08:00
parent 979fea0d60
commit 12cee4027a
3 changed files with 41 additions and 5 deletions
+3 -2
View File
@@ -115,6 +115,7 @@ class ConnectionHandler:
self.client_have_voice_last_time = 0.0
self.client_no_voice_last_time = 0.0
self.client_voice_stop = False
self.client_voice_frame_count = 0
# asr相关变量
# 因为实际部署时可能会用到公共的本地ASR,不能把变量暴露给公共ASR
@@ -627,8 +628,8 @@ class ConnectionHandler:
)
memory_str = future.result()
uuid_str = str(uuid.uuid4()).replace("-", "")
self.sentence_id = uuid_str
self.sentence_id = str(uuid.uuid4().hex)
if self.intent_type == "function_call" and functions is not None:
# 使用支持functions的streaming接口
@@ -12,6 +12,8 @@ from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase
from core.handle.abortHandle import handleAbortMessage
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
from asyncio import Task
TAG = __name__
logger = setup_logging()
@@ -141,6 +143,7 @@ class TTSProvider(TTSProviderBase):
super().__init__(config, delete_audio_file)
self.ws = None
self.interface_type = InterfaceType.DUAL_STREAM
self._monitor_task = None # 监听任务引用
self.appId = config.get("appid")
self.access_token = config.get("access_token")
self.cluster = config.get("cluster")
@@ -270,8 +273,7 @@ class TTSProvider(TTSProviderBase):
try:
# 建立新连接
if self.ws is None:
await handleAbortMessage(self.conn)
logger.bind(tag=TAG).error(f"WebSocket连接不存在,终止发送文本")
logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本")
return
# 过滤Markdown
@@ -293,6 +295,25 @@ class TTSProvider(TTSProviderBase):
async def start_session(self, session_id):
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
try:
task = self._monitor_task
if (
task is not None
and isinstance(task, Task)
and not task.done()
):
logger.bind(tag=TAG).info("等待上一个监听任务结束...")
if self.ws is not None:
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()
@@ -463,6 +484,8 @@ class TTSProvider(TTSProviderBase):
except:
pass
self.ws = None
# 监听任务退出时清理引用
self._monitor_task = None
async def send_event(
self,
@@ -35,6 +35,10 @@ class VADProvider(VADProviderBase):
pcm_frame = self.decoder.decode(opus_packet, 960)
conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区
# 初始化帧计数器
if not hasattr(conn, "client_voice_frame_count"):
conn.client_voice_frame_count = 0
# 处理缓冲区中的完整帧(每次处理512采样点)
client_have_voice = False
while len(conn.client_audio_buffer) >= 512 * 2:
@@ -50,7 +54,15 @@ class VADProvider(VADProviderBase):
# 检测语音活动
with torch.no_grad():
speech_prob = self.model(audio_tensor, 16000).item()
client_have_voice = speech_prob >= self.vad_threshold
is_voice = speech_prob >= self.vad_threshold
if is_voice:
conn.client_voice_frame_count += 1
else:
conn.client_voice_frame_count = 0
# 只有连续4帧检测到语音才认为有语音
client_have_voice = conn.client_voice_frame_count >= 4
# 如果之前有声音,但本次没有声音,且与上次有声音的时间差已经超过了静默阈值,则认为已经说完一句话
if conn.client_have_voice and not client_have_voice: