如何在FastAPI应用外部触发WebSocket消息发送?
现有代码
WebSocket连接管理器
class ConnectionManager: def __init__(self) -> None: self.connections = {} async def connect(self, user_id: str, websocket: WebSocket): await websocket.accept() self.connections[user_id] = websocket async def disconnect(self, user_id): websocket: WebSocket = self.connections[user_id] await websocket.close() del self.connections[user_id] async def send_messages(self, user_ids, message): for user_id in user_ids: websocket: WebSocket = self.connections[user_id] await websocket.send_json(message)
WebSocket路由
@router.websocket("/ws/{token}") async def ws(websocket: WebSocket, token: str, redis :Annotated [Redis, Depends(get_redis)]): user_id = redis.get(token) if user_id: redis.expire(user_id) else: raise redis_error try: manager.connect(user_id, WebSocket) except WebSocketException: manager.disconnect(user_id)
需求说明
需要存储用户WebSocket连接,当Redis Pub/Sub消息到达时向指定用户推送消息,且消息处理模块独立于FastAPI应用;曾尝试在FastAPI内用threading+asyncio实现,但会干扰应用运行,寻求外部触发WebSocket消息的可行方案。
已尝试的方案及问题
全局Redis Pub/Sub订阅失败方案
redis = Redis(redis_host, redis_port) pubsub = redis.pubsub() pubsub.subscribe("channel_signal") @router.websocket("/ws/{token}") async def ws(websocket: WebSocket, token: str): message = await pubsub.get_message(ignore_subscribe_messages=True) if message is not None: # do something try: manager.connect(user_id, WebSocket) except WebSocketException: manager.disconnect(user_id)
问题:收到Redis Pub/Sub未订阅的错误提示。
每个连接创建Redis连接方案
@router.websocket("/ws/{token}") async def ws(websocket: WebSocket, token: str): redis = Redis(redis_host, redis_port) pubsub = redis.pubsub() pubsub.subscribe("channel_signal") message = await pubsub.get_message(ignore_subscribe_messages=True) if message is not None: # do something try: manager.connect(user_id, WebSocket) except WebSocketException: manager.disconnect(user_id)
问题:每个WebSocket连接都会创建新的Redis连接,造成资源浪费,需要全局复用Redis连接的方案。
更新1:基于建议的实现及新问题
使用redispy的redis.asyncio编写了如下代码:
from fastapi import FastAPI, WebSocket, WebSocketException from v1.endpoints.user.auth import router as auth_router from v1.endpoints.signals import router as signals_router from configs.connection_config import redis_host, redis_port import redis.asyncio as aioredis import threading import asyncio import uuid app = FastAPI() app.include_router(auth_router, prefix="/users/auth", tags = ["auth"]) app.include_router(signals_router, prefix="/signals", tags = ["signals"]) class ConnectionManager: last_message = "" def __init__(self) -> None: self.connections = {} async def connect(self, user_id: str, websocket: WebSocket): await websocket.accept() self.connections[user_id] = websocket async def disconnect(self, user_id): websocket: WebSocket = self.connections[user_id] await websocket.close() del self.connections[user_id] async def send_messages(self, user_ids, message): for user_id in user_ids: websocket: WebSocket = self.connections[user_id] await websocket.send_json(message) manager = ConnectionManager() @app.websocket("/ws") async def ws(websocket: WebSocket): try: await manager.connect(str(uuid.uuid4()), websocket) except WebSocketException: await manager.disconnect(str(uuid.uuid4())) redis_client = None @app.on_event("startup") async def startup_event_connect_redis(): global redis_client redis_client = aioredis.Redis(host=redis_host, port=redis_port) def listen_to_redis(): pubsub = redis_client.pubsub() pubsub.subscribe("channel_signal") while True: message = pubsub.get_message(ignore_subscribe_messages=True) if message: print(message["data"]) @app.on_event("startup") async def startup_event_listen_redis(): threading.Thread(target=listen_to_redis, daemon=True).start()
问题:if message判断始终为真,导致无限打印"hi",事件被无限触发。
更新2:最终优化方向
虽因多用户WebSocket连接管理问题未完成完整测试,但已实现全局Redis连接并独立于FastAPI生命周期监听,最终改用官方aioredis库而非redispy版本。
解决方案
1. 全局Redis连接+异步Pub/Sub监听
使用官方aioredis库实现异步Redis连接,避免阻塞FastAPI事件循环,同时在应用启动时启动独立的异步任务监听Pub/Sub(而非线程,线程会导致异步Redis调用的兼容性问题)。
完整示例代码:
from fastapi import FastAPI, WebSocket, WebSocketException import aioredis import asyncio import json app = FastAPI() # 全局WebSocket连接管理器 class ConnectionManager: def __init__(self) -> None: self.connections = {} async def connect(self, user_id: str, websocket: WebSocket): await websocket.accept() self.connections[user_id] = websocket async def disconnect(self, user_id): # 避免KeyError,先弹出再判断 websocket = self.connections.pop(user_id, None) if websocket: await websocket.close() async def send_messages(self, user_ids, message): for user_id in user_ids: websocket = self.connections.get(user_id) if websocket: await websocket.send_json(message) manager = ConnectionManager() redis_client = None # 应用启动时初始化Redis连接并启动监听任务 @app.on_event("startup") async def startup(): global redis_client # 替换为你的Redis地址 redis_client = aioredis.from_url(f"redis://{redis_host}:{redis_port}") # 启动异步监听任务,不阻塞FastAPI事件循环 asyncio.create_task(listen_to_redis()) # 异步监听Redis Pub/Sub频道 async def listen_to_redis(): pubsub = redis_client.pubsub() await pubsub.subscribe("channel_signal") # 官方推荐的异步遍历方式,自动处理消息 async for message in pubsub.listen(): if message["type"] == "message": # 解析消息(假设消息是JSON格式,包含目标用户ID和内容) try: msg_data = json.loads(message["data"]) user_ids = msg_data.get("user_ids", []) content = msg_data.get("content", {}) await manager.send_messages(user_ids, content) except json.JSONDecodeError: # 跳过格式错误的消息 continue # WebSocket路由逻辑 @app.websocket("/ws/{token}") async def websocket_endpoint(websocket: WebSocket, token: str): # 从Redis获取用户ID user_id = await redis_client.get(token) if not user_id: await websocket.close(code=1008) return user_id = user_id.decode("utf-8") # 延长token有效期 await redis_client.expire(token, 3600) try: await manager.connect(user_id, websocket) # 保持连接,监听客户端消息(可选,若不需要可删除) while True: await websocket.receive_text() except WebSocketException: await manager.disconnect(user_id)
2. 解决get_message()无限触发问题的原因及修复
原问题中使用redispy的redis.asyncio时,get_message()在无消息时返回None,但如果误处理了订阅相关消息(即使设置了ignore_subscribe_messages=True),可能导致异常判断。改用官方推荐的async for message in pubsub.listen()方式,可自动过滤订阅/取消订阅消息,避免轮询导致的空消息判断问题。
3. 独立于FastAPI的消息触发方式
如果需要在FastAPI外部触发WebSocket消息,只需向Redis的channel_signal频道发布符合格式的JSON消息即可,示例脚本:
import aioredis import json import asyncio async def trigger_websocket_message(): # 连接Redis redis = aioredis.from_url("redis://localhost:6379") # 发布消息:指定目标用户ID和推送内容 await redis.publish( "channel_signal", json.dumps({ "user_ids": ["user_123", "user_456"], "content": {"type": "alert", "text": "您有一条新通知"} }) ) await redis.close() asyncio.run(trigger_websocket_message())
内容的提问来源于stack exchange,提问作者Anthraxff

