186 lines
6.9 KiB
Python
186 lines
6.9 KiB
Python
import asyncio
|
||
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
|
||
import time
|
||
|
||
|
||
async def prepare_asr_service(session, device_id, session_id):
|
||
if hasattr(session.asr_service, "device_id"):
|
||
session.asr_service.device_id = device_id
|
||
if hasattr(session.asr_service, "session_id"):
|
||
session.asr_service.session_id = session_id
|
||
await establish_asr_connection(session)
|
||
|
||
|
||
async def prepare_llm_service(session, device_id, session_id, selected_role):
|
||
if hasattr(session.llm_service, "device_id"):
|
||
session.llm_service.device_id = device_id
|
||
if hasattr(session.llm_service, "session_id"):
|
||
session.llm_service.session_id = session_id
|
||
if hasattr(session.llm_service, "connect"):
|
||
await session.llm_service.connect()
|
||
|
||
from services.conversation_history import conversation_history_manager
|
||
from config import settings
|
||
device_history = await conversation_history_manager.get_history(device_id)
|
||
if device_history:
|
||
session.llm_service.history = device_history.history[
|
||
-settings.max_conversation_history :
|
||
]
|
||
|
||
|
||
async def prepare_tts_service(session, device_id, session_id, selected_role):
|
||
session.tts_service.selected_role = selected_role
|
||
if hasattr(session.tts_service, "device_id"):
|
||
session.tts_service.device_id = device_id
|
||
if hasattr(session.tts_service, "session_id"):
|
||
session.tts_service.session_id = session_id
|
||
await establish_tts_connection(session)
|
||
|
||
|
||
async def establish_asr_connection(session: AudioSession):
|
||
|
||
try:
|
||
if session.device_id and session.session_id:
|
||
session_key = (session.device_id, session.session_id)
|
||
if interrupt_handler.is_interrupted(session_key):
|
||
session_logger.info(
|
||
session.device_id, session.session_id, "会话已中断,跳过建立ASR连接"
|
||
)
|
||
return
|
||
session.send_task = await task_manager.create_task(
|
||
send_audio_task(session),
|
||
device_id=session.device_id,
|
||
session_key=(session.device_id, session.session_id) if session.device_id and session.session_id else None,
|
||
task_type="audio_send"
|
||
)
|
||
await session.asr_service.connect()
|
||
if not session.asr_service.connected:
|
||
session_logger.error("unknown", "unknown", "语音识别 WebSocket 连接未建立")
|
||
return
|
||
session_logger.info(
|
||
"unknown", "unknown", f"语音识别连接已建立,采样率: {session.sample_rate}Hz"
|
||
)
|
||
except Exception as e:
|
||
session_logger.error(
|
||
"unknown", "unknown", f"建立语音识别 WebSocket 连接失败: {e}", exc_info=True
|
||
)
|
||
|
||
|
||
async def establish_tts_connection(session: AudioSession):
|
||
|
||
try:
|
||
if session.device_id and session.session_id:
|
||
session_key = (session.device_id, session.session_id)
|
||
if interrupt_handler.is_interrupted(session_key):
|
||
session_logger.info(
|
||
session.device_id, session.session_id, "会话已中断,跳过建立TTS连接"
|
||
)
|
||
return
|
||
|
||
await session.tts_service.connect()
|
||
|
||
if not session.tts_service.connected:
|
||
session_logger.error(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
"语音合成连接未建立",
|
||
)
|
||
return
|
||
|
||
session_logger.info(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
"语音合成连接已建立",
|
||
)
|
||
except Exception as e:
|
||
session_logger.error(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
f"建立语音合成连接失败: {e}",
|
||
)
|
||
|
||
|
||
async def send_audio_task(session: AudioSession):
|
||
try:
|
||
while True:
|
||
if session.device_id and session.session_id:
|
||
session_key = (session.device_id, session.session_id)
|
||
if interrupt_handler.is_interrupted(session_key):
|
||
session_logger.info(
|
||
session.device_id, session.session_id, "发送音频任务被中断"
|
||
)
|
||
break
|
||
audio_data = await session.send_queue.get()
|
||
if audio_data is None: # 结束信号
|
||
break
|
||
|
||
await session.asr_service.send_audio(audio_data)
|
||
|
||
except asyncio.CancelledError:
|
||
session_logger.info(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
"发送音频数据任务被取消",
|
||
)
|
||
except Exception as e:
|
||
session_logger.error(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
f"发送音频数据时出错: {str(e)}",
|
||
exc_info=True,
|
||
)
|
||
|
||
|
||
async def end_asr_session(session: AudioSession):
|
||
|
||
if session.asr_service.connected:
|
||
try:
|
||
asr_end_start_time = time.perf_counter()
|
||
await session.asr_service.send_end()
|
||
await session.asr_service.receive_results()
|
||
session.transcript = session.asr_service.transcript
|
||
asr_end_elapsed = time.perf_counter() - asr_end_start_time
|
||
session_logger.info(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
f"ASR结束延迟: 从发送结束包到获取最终结果耗时 {asr_end_elapsed:.2f} 秒",
|
||
)
|
||
except Exception as e:
|
||
session_logger.error(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
f"ASR结束处理出错: {e}",
|
||
exc_info=True,
|
||
)
|
||
finally:
|
||
if session.send_task:
|
||
session.send_task.cancel()
|
||
try:
|
||
await asyncio.wait_for(session.send_task, timeout=2.0)
|
||
except (asyncio.CancelledError, asyncio.TimeoutError):
|
||
pass
|
||
except Exception as e:
|
||
session_logger.error(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
f"取消发送任务时出错: {e}",
|
||
)
|
||
finally:
|
||
session.send_task = None
|
||
|
||
await session.asr_service.close()
|
||
session_logger.info(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
"语音识别 WebSocket 连接和 ClientSession 已关闭",
|
||
)
|
||
else:
|
||
session_logger.error(
|
||
session.device_id or "unknown",
|
||
session.session_id or "unknown",
|
||
"语音识别 WebSocket 连接未建立",
|
||
)
|