394 lines
19 KiB
Python
394 lines
19 KiB
Python
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 handlers.audio_file_handler import upload_message_audio
|
||
from banban.service.device_audio_cache import device_audio_cache_service
|
||
from banban.service.device_voice_archive import device_voice_archive_service
|
||
from banban.service.pending_voice_message import pending_voice_message_service
|
||
from services.audio_session import audio_session_manager
|
||
from services.interrupt_handler import interrupt_handler
|
||
from services.task_manager import task_manager
|
||
from services.connection_manager import connection_manager
|
||
from services.device_target_cache import device_target_cache
|
||
from services.offline_audio_cache import offline_audio_cache
|
||
from services.target_audio_cache import target_audio_cache
|
||
from services.card_service import card_service
|
||
from services.device_target_cache import VOICE_TARGET_KIND_PARENT
|
||
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
|
||
from config import settings
|
||
from banban.service.im import im_service as im_conversation_service
|
||
from utils.audio_format import detect_audio_format, wrap_pcm_as_wav
|
||
from fastapi import HTTPException
|
||
|
||
|
||
def prepare_message_audio(audio_data: bytes) -> tuple[bytes, str, str, str]:
|
||
source_format = detect_audio_format(audio_data)
|
||
if source_format == "wav":
|
||
return audio_data, "audio/wav", "wav", source_format
|
||
if source_format == "mp3":
|
||
return audio_data, "audio/mpeg", "mp3", source_format
|
||
return wrap_pcm_as_wav(audio_data), "audio/wav", "wav", source_format
|
||
|
||
|
||
async def handle_websocket_messages(websocket: WebSocket, device_id: str, serial_number: 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,
|
||
serial_number,
|
||
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("REGISTER_TARGET_DEVICE:"):
|
||
card_uuid = text_data.split(":", 1)[1].strip()
|
||
|
||
# 检查卡片是否存在
|
||
existing_card = await card_service.get_card_by_uuid(card_uuid)
|
||
|
||
if existing_card:
|
||
# 卡片已存在,使用卡片绑定的设备ID作为目标设备ID
|
||
target_device_id = existing_card.device_id
|
||
session_logger.info(device_id, "card", f"卡片已存在,绑定的设备ID: {target_device_id}")
|
||
else:
|
||
# 卡片不存在,创建新卡片并绑定到当前设备
|
||
new_card = await card_service.activate_card(card_uuid, device_id)
|
||
session_logger.info(device_id, "card", f"新卡片{card_uuid}已创建并激活,绑定到设备: {device_id}")
|
||
return
|
||
|
||
# 设置目标设备
|
||
await device_target_cache.set_target(device_id, target_device_id)
|
||
|
||
# 检查目标设备是否在线
|
||
target_websocket = await connection_manager.get_connection(target_device_id)
|
||
if target_websocket and target_websocket.client_state.name == "CONNECTED":
|
||
sound_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/welcome.mp3"
|
||
await websocket.send_text(f"TARGET_DEVICE_REGISTERED_URL:{sound_url}")
|
||
session_logger.info(device_id, "target", f"成功注册目标设备: {target_device_id}")
|
||
else:
|
||
sound_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/offline.mp3"
|
||
await websocket.send_text(f"TARGET_DEVICE_REGISTERED_URL:{sound_url}")
|
||
session_logger.warning(device_id, "target", f"目标设备 {target_device_id} 不在线")
|
||
elif 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, serial_number: 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: # 开始包
|
||
voice_target = await device_target_cache.get_voice_target(device_id)
|
||
if voice_target:
|
||
await target_audio_cache.clear_audio_data(voice_target.audio_cache_key)
|
||
return None, None
|
||
current_active_session = await handle_start_packet(
|
||
device_id, session_id, session_key, session, current_active_session
|
||
)
|
||
elif packet_type == 3: # 中断包
|
||
voice_target = await device_target_cache.get_voice_target(device_id)
|
||
if voice_target:
|
||
return None, None
|
||
current_active_session = await handle_interrupt_packet(
|
||
device_id, session_id, session_key, session, current_active_session
|
||
)
|
||
elif packet_type == 2: # 结束包
|
||
voice_target = await device_target_cache.get_voice_target(device_id)
|
||
if voice_target:
|
||
session_logger.info(device_id, session_id, f"收到结束包,开始处理缓存音频数据")
|
||
if voice_target.kind == VOICE_TARGET_KIND_PARENT:
|
||
await process_parent_leave_message(device_id, voice_target.audio_cache_key)
|
||
else:
|
||
await process_cached_audio(device_id, voice_target.target_device_id, serial_number)
|
||
await device_target_cache.remove_target(device_id)
|
||
return None, None
|
||
if current_active_session == session_key:
|
||
current_active_session = None
|
||
elif packet_type == 4: # 发送给目标设备的音频包
|
||
await handle_target_audio_packet(device_id, audio_data)
|
||
return None, 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
|
||
|
||
|
||
async def handle_target_audio_packet(device_id: str, audio_data: bytes):
|
||
"""处理发送给目标设备的音频包"""
|
||
try:
|
||
voice_target = await device_target_cache.get_voice_target(device_id)
|
||
|
||
if not voice_target:
|
||
session_logger.warning(device_id, "target", "未设置目标设备,无法发送音频")
|
||
return
|
||
|
||
await target_audio_cache.add_audio_data(voice_target.audio_cache_key, audio_data)
|
||
if voice_target.kind == VOICE_TARGET_KIND_PARENT:
|
||
session_logger.info(device_id, "target", "已缓存发给家长的留言音频数据")
|
||
else:
|
||
session_logger.info(device_id, "target", f"已缓存音频数据到目标设备 {voice_target.target_device_id}")
|
||
|
||
except Exception as e:
|
||
session_logger.error(device_id, "target", f"处理目标音频包时出错: {e}", exc_info=True)
|
||
|
||
|
||
async def process_parent_leave_message(device_id: str, audio_cache_key: str):
|
||
"""处理设备发给家长的留言音频并写入家长会话。"""
|
||
try:
|
||
cached_audio = await target_audio_cache.get_audio_data(audio_cache_key)
|
||
if not cached_audio:
|
||
session_logger.info(device_id, "parent", "发给家长的留言没有缓存音频数据")
|
||
return
|
||
|
||
archive_audio, media_mime_type, extension, source_format = prepare_message_audio(cached_audio)
|
||
stored_audio = await upload_message_audio(
|
||
archive_audio,
|
||
device_id,
|
||
content_type=media_mime_type,
|
||
extension=extension,
|
||
)
|
||
await im_conversation_service.create_device_parent_leave_message(
|
||
device_id=device_id,
|
||
media_file_key=stored_audio.file_key,
|
||
media_mime_type=media_mime_type,
|
||
media_size_bytes=len(archive_audio),
|
||
ext_json={
|
||
"source": "device_ws_parent_leave_message",
|
||
"storage": "cos",
|
||
"source_format": source_format,
|
||
"archive_format": extension,
|
||
},
|
||
)
|
||
session_logger.info(
|
||
device_id,
|
||
"parent",
|
||
f"发给家长的留言已上传COS并写入家长会话: {stored_audio.file_key}",
|
||
)
|
||
|
||
websocket = await connection_manager.get_connection(device_id)
|
||
if websocket and websocket.client_state.name == "CONNECTED":
|
||
success_audio_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/message_ok.mp3"
|
||
await websocket.send_text(f"PROMPT_SOUND_URL:{success_audio_url}")
|
||
session_logger.info(device_id, "device", f"留言成功提示音已发送给设备 {device_id}")
|
||
else:
|
||
session_logger.warning(device_id, "device", f"设备 {device_id} 不在线,暂不发送留言成功提示音")
|
||
except Exception as e:
|
||
if isinstance(e, HTTPException):
|
||
session_logger.error(device_id, "parent", f"发给家长的留言保存失败,返回HTTPException: {e}")
|
||
websocket = await connection_manager.get_connection(device_id)
|
||
if e.status_code == 404 or e.status_code == 400:
|
||
if websocket and websocket.client_state.name == "CONNECTED":
|
||
fail_audio_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/save_audio_fail.mp3"
|
||
await websocket.send_text(f"PROMPT_SOUND_URL:{fail_audio_url}")
|
||
session_logger.info(device_id, "device", f"留言失败提示音已发送给设备 {device_id}")
|
||
else:
|
||
session_logger.warning(device_id, "device", f"设备 {device_id} 不在线,暂不发送留言失败提示音")
|
||
else:
|
||
session_logger.error(device_id, "parent", f"处理家长留言音频时出错: {e}", exc_info=True)
|
||
finally:
|
||
await target_audio_cache.clear_audio_data(audio_cache_key)
|
||
|
||
|
||
async def process_cached_audio(device_id: str, target_device_id: str, serial_number: str):
|
||
"""处理缓存的音频数据并发送音频URL"""
|
||
try:
|
||
cached_audio = await target_audio_cache.get_audio_data(target_device_id)
|
||
if not cached_audio:
|
||
session_logger.info(device_id, "target", f"目标设备 {target_device_id} 没有缓存的音频数据")
|
||
return
|
||
|
||
audio_url, local_audio_path = await device_audio_cache_service.save_device_audio(
|
||
cached_audio,
|
||
device_id=target_device_id,
|
||
)
|
||
archive_result = await device_voice_archive_service.archive_peer_voice_message(
|
||
sender_device_id=device_id,
|
||
receiver_device_id=target_device_id,
|
||
local_audio_path=str(local_audio_path),
|
||
)
|
||
if archive_result is not None:
|
||
await pending_voice_message_service.add_pending_message(
|
||
target_device_id=target_device_id,
|
||
sender_device_id=device_id,
|
||
im_message_id=archive_result.message_id,
|
||
media_file_key=archive_result.media_file_key,
|
||
audio_url=audio_url,
|
||
source="device_peer_voice",
|
||
)
|
||
session_logger.info(
|
||
device_id,
|
||
"target",
|
||
f"设备留言已加入待收听队列: target_device_id={target_device_id}, audio_url={audio_url}",
|
||
)
|
||
else:
|
||
await offline_audio_cache.add_audio_url(target_device_id, audio_url)
|
||
session_logger.warning(
|
||
device_id,
|
||
"target",
|
||
"设备留言已写入内存待收听队列,但归档失败,重启后无法恢复这条留言",
|
||
)
|
||
|
||
websocket = await connection_manager.get_connection(device_id)
|
||
if websocket and websocket.client_state.name == "CONNECTED":
|
||
success_audio_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/message_ok.mp3"
|
||
await websocket.send_text(f"PROMPT_SOUND_URL:{success_audio_url}")
|
||
session_logger.info(device_id, "device", f"留言已收到音频URL给设备 {device_id}")
|
||
else:
|
||
session_logger.warning(device_id, "device", f"设备 {device_id} 不在线,暂不发送留言已收到音频URL")
|
||
except Exception as e:
|
||
if isinstance(e, HTTPException):
|
||
session_logger.error(device_id, "target", f"保存音频文件时出错,返回HTTPException: {e}")
|
||
websocket = await connection_manager.get_connection(device_id)
|
||
if e.status_code == 404 or e.status_code == 400:
|
||
if websocket and websocket.client_state.name == "CONNECTED":
|
||
success_audio_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/save_audio_fail.mp3"
|
||
await websocket.send_text(f"PROMPT_SOUND_URL:{success_audio_url}")
|
||
session_logger.info(device_id, "device", f"留言已收到音频URL给设备 {device_id}")
|
||
else:
|
||
session_logger.warning(device_id, "device", f"设备 {device_id} 不在线,暂不发送留言已收到音频URL")
|
||
else:
|
||
session_logger.error(device_id, "target", f"处理缓存音频时出错: {e}", exc_info=True)
|
||
finally:
|
||
# 清除缓存
|
||
await target_audio_cache.clear_audio_data(target_device_id)
|