143 lines
5.8 KiB
Python
143 lines
5.8 KiB
Python
# ws_message_manager.py
|
||
import asyncio
|
||
import logging
|
||
from enum import IntEnum
|
||
from typing import Dict, Any, Optional
|
||
from fastapi import WebSocket
|
||
|
||
# 配置日志
|
||
logging.basicConfig(level=logging.INFO)
|
||
logger = logging.getLogger("ws_message_manager")
|
||
|
||
|
||
# ==================== 1. 前端→服务端 枚举(接收消息类型) ====================
|
||
# 前端发给服务端的消息类型(如音频、文本、操作指令)
|
||
class ClientMsgType(IntEnum):
|
||
AUDIO = 0x01 # 前端发送的音频二进制
|
||
TEXT = 0x02 # 前端发送的文本消息(如手动输入)
|
||
ACTION = 0x03 # 前端发送的操作指令(如开始录音、停止播放)
|
||
PING = 0x04 # 前端心跳
|
||
# 扩展:前端新增消息类型(如配置修改)
|
||
# CONFIG = 0x05
|
||
|
||
# ==================== 2. 服务端→前端 枚举(发送消息类型) ====================
|
||
# 服务端推给前端的消息类型(如文字回复、语音、动画)
|
||
class ServerMsgType(IntEnum):
|
||
TEXT = 0x01 # 服务端推送的文字回复(和前端TEXT枚举值相同但语义不同)
|
||
VOICE = 0x02 # 服务端推送的语音二进制
|
||
ANIMATION = 0x03 # 服务端推送的数字人动画
|
||
ACTION = 0x04 # 服务端推送的前端操作指令
|
||
ERROR = 0x05 # 服务端推送的错误信息
|
||
PONG = 0x06 # 服务端心跳响应(前端PING对应)
|
||
# 扩展:服务端新增推送类型(如视频流)
|
||
# VIDEO = 0x07
|
||
|
||
|
||
# 标准化消息结构
|
||
class WsMessage:
|
||
def __init__(self, type: ServerMsgType, data: Any):
|
||
self.type = type
|
||
self.data = data
|
||
|
||
def to_dict(self) -> Dict[str, Any]:
|
||
"""转换为可序列化的字典(前端解析用)"""
|
||
return {
|
||
"type": self.type.value,
|
||
"data": self.data
|
||
}
|
||
|
||
|
||
# 全局连接-队列映射(每个连接对应一个队列)
|
||
class WsQueueManager:
|
||
def __init__(self):
|
||
# 结构:{websocket对象: (消息队列, 推送协程任务)}
|
||
self.queue_map: Dict[WebSocket, tuple[asyncio.Queue, asyncio.Task]] = {}
|
||
|
||
async def create_queue(self, websocket: WebSocket) -> asyncio.Queue:
|
||
"""为新连接创建专属队列,并启动推送协程"""
|
||
if websocket in self.queue_map:
|
||
logger.warning(f"连接已存在队列: {id(websocket)}")
|
||
return self.queue_map[websocket][0]
|
||
|
||
# 创建队列(最大100条,避免堆积)
|
||
msg_queue = asyncio.Queue(maxsize=100)
|
||
# 启动推送协程(单协程推送,避免并发send)
|
||
push_task = asyncio.create_task(self._push_worker(websocket, msg_queue))
|
||
self.queue_map[websocket] = (msg_queue, push_task)
|
||
logger.info(f"为连接 {id(websocket)} 创建队列,推送协程启动")
|
||
return msg_queue
|
||
|
||
async def _push_worker(self, websocket: WebSocket, msg_queue: asyncio.Queue):
|
||
"""推送协程:从队列读取消息并发送给前端"""
|
||
try:
|
||
while True:
|
||
# 阻塞获取队列中的消息
|
||
message: Optional[WsMessage] = await msg_queue.get()
|
||
|
||
if message is None: # 连接关闭标记
|
||
logger.info(f"连接 {id(websocket)} 推送协程收到关闭标记")
|
||
break
|
||
|
||
# 根据消息类型选择发送方式(文本/二进制)
|
||
try:
|
||
if message.type == ServerMsgType.VOICE:
|
||
# 语音是二进制数据,直接发送bytes
|
||
await websocket.send_bytes(message.data)
|
||
else:
|
||
# 其他类型序列化为JSON字符串
|
||
await websocket.send_json(message.to_dict())
|
||
except Exception as e:
|
||
logger.error(f"推送消息失败: {e},消息类型: {message.type}")
|
||
finally:
|
||
msg_queue.task_done() # 标记队列任务完成
|
||
except asyncio.CancelledError:
|
||
logger.info(f"连接 {id(websocket)} 推送协程被取消")
|
||
except Exception as e:
|
||
logger.error(f"推送协程异常: {e}")
|
||
|
||
async def send_message(self, websocket: WebSocket, msg_type: ServerMsgType, data: Any):
|
||
"""对外暴露的发送方法:封装消息类型并入队"""
|
||
|
||
if websocket not in self.queue_map:
|
||
raise ValueError(f"连接 {id(websocket)} 未创建队列")
|
||
|
||
msg_queue, _ = self.queue_map[websocket]
|
||
# 封装标准化消息
|
||
message = WsMessage(type=msg_type, data=data)
|
||
|
||
try:
|
||
# 非阻塞入队(避免处理协程被阻塞)
|
||
await asyncio.wait_for(msg_queue.put(message), timeout=1.0)
|
||
except asyncio.QueueFull:
|
||
logger.error(f"队列已满,丢弃消息(类型: {msg_type})")
|
||
except asyncio.TimeoutError:
|
||
logger.error(f"消息入队超时(类型: {msg_type})")
|
||
|
||
async def close_queue(self, websocket: WebSocket):
|
||
"""关闭连接对应的队列和推送协程"""
|
||
if websocket not in self.queue_map:
|
||
return
|
||
|
||
msg_queue, push_task = self.queue_map.pop(websocket)
|
||
# 发送关闭标记,终止推送协程
|
||
try:
|
||
await msg_queue.put(None)
|
||
except:
|
||
pass
|
||
# 取消推送协程并等待结束
|
||
push_task.cancel()
|
||
try:
|
||
await push_task
|
||
except asyncio.CancelledError:
|
||
pass
|
||
# 清空队列
|
||
while not msg_queue.empty():
|
||
try:
|
||
msg_queue.get_nowait()
|
||
except asyncio.QueueEmpty:
|
||
break
|
||
logger.info(f"连接 {id(websocket)} 队列已清理")
|
||
|
||
|
||
# 创建全局单例(所有连接共用一个管理器)
|
||
ws_queue_manager = WsQueueManager() |