182 lines
7.2 KiB
Python
182 lines
7.2 KiB
Python
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, "已完成中断处理")
|