196 lines
6.9 KiB
Python
196 lines
6.9 KiB
Python
import asyncio
|
|
from typing import Dict, Optional
|
|
from interfaces.asr import ASR
|
|
from interfaces.llm import LLM
|
|
from interfaces.tts import TTS
|
|
from services.interrupt_handler import interrupt_handler
|
|
|
|
|
|
class AudioSession:
|
|
def __init__(
|
|
self,
|
|
asr_service: Optional[ASR] = None,
|
|
llm_service: Optional[LLM] = None,
|
|
tts_service: Optional[TTS] = None,
|
|
sample_rate: int = 16000,
|
|
):
|
|
self.start_time = None
|
|
self.end_time = None
|
|
self.last_received_seq = -1
|
|
self.asr_service = asr_service
|
|
self.send_queue = None
|
|
self.send_task = None
|
|
self.transcript = ""
|
|
self.last_activity_time = asyncio.get_running_loop().time()
|
|
self.websocket = None
|
|
self.llm_service = llm_service
|
|
self.tts_service = tts_service
|
|
self.device_id = None
|
|
self.session_id = None
|
|
self.sample_rate = sample_rate
|
|
self.audio_queue = None
|
|
|
|
async def set_interrupted(self, interrupted=True):
|
|
|
|
if self.device_id and self.session_id:
|
|
session_key = (self.device_id, self.session_id)
|
|
await interrupt_handler.set_interrupt_state(session_key, interrupted)
|
|
if hasattr(self.llm_service, "closed"):
|
|
self.llm_service.closed = interrupted
|
|
from utils.logger import session_logger
|
|
|
|
session_logger.info(
|
|
self.device_id, self.session_id, f"会话中断状态已设置为: {interrupted}"
|
|
)
|
|
if interrupted:
|
|
from services.task_manager import task_manager
|
|
await task_manager.create_task(
|
|
interrupt_handler.handle_interrupt(session_key),
|
|
device_id=self.device_id,
|
|
session_key=session_key,
|
|
task_type="interrupt"
|
|
)
|
|
|
|
async def register_interrupt_handlers(self):
|
|
|
|
if not self.device_id or not self.session_id:
|
|
return
|
|
session_key = (self.device_id, self.session_id)
|
|
if hasattr(self, "audio_queue") and self.audio_queue:
|
|
await interrupt_handler.register_cleanup_handler(
|
|
session_key,
|
|
self.clear_audio_queue,
|
|
priority=interrupt_handler.PRIORITY["QUEUE_CLEANUP"],
|
|
)
|
|
if hasattr(self, "send_queue") and self.send_queue:
|
|
await interrupt_handler.register_cleanup_handler(
|
|
session_key,
|
|
self.clear_send_queue,
|
|
priority=interrupt_handler.PRIORITY["QUEUE_CLEANUP"],
|
|
)
|
|
if self.send_task:
|
|
await interrupt_handler.register_cleanup_handler(
|
|
session_key,
|
|
self.cancel_send_task,
|
|
priority=interrupt_handler.PRIORITY["TASK_CANCELLATION"],
|
|
)
|
|
|
|
async def clear_audio_queue(self):
|
|
|
|
try:
|
|
if hasattr(self, "audio_queue") and self.audio_queue:
|
|
queue_size = self.audio_queue.qsize()
|
|
cleared_count = 0
|
|
|
|
# 发送结束信号
|
|
await self.audio_queue.put((None, 0))
|
|
|
|
# 清空队列中的所有项目
|
|
while not self.audio_queue.empty():
|
|
try:
|
|
item = await self.audio_queue.get_nowait()
|
|
self.audio_queue.task_done()
|
|
cleared_count += 1
|
|
# 如果是大对象,显式删除引用
|
|
del item
|
|
except asyncio.QueueEmpty:
|
|
break
|
|
except Exception:
|
|
pass
|
|
|
|
from utils.logger import session_logger
|
|
session_logger.info(
|
|
self.device_id, self.session_id,
|
|
f"音频队列已清空,原有项目数: {queue_size},清理项目数: {cleared_count}"
|
|
)
|
|
|
|
# 清除队列引用
|
|
self.audio_queue = None
|
|
except Exception as e:
|
|
from utils.logger import session_logger
|
|
|
|
session_logger.error(
|
|
self.device_id, self.session_id, f"清空音频队列时出错: {e}"
|
|
)
|
|
|
|
async def clear_send_queue(self):
|
|
|
|
try:
|
|
if hasattr(self, "send_queue") and self.send_queue:
|
|
queue_size = self.send_queue.qsize()
|
|
cleared_count = 0
|
|
|
|
# 清空队列中的所有项目
|
|
while not self.send_queue.empty():
|
|
try:
|
|
item = self.send_queue.get_nowait()
|
|
self.send_queue.task_done()
|
|
cleared_count += 1
|
|
# 如果是大对象,显式删除引用
|
|
del item
|
|
except asyncio.QueueEmpty:
|
|
break
|
|
|
|
# 发送结束信号
|
|
await self.send_queue.put(None)
|
|
|
|
from utils.logger import session_logger
|
|
session_logger.info(
|
|
self.device_id, self.session_id,
|
|
f"发送队列已清空,原有项目数: {queue_size},清理项目数: {cleared_count}"
|
|
)
|
|
|
|
# 清除队列引用
|
|
self.send_queue = None
|
|
except Exception as e:
|
|
from utils.logger import session_logger
|
|
|
|
session_logger.error(
|
|
self.device_id, self.session_id, f"清空发送队列时出错: {e}"
|
|
)
|
|
|
|
async def cancel_send_task(self):
|
|
|
|
if self.send_task:
|
|
self.send_task.cancel()
|
|
try:
|
|
await asyncio.wait_for(self.send_task, timeout=1.0)
|
|
except (asyncio.CancelledError, asyncio.TimeoutError):
|
|
from utils.logger import session_logger
|
|
|
|
session_logger.info(self.device_id, self.session_id, "发送任务已取消")
|
|
except Exception as e:
|
|
from utils.logger import session_logger
|
|
|
|
session_logger.error(
|
|
self.device_id, self.session_id, f"取消发送任务时出错: {e}"
|
|
)
|
|
finally:
|
|
self.send_task = None
|
|
|
|
|
|
class AudioSessionManager:
|
|
def __init__(self):
|
|
self.sessions: Dict[tuple, AudioSession] = {}
|
|
self.lock = asyncio.Lock()
|
|
|
|
async def get_session(self, session_key):
|
|
async with self.lock:
|
|
return self.sessions.get(session_key)
|
|
|
|
async def set_session(self, session_key, session):
|
|
async with self.lock:
|
|
self.sessions[session_key] = session
|
|
|
|
async def remove_session(self, session_key):
|
|
async with self.lock:
|
|
if session_key in self.sessions:
|
|
del self.sessions[session_key]
|
|
|
|
async def get_all_sessions(self):
|
|
async with self.lock:
|
|
return list(self.sessions.items())
|
|
|
|
|
|
audio_session_manager = AudioSessionManager()
|