add banbanmini backend
This commit is contained in:
181
talkingq-url/handlers/audio_session_handler.py
Normal file
181
talkingq-url/handlers/audio_session_handler.py
Normal file
@@ -0,0 +1,181 @@
|
||||
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, "已完成中断处理")
|
||||
Reference in New Issue
Block a user