Merge pull request #2695 from xinnan-tech/py_device_bind

update: 未绑定设备策略优化
This commit is contained in:
hrz
2025-12-13 23:30:44 +08:00
committed by GitHub
7 changed files with 139 additions and 70 deletions
+50 -23
View File
@@ -69,6 +69,7 @@ class ConnectionHandler:
self.server = server # 保存server实例的引用 self.server = server # 保存server实例的引用
self.need_bind = False # 是否需要绑定设备 self.need_bind = False # 是否需要绑定设备
self.bind_completed_event = asyncio.Event()
self.bind_code = None # 绑定设备的验证码 self.bind_code = None # 绑定设备的验证码
self.last_bind_prompt_time = 0 # 上次播放绑定提示的时间戳(秒) self.last_bind_prompt_time = 0 # 上次播放绑定提示的时间戳(秒)
self.bind_prompt_interval = 60 # 绑定提示播放间隔(秒) self.bind_prompt_interval = 60 # 绑定提示播放间隔(秒)
@@ -266,28 +267,41 @@ class ConnectionHandler:
f"保存记忆后关闭连接失败: {close_error}" f"保存记忆后关闭连接失败: {close_error}"
) )
async def _route_message(self, message): async def _discard_message_with_bind_prompt(self):
"""消息路由""" """丢弃消息并检查是否需要播放绑定提示"""
if isinstance(message, str):
await handleTextMessage(self, message)
elif isinstance(message, bytes):
if self.vad is None or self.asr is None:
return
# 未绑定设备直接丢弃所有音频,不进行ASR处理
if self.need_bind:
current_time = time.time() current_time = time.time()
# 检查是否需要播放绑定提示 # 检查是否需要播放绑定提示
if ( if current_time - self.last_bind_prompt_time >= self.bind_prompt_interval:
current_time - self.last_bind_prompt_time
>= self.bind_prompt_interval
):
self.last_bind_prompt_time = current_time self.last_bind_prompt_time = current_time
# 复用现有的绑定提示逻辑 # 复用现有的绑定提示逻辑
from core.handle.receiveAudioHandle import check_bind_device from core.handle.receiveAudioHandle import check_bind_device
asyncio.create_task(check_bind_device(self)) asyncio.create_task(check_bind_device(self))
# 直接丢弃音频,不进行ASR处理
async def _route_message(self, message):
"""消息路由"""
# 检查是否已经获取到真实的绑定状态
if not self.bind_completed_event.is_set():
# 还没有获取到真实状态,等待直到获取到真实状态或超时
try:
await asyncio.wait_for(self.bind_completed_event.wait(), timeout=1)
except asyncio.TimeoutError:
# 超时仍未获取到真实状态,丢弃消息
await self._discard_message_with_bind_prompt()
return
# 已经获取到真实状态,检查是否需要绑定
if self.need_bind:
# 需要绑定,丢弃消息
await self._discard_message_with_bind_prompt()
return
# 不需要绑定,继续处理消息
if isinstance(message, str):
await handleTextMessage(self, message)
elif isinstance(message, bytes):
if self.vad is None or self.asr is None:
return return
# 处理来自MQTT网关的音频包 # 处理来自MQTT网关的音频包
@@ -413,6 +427,14 @@ class ConnectionHandler:
def _initialize_components(self): def _initialize_components(self):
try: try:
if self.tts is None:
self.tts = self._initialize_tts()
# 打开语音合成通道
asyncio.run_coroutine_threadsafe(
self.tts.open_audio_channels(self), self.loop
)
if self.need_bind:
return
self.selected_module_str = build_module_string( self.selected_module_str = build_module_string(
self.config.get("selected_module", {}) self.config.get("selected_module", {})
) )
@@ -436,17 +458,10 @@ class ConnectionHandler:
# 初始化声纹识别 # 初始化声纹识别
self._initialize_voiceprint() self._initialize_voiceprint()
# 打开语音识别通道 # 打开语音识别通道
asyncio.run_coroutine_threadsafe( asyncio.run_coroutine_threadsafe(
self.asr.open_audio_channels(self), self.loop self.asr.open_audio_channels(self), self.loop
) )
if self.tts is None:
self.tts = self._initialize_tts()
# 打开语音合成通道
asyncio.run_coroutine_threadsafe(
self.tts.open_audio_channels(self), self.loop
)
"""加载记忆""" """加载记忆"""
self._initialize_memory() self._initialize_memory()
@@ -461,6 +476,7 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error(f"实例化组件失败: {e}") self.logger.bind(tag=TAG).error(f"实例化组件失败: {e}")
def _init_prompt_enhancement(self): def _init_prompt_enhancement(self):
# 更新上下文信息 # 更新上下文信息
self.prompt_manager.update_context_info(self, self.client_ip) self.prompt_manager.update_context_info(self, self.client_ip)
enhanced_prompt = self.prompt_manager.build_enhanced_prompt( enhanced_prompt = self.prompt_manager.build_enhanced_prompt(
@@ -496,7 +512,11 @@ class ConnectionHandler:
def _initialize_asr(self): def _initialize_asr(self):
"""初始化ASR""" """初始化ASR"""
if self._asr is not None and hasattr(self._asr, "interface_type") and self._asr.interface_type == InterfaceType.LOCAL: if (
self._asr is not None
and hasattr(self._asr, "interface_type")
and self._asr.interface_type == InterfaceType.LOCAL
):
# 如果公共ASR是本地服务,则直接返回 # 如果公共ASR是本地服务,则直接返回
# 因为本地一个实例ASR,可以被多个连接共享 # 因为本地一个实例ASR,可以被多个连接共享
asr = self._asr asr = self._asr
@@ -536,6 +556,8 @@ class ConnectionHandler:
async def _initialize_private_config_async(self): async def _initialize_private_config_async(self):
"""从接口异步获取差异化配置(异步版本,不阻塞主循环)""" """从接口异步获取差异化配置(异步版本,不阻塞主循环)"""
if not self.read_config_from_api: if not self.read_config_from_api:
self.need_bind = False
self.bind_completed_event.set()
return return
try: try:
begin_time = time.time() begin_time = time.time()
@@ -548,15 +570,20 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).info( self.logger.bind(tag=TAG).info(
f"{time.time() - begin_time} 秒,异步获取差异化配置成功: {json.dumps(filter_sensitive_info(private_config), ensure_ascii=False)}" f"{time.time() - begin_time} 秒,异步获取差异化配置成功: {json.dumps(filter_sensitive_info(private_config), ensure_ascii=False)}"
) )
self.need_bind = False
self.bind_completed_event.set()
except DeviceNotFoundException as e: except DeviceNotFoundException as e:
self.need_bind = True self.need_bind = True
self.bind_completed_event.set() # 状态已确定,设置事件
private_config = {} private_config = {}
except DeviceBindException as e: except DeviceBindException as e:
self.need_bind = True self.need_bind = True
self.bind_code = e.bind_code self.bind_code = e.bind_code
self.bind_completed_event.set() # 状态已确定,设置事件
private_config = {} private_config = {}
except Exception as e: except Exception as e:
self.need_bind = True self.need_bind = True
self.bind_completed_event.set() # 状态已确定,设置事件
self.logger.bind(tag=TAG).error(f"异步获取差异化配置失败: {e}") self.logger.bind(tag=TAG).error(f"异步获取差异化配置失败: {e}")
private_config = {} private_config = {}
@@ -101,7 +101,7 @@ async def checkWakeupWords(conn, text):
} }
# 获取音频数据 # 获取音频数据
opus_packets = audio_to_data(response.get("file_path")) opus_packets = await audio_to_data(response.get("file_path"), use_cache=False)
# 播放唤醒词回复 # 播放唤醒词回复
conn.client_abort = False conn.client_abort = False
@@ -123,7 +123,7 @@ 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) opus_packets = await audio_to_data(file_path)
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text)) conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
conn.close_after_chat = True conn.close_after_chat = True
@@ -142,7 +142,7 @@ 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) opus_packets = await audio_to_data(music_path)
conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text)) conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text))
# 逐个播放数字 # 逐个播放数字
@@ -150,7 +150,7 @@ async def check_bind_device(conn):
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) num_packets = await audio_to_data(num_path)
conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, num_packets, None)) 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}")
@@ -162,5 +162,5 @@ 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) opus_packets = await audio_to_data(music_path)
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text)) conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
@@ -17,7 +17,12 @@ async def sendAudioMessage(conn, sentenceType, audios, text):
if sentenceType == SentenceType.FIRST: if sentenceType == SentenceType.FIRST:
# 同一句子的后续消息加入流控队列,其他情况立即发送 # 同一句子的后续消息加入流控队列,其他情况立即发送
if hasattr(conn, "audio_rate_controller") and conn.audio_rate_controller and getattr(conn, "audio_flow_control", {}).get("sentence_id") == conn.sentence_id: if (
hasattr(conn, "audio_rate_controller")
and conn.audio_rate_controller
and getattr(conn, "audio_flow_control", {}).get("sentence_id")
== conn.sentence_id
):
conn.audio_rate_controller.add_message( conn.audio_rate_controller.add_message(
lambda: send_tts_message(conn, "sentence_start", text) lambda: send_tts_message(conn, "sentence_start", text)
) )
@@ -120,7 +125,8 @@ def _get_or_create_rate_controller(conn, frame_duration, is_single_packet):
# 判断是否需要重置:单包模式且 sentence_id 变化,或者控制器不存在 # 判断是否需要重置:单包模式且 sentence_id 变化,或者控制器不存在
need_reset = ( need_reset = (
is_single_packet is_single_packet
and getattr(conn, "audio_flow_control", {}).get("sentence_id") != conn.sentence_id and getattr(conn, "audio_flow_control", {}).get("sentence_id")
!= conn.sentence_id
) or not hasattr(conn, "audio_rate_controller") ) or not hasattr(conn, "audio_rate_controller")
if need_reset: if need_reset:
@@ -138,7 +144,9 @@ def _get_or_create_rate_controller(conn, frame_duration, is_single_packet):
} }
# 启动后台发送循环 # 启动后台发送循环
_start_background_sender(conn, conn.audio_rate_controller, conn.audio_flow_control) _start_background_sender(
conn, conn.audio_rate_controller, conn.audio_flow_control
)
return conn.audio_rate_controller, conn.audio_flow_control return conn.audio_rate_controller, conn.audio_flow_control
@@ -152,6 +160,7 @@ def _start_background_sender(conn, rate_controller, flow_control):
rate_controller: 速率控制器 rate_controller: 速率控制器
flow_control: 流控状态 flow_control: 流控状态
""" """
async def send_callback(packet): async def send_callback(packet):
# 检查是否应该中止 # 检查是否应该中止
if conn.client_abort: if conn.client_abort:
@@ -165,7 +174,9 @@ def _start_background_sender(conn, rate_controller, flow_control):
rate_controller.start_sending(send_callback) rate_controller.start_sending(send_callback)
async def _send_audio_with_rate_control(conn, audio_list, rate_controller, flow_control, send_delay): async def _send_audio_with_rate_control(
conn, audio_list, rate_controller, flow_control, send_delay
):
""" """
使用 rate_controller 发送音频包 使用 rate_controller 发送音频包
@@ -235,7 +246,7 @@ async def send_tts_message(conn, state, text=None):
stop_tts_notify_voice = conn.config.get( stop_tts_notify_voice = conn.config.get(
"stop_tts_notify_voice", "config/assets/tts_notify.mp3" "stop_tts_notify_voice", "config/assets/tts_notify.mp3"
) )
audios = audio_to_data(stop_tts_notify_voice, is_opus=True) audios = await audio_to_data(stop_tts_notify_voice, is_opus=True)
await sendAudio(conn, audios) await sendAudio(conn, audios)
# 等待所有音频包发送完成 # 等待所有音频包发送完成
await _wait_for_audio_completion(conn) await _wait_for_audio_completion(conn)
@@ -118,7 +118,6 @@ class AudioRateController:
self.queue_empty_event.set() self.queue_empty_event.set()
def start_sending(self, send_audio_callback): def start_sending(self, send_audio_callback):
""" """
启动异步发送任务 启动异步发送任务
@@ -129,6 +128,7 @@ class AudioRateController:
Returns: Returns:
asyncio.Task: 发送任务 asyncio.Task: 发送任务
""" """
async def _send_loop(): async def _send_loop():
try: try:
while True: while True:
+4
View File
@@ -19,6 +19,7 @@ class CacheType(Enum):
CONFIG = "config" CONFIG = "config"
DEVICE_PROMPT = "device_prompt" DEVICE_PROMPT = "device_prompt"
VOICEPRINT_HEALTH = "voiceprint_health" # 声纹识别健康检查 VOICEPRINT_HEALTH = "voiceprint_health" # 声纹识别健康检查
AUDIO_DATA = "audio_data" # 音频数据缓存
@dataclass @dataclass
@@ -58,5 +59,8 @@ class CacheConfig:
CacheType.VOICEPRINT_HEALTH: cls( CacheType.VOICEPRINT_HEALTH: cls(
strategy=CacheStrategy.TTL, ttl=600, max_size=100 # 10分钟过期 strategy=CacheStrategy.TTL, ttl=600, max_size=100 # 10分钟过期
), ),
CacheType.AUDIO_DATA: cls(
strategy=CacheStrategy.TTL, ttl=600, max_size=100 # 10分钟过期
),
} }
return configs.get(cache_type, cls()) return configs.get(cache_type, cls())
+28 -1
View File
@@ -4,6 +4,7 @@ import json
import copy import copy
import wave import wave
import socket import socket
import asyncio
import requests import requests
import subprocess import subprocess
import numpy as np import numpy as np
@@ -268,13 +269,29 @@ def audio_to_data_stream(
pcm_to_data_stream(raw_data, is_opus, callback) pcm_to_data_stream(raw_data, is_opus, callback)
def audio_to_data(audio_file_path: str, is_opus: bool = True) -> list[bytes]: async def audio_to_data(
audio_file_path: str, is_opus: bool = True, use_cache: bool = True
) -> list[bytes]:
""" """
将音频文件转换为Opus/PCM编码的帧列表 将音频文件转换为Opus/PCM编码的帧列表
Args: Args:
audio_file_path: 音频文件路径 audio_file_path: 音频文件路径
is_opus: 是否进行Opus编码 is_opus: 是否进行Opus编码
use_cache: 是否使用缓存
""" """
from core.utils.cache.manager import cache_manager
from core.utils.cache.config import CacheType
# 生成缓存键,包含文件路径和编码类型
cache_key = f"{audio_file_path}:{is_opus}"
# 尝试从缓存获取结果
if use_cache:
cached_result = cache_manager.get(CacheType.AUDIO_DATA, cache_key)
if cached_result is not None:
return cached_result
def _sync_audio_to_data():
# 获取文件后缀名 # 获取文件后缀名
file_type = os.path.splitext(audio_file_path)[1] file_type = os.path.splitext(audio_file_path)[1]
if file_type: if file_type:
@@ -319,6 +336,16 @@ def audio_to_data(audio_file_path: str, is_opus: bool = True) -> list[bytes]:
return datas return datas
loop = asyncio.get_running_loop()
# 在单独的线程中执行同步的音频处理操作
result = await loop.run_in_executor(None, _sync_audio_to_data)
# 将结果存入缓存,使用配置中定义的TTL(10分钟)
if use_cache:
cache_manager.set(CacheType.AUDIO_DATA, cache_key, result)
return result
def audio_bytes_to_data_stream( def audio_bytes_to_data_stream(
audio_bytes, file_type, is_opus, callback: Callable[[Any], Any] audio_bytes, file_type, is_opus, callback: Callable[[Any], Any]