67 lines
3.0 KiB
Python
67 lines
3.0 KiB
Python
# session_manager.py(最终版)
|
|
import asyncio
|
|
import logging
|
|
from dataclasses import dataclass, field
|
|
from fastapi import WebSocket
|
|
from typing import Optional, List, Any
|
|
from asyncio import Queue
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
@dataclass
|
|
class WsSession:
|
|
"""WebSocket 会话类:管理单个连接的所有上下文(无用户标识,仅做资源隔离)"""
|
|
websocket: WebSocket # 前端 WS 连接
|
|
session_id: str # 会话ID(用 conn_id 即可,无需用户ID)
|
|
is_closed: bool = False # 会话是否关闭
|
|
tasks: List[asyncio.Task] = field(default_factory=list) # 转发协程列表
|
|
# ASR/TTS 客户端连接(每个会话独立创建,天然隔离)
|
|
asr_client: Optional[Any] = None # ASR 客户端任务(存储协程任务)
|
|
tts_client: Optional[Any] = None # TTS 客户端实例(无用户态)
|
|
# 新增:ASR 依赖的核心属性
|
|
audio_queue: Queue = field(default_factory=lambda: Queue(maxsize=30)) # 音频队列(防积压)
|
|
llm_queue: Queue = field(default_factory=Queue) # LLM 队列(转发ASR最终结果)
|
|
asr_ws_task: Optional[asyncio.Task] = None # ASR 核心协程任务
|
|
|
|
async def create_session(websocket: WebSocket) -> WsSession:
|
|
"""创建会话(仅初始化资源,无用户相关逻辑)"""
|
|
conn_id = str(id(websocket))
|
|
session = WsSession(
|
|
websocket=websocket,
|
|
session_id=conn_id # 会话ID = 连接ID,无需用户标识
|
|
)
|
|
logger.info(f"会话 [{session.session_id}] 创建成功")
|
|
return session
|
|
|
|
async def close_session(session: WsSession):
|
|
"""关闭会话:清理所有协程和ASR/TTS连接(核心:资源释放)"""
|
|
session.is_closed = True
|
|
# 1. 取消所有转发协程
|
|
for task in session.tasks:
|
|
if not task.done():
|
|
task.cancel()
|
|
try:
|
|
await task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
# 2. 单独取消ASR核心任务(新增)
|
|
if session.asr_ws_task and not session.asr_ws_task.done():
|
|
session.asr_ws_task.cancel()
|
|
try:
|
|
await session.asr_ws_task
|
|
except asyncio.CancelledError:
|
|
logger.info(f"会话 [{session.session_id}] ASR核心协程已取消")
|
|
# 3. 关闭ASR/TTS客户端(适配:asr_client 是任务,无需close,仅日志提示)
|
|
if session.asr_client:
|
|
logger.info(f"会话 [{session.session_id}] ASR 客户端任务已清理")
|
|
if session.tts_client and hasattr(session.tts_client, 'closed') and not session.tts_client.closed:
|
|
await session.tts_client.close()
|
|
logger.info(f"会话 [{session.session_id}] TTS 客户端已关闭")
|
|
# 4. 清空音频队列(防止内存泄漏)
|
|
while not session.audio_queue.empty():
|
|
try:
|
|
session.audio_queue.get_nowait()
|
|
session.audio_queue.task_done()
|
|
except asyncio.QueueEmpty:
|
|
break
|
|
logger.info(f"会话 [{session.session_id}] 已完全清理") |