Files
banban/talkingq-url/handlers/websocket_message_handler.py
2026-03-24 15:04:36 +08:00

181 lines
7.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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