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
+51 -24
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,30 +267,43 @@ class ConnectionHandler:
f"保存记忆后关闭连接失败: {close_error}" f"保存记忆后关闭连接失败: {close_error}"
) )
async def _discard_message_with_bind_prompt(self):
"""丢弃消息并检查是否需要播放绑定提示"""
current_time = time.time()
# 检查是否需要播放绑定提示
if current_time - self.last_bind_prompt_time >= self.bind_prompt_interval:
self.last_bind_prompt_time = current_time
# 复用现有的绑定提示逻辑
from core.handle.receiveAudioHandle import check_bind_device
asyncio.create_task(check_bind_device(self))
async def _route_message(self, message): 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): if isinstance(message, str):
await handleTextMessage(self, message) await handleTextMessage(self, message)
elif isinstance(message, bytes): elif isinstance(message, bytes):
if self.vad is None or self.asr is None: if self.vad is None or self.asr is None:
return return
# 未绑定设备直接丢弃所有音频,不进行ASR处理
if self.need_bind:
current_time = time.time()
# 检查是否需要播放绑定提示
if (
current_time - self.last_bind_prompt_time
>= self.bind_prompt_interval
):
self.last_bind_prompt_time = current_time
# 复用现有的绑定提示逻辑
from core.handle.receiveAudioHandle import check_bind_device
asyncio.create_task(check_bind_device(self))
# 直接丢弃音频,不进行ASR处理
return
# 处理来自MQTT网关的音频包 # 处理来自MQTT网关的音频包
if self.conn_from_mqtt_gateway and len(message) >= 16: if self.conn_from_mqtt_gateway and len(message) >= 16:
handled = await self._process_mqtt_audio_message(message) handled = await self._process_mqtt_audio_message(message)
@@ -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())
+62 -35
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,56 +269,82 @@ 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
file_type = os.path.splitext(audio_file_path)[1] from core.utils.cache.config import CacheType
if file_type:
file_type = file_type.lstrip(".")
# 读取音频文件,-nostdin 参数:不要从标准输入读取数据,否则FFmpeg会阻塞
audio = AudioSegment.from_file(
audio_file_path, format=file_type, parameters=["-nostdin"]
)
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配) # 生成缓存键,包含文件路径和编码类型
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2) cache_key = f"{audio_file_path}:{is_opus}"
# 获取原始PCM数据(16位小端) # 尝试从缓存获取结果
raw_data = audio.raw_data if use_cache:
cached_result = cache_manager.get(CacheType.AUDIO_DATA, cache_key)
if cached_result is not None:
return cached_result
# 初始化Opus编码器 def _sync_audio_to_data():
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO) # 获取文件后缀名
file_type = os.path.splitext(audio_file_path)[1]
if file_type:
file_type = file_type.lstrip(".")
# 读取音频文件,-nostdin 参数:不要从标准输入读取数据,否则FFmpeg会阻塞
audio = AudioSegment.from_file(
audio_file_path, format=file_type, parameters=["-nostdin"]
)
# 编码参数 # 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
frame_duration = 60 # 60ms per frame audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
datas = [] # 获取原始PCM数据(16位小端)
# 按帧处理所有音频数据(包括最后一帧可能补零) raw_data = audio.raw_data
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
# 获取当前帧的二进制数据
chunk = raw_data[i : i + frame_size * 2]
# 如果最后一帧不足,补零 # 初始化Opus编码器
if len(chunk) < frame_size * 2: encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
chunk += b"\x00" * (frame_size * 2 - len(chunk))
if is_opus: # 编码参数
# 转换为numpy数组处理 frame_duration = 60 # 60ms per frame
np_frame = np.frombuffer(chunk, dtype=np.int16) frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
# 编码Opus数据
frame_data = encoder.encode(np_frame.tobytes(), frame_size)
else:
frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
datas.append(frame_data) datas = []
# 按帧处理所有音频数据(包括最后一帧可能补零)
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
# 获取当前帧的二进制数据
chunk = raw_data[i : i + frame_size * 2]
return datas # 如果最后一帧不足,补零
if len(chunk) < frame_size * 2:
chunk += b"\x00" * (frame_size * 2 - len(chunk))
if is_opus:
# 转换为numpy数组处理
np_frame = np.frombuffer(chunk, dtype=np.int16)
# 编码Opus数据
frame_data = encoder.encode(np_frame.tobytes(), frame_size)
else:
frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
datas.append(frame_data)
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(