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

186 lines
6.9 KiB
Python
Raw Permalink 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
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 连接未建立",
)