add banbanmini backend
This commit is contained in:
180
talkingq-url/handlers/websocket_message_handler.py
Normal file
180
talkingq-url/handlers/websocket_message_handler.py
Normal file
@@ -0,0 +1,180 @@
|
||||
import json
|
||||
import time
|
||||
import asyncio
|
||||
from fastapi import WebSocket
|
||||
from handlers.audio_packet_parser import parse_packet
|
||||
from handlers.audio_session_handler import handle_websocket_data
|
||||
from services.audio_session import audio_session_manager
|
||||
from services.interrupt_handler import interrupt_handler
|
||||
from services.task_manager import task_manager
|
||||
from utils.logger import session_logger
|
||||
from handlers.prompt_sound_handler import handle_prompt_sound_request
|
||||
from handlers.session_cleanup_handler import handle_old_session_cleanup
|
||||
|
||||
async def handle_websocket_messages(websocket: WebSocket, device_id: str):
|
||||
"""
|
||||
处理WebSocket连接中的所有消息
|
||||
|
||||
Args:
|
||||
websocket: WebSocket连接
|
||||
device_id: 设备ID
|
||||
"""
|
||||
current_active_session = None
|
||||
first_audio_received_time = None
|
||||
|
||||
while True:
|
||||
message = await websocket.receive()
|
||||
|
||||
if not websocket.client_state.name == "CONNECTED":
|
||||
session_logger.info(
|
||||
device_id,
|
||||
"connection",
|
||||
"WebSocket连接已关闭,退出消息循环",
|
||||
)
|
||||
break
|
||||
|
||||
if message["type"] == "websocket.receive":
|
||||
data = message["text"] if "text" in message else message["bytes"]
|
||||
|
||||
if isinstance(data, str):
|
||||
await handle_text_message(websocket, device_id, data)
|
||||
else: # 处理二进制音频数据
|
||||
current_active_session, first_audio_received_time = await handle_binary_message(
|
||||
websocket,
|
||||
device_id,
|
||||
data,
|
||||
current_active_session,
|
||||
first_audio_received_time
|
||||
)
|
||||
elif message["type"] == "websocket.disconnect":
|
||||
break
|
||||
|
||||
|
||||
async def handle_text_message(websocket: WebSocket, device_id: str, text_data: str):
|
||||
"""处理文本消息"""
|
||||
if text_data.startswith("REQUEST_PROMPT_SOUND:"):
|
||||
prompt_type = text_data.split(":", 1)[1]
|
||||
if device_id:
|
||||
await handle_prompt_sound_request(device_id, prompt_type)
|
||||
elif text_data == "NETWORK_RESET_ACKNOWLEDGED":
|
||||
session_logger.info(device_id, "network", f"设备 {device_id} 确认网络重置")
|
||||
elif text_data.startswith("FIRMWARE_UPDATE_STATUS:"):
|
||||
status_info = text_data.split(":", 1)[1]
|
||||
parts = status_info.split(",")
|
||||
status_dict = {}
|
||||
for part in parts:
|
||||
if "=" in part:
|
||||
key, value = part.split("=", 1)
|
||||
status_dict[key.strip()] = value.strip()
|
||||
|
||||
from services.device_update_manager import device_firmware_update_manager
|
||||
|
||||
update_status = status_dict.get("status", "updating")
|
||||
progress = float(status_dict.get("progress", "0")) if "progress" in status_dict else 0.0
|
||||
|
||||
await device_firmware_update_manager.update_firmware_progress(device_id, progress)
|
||||
|
||||
if update_status in ["success", "failed", "completed"]:
|
||||
version = status_dict.get("version", "unknown")
|
||||
await device_firmware_update_manager.update_firmware_update(
|
||||
device_id,
|
||||
firmware_version=version,
|
||||
update_status=update_status
|
||||
)
|
||||
|
||||
session_logger.info(device_id, "update", f"固件更新状态已保存: {update_status}, 进度: {progress}")
|
||||
elif text_data.startswith("FIRMWARE_VERSION:"):
|
||||
version_info = text_data.split(":", 1)[1].strip()
|
||||
session_logger.info(device_id, "update", f"收到设备固件版本: {version_info}")
|
||||
from services.device_update_manager import device_firmware_update_manager
|
||||
await device_firmware_update_manager.update_device_firmware_version(
|
||||
device_id,
|
||||
version_info
|
||||
)
|
||||
await device_firmware_update_manager.update_firmware_update(
|
||||
device_id,
|
||||
version_info,
|
||||
"success" # 收到版本信息说明设备当前状态正常
|
||||
)
|
||||
|
||||
|
||||
async def handle_binary_message(websocket: WebSocket, device_id: str, binary_data, current_active_session, first_audio_received_time):
|
||||
"""处理二进制音频消息"""
|
||||
session_key, session, packet_type, audio_data, sample_rate = await parse_packet(binary_data, websocket)
|
||||
|
||||
if session:
|
||||
session.websocket = websocket
|
||||
session_id = session_key[1]
|
||||
|
||||
if packet_type == 1: # 开始包
|
||||
current_active_session = await handle_start_packet(
|
||||
device_id, session_id, session_key, session, current_active_session
|
||||
)
|
||||
elif packet_type == 3: # 中断包
|
||||
current_active_session = await handle_interrupt_packet(
|
||||
device_id, session_id, session_key, session, current_active_session
|
||||
)
|
||||
elif packet_type == 2: # 结束包
|
||||
if current_active_session == session_key:
|
||||
current_active_session = None
|
||||
elif packet_type == 0 and first_audio_received_time is None:
|
||||
first_audio_received_time = time.perf_counter()
|
||||
session.start_time = first_audio_received_time
|
||||
session_logger.info(
|
||||
device_id, session_id, f"收到第一个音频数据包时间: {first_audio_received_time}"
|
||||
)
|
||||
|
||||
await handle_websocket_data(websocket, session_key, session, packet_type, audio_data, sample_rate)
|
||||
|
||||
return current_active_session, first_audio_received_time
|
||||
|
||||
|
||||
async def handle_start_packet(device_id, session_id, session_key, session, current_active_session):
|
||||
"""处理会话开始包"""
|
||||
if current_active_session and current_active_session[1] != session_id:
|
||||
old_session = await audio_session_manager.get_session(current_active_session)
|
||||
if old_session:
|
||||
session_logger.info(
|
||||
device_id,
|
||||
current_active_session[1],
|
||||
f"检测到新会话 {session_id},优先中断旧会话 {current_active_session[1]}",
|
||||
)
|
||||
await interrupt_handler.set_interrupt_state(current_active_session, True)
|
||||
if hasattr(old_session.llm_service, "close"):
|
||||
await old_session.llm_service.close()
|
||||
if hasattr(old_session.tts_service, "close"):
|
||||
await old_session.tts_service.close()
|
||||
await task_manager.create_task(
|
||||
handle_old_session_cleanup(old_session, current_active_session),
|
||||
device_id=device_id,
|
||||
session_key=current_active_session,
|
||||
task_type="cleanup"
|
||||
)
|
||||
|
||||
current_active_session = session_key
|
||||
await interrupt_handler.set_interrupt_state(session_key, False)
|
||||
if hasattr(session.llm_service, "closed"):
|
||||
session.llm_service.closed = False
|
||||
session_logger.info(device_id, session_id, f"设置当前活跃会话: {session_id}")
|
||||
await interrupt_handler.register_session(session_key)
|
||||
|
||||
return current_active_session
|
||||
|
||||
|
||||
async def handle_interrupt_packet(device_id, session_id, session_key, session, current_active_session):
|
||||
"""处理中断包"""
|
||||
if current_active_session == session_key:
|
||||
current_active_session = None
|
||||
session_logger.info(device_id, session_id, "收到中断包,立即设置中断状态")
|
||||
await interrupt_handler.set_interrupt_state(session_key, True)
|
||||
if hasattr(session.llm_service, "closed"):
|
||||
session.llm_service.closed = True
|
||||
if interrupt_handler.is_processing_interrupt(session_key):
|
||||
session_logger.info(device_id, session_id, "检测到连续中断,中断请求将被排队")
|
||||
await task_manager.create_task(
|
||||
interrupt_handler.handle_interrupt(session_key),
|
||||
device_id=device_id,
|
||||
session_key=session_key,
|
||||
task_type="interrupt"
|
||||
)
|
||||
return current_active_session
|
||||
Reference in New Issue
Block a user