Files
aistream-test/python/session_manager.py
2025-12-01 03:42:34 +08:00

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}] 已完全清理")