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

182 lines
7.2 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 asyncio
import time
from fastapi import WebSocket
from services.audio_session import AudioSession
from services.task_manager import task_manager
from utils.logger import session_logger
from services.interrupt_handler import interrupt_handler
async def handle_websocket_data(
websocket: WebSocket,
session_key,
session,
packet_type,
audio_data,
sample_rate=16000,
):
if not session:
return
device_id, session_id = session_key
session.last_activity_time = asyncio.get_running_loop().time()
if packet_type == 1: # 开始包
try:
await websocket.send_text("START_ACK")
session_logger.info(device_id, session_id, "已发送开始包确认")
except Exception as e:
session_logger.error(device_id, session_id, f"发送开始包确认失败: {str(e)}")
session.sample_rate = sample_rate
await handle_start_packet(session, device_id, session_id)
session_logger.info(device_id, session_id, "语音识别开始使用16K采样率")
elif packet_type == 0: # 音频数据包
await handle_audio_packet(session, audio_data)
elif packet_type == 2: # 结束包
try:
await websocket.send_text("END_ACK")
session_logger.info(device_id, session_id, "已发送结束包确认")
except Exception as e:
session_logger.error(device_id, session_id, f"发送结束包确认失败: {str(e)}")
await handle_end_packet(session, session_key)
session_logger.info(device_id, session_id, "语音识别完成")
elif packet_type == 3: # 中断包
try:
await websocket.send_text("INTERRUPT_ACK")
session_logger.info(device_id, session_id, "已发送中断包确认")
except Exception as e:
session_logger.error(device_id, session_id, f"发送中断包确认失败: {str(e)}")
await handle_interrupt_packet(session, session_key)
session_logger.info(device_id, session_id, "接收到中断包")
async def handle_start_packet(session: AudioSession, device_id: str, session_id: str):
session_key = (device_id, session_id)
await interrupt_handler.set_interrupt_state(session_key, False)
session.start_time = time.perf_counter()
session_logger.info(device_id, session_id, "会话开始")
await interrupt_handler.register_session(session_key)
from handlers.command_handler import get_device_role
selected_role = await get_device_role(device_id)
session.send_queue = asyncio.Queue()
session.device_id = device_id
session.session_id = session_id
if hasattr(session, "send_task") and session.send_task:
session.send_task.cancel()
try:
await session.send_task
except asyncio.CancelledError:
pass
session.send_task = None
from handlers.service_connection_handler import prepare_asr_service
await prepare_asr_service(session, device_id, session_id)
from handlers.service_connection_handler import (
prepare_llm_service,
prepare_tts_service,
)
async def background_warmup():
try:
llm_task = await task_manager.create_task(
prepare_llm_service(session, device_id, session_id, selected_role),
device_id=device_id,
session_key=(device_id, session_id),
task_type="service_warmup"
)
tts_task = await task_manager.create_task(
prepare_tts_service(session, device_id, session_id, selected_role),
device_id=device_id,
session_key=(device_id, session_id),
task_type="service_warmup"
)
await asyncio.gather(llm_task, tts_task)
session_logger.info(device_id, session_id, "LLM和TTS服务预热完成")
except Exception as e:
session_logger.error(
device_id, session_id, f"服务预热过程中发生错误: {e}", exc_info=True
)
await task_manager.create_task(
background_warmup(),
device_id=device_id,
session_key=(device_id, session_id),
task_type="background_warmup"
)
session_logger.info(
device_id, session_id, "ASR服务已准备就绪后台继续预热其他服务"
)
async def handle_audio_packet(session: AudioSession, audio_data: bytes):
if session.asr_service.connected:
# if len(audio_data) == 0:
# print("Received empty audio data packet, ignoring.")
# return
# print("Sending audio data to ASR service...")
# print(f"Audio data length: {len(audio_data)} bytes")
# print(audio_data)
if session.send_queue.qsize() < 100: # 设置合理的上限
await session.send_queue.put(audio_data)
else:
session_logger.warning("unknown", "unknown", "音频队列过大,丢弃部分数据")
async def handle_end_packet(session: AudioSession, session_key: tuple):
device_id, session_id = session_key
session.end_time = time.perf_counter()
session_logger.info(device_id, session_id, "会话结束")
if session.send_task:
session.send_task.cancel()
try:
await session.send_task
except asyncio.CancelledError:
pass
from handlers.service_connection_handler import end_asr_session
from handlers.transcription_handler import handle_transcription
await end_asr_session(session)
await handle_transcription(session, device_id, session_id)
if session.llm_service and hasattr(session.llm_service, "close"):
await session.llm_service.close()
if session.tts_service:
await session.tts_service.close()
await interrupt_handler.remove_session(session_key)
async def handle_interrupt_packet(session: AudioSession, session_key: tuple):
"""处理中断包"""
device_id, session_id = session_key
session_logger.info(device_id, session_id, "处理中断包")
if not interrupt_handler.is_interrupted(session_key):
await interrupt_handler.set_interrupt_state(session_key, True)
if session.llm_service and hasattr(session.llm_service, "closed"):
session.llm_service.closed = True
is_processing = interrupt_handler.is_processing_interrupt(session_key)
if is_processing:
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"
)
if session.websocket and not is_processing:
await task_manager.create_task(
interrupt_handler.notify_client_interrupt_processed(session_key, session.websocket),
device_id=device_id,
session_key=session_key,
task_type="interrupt_notification"
)
else:
session_logger.info(
device_id, session_id, "会话已处于中断状态,重复的中断请求已被记录"
)
session_logger.info(device_id, session_id, "已完成中断处理")