FastAPI中如何获取WebSocket连接的客户端IP并改造连接管理器存储结构
获取WebSocket客户端IP及改造连接管理器实现方案
1. 获取客户端IP地址
FastAPI的WebSocket实例内置client属性,为(主机地址, 端口)格式的元组,直接通过websocket.client.host即可获取原生客户端IP。
如果服务部署在Nginx等反向代理之后,需要先配置代理透传客户端IP,再从请求头中获取:
# 反向代理场景获取真实IP,第二个参数为兜底的直连IP client_ip = websocket.headers.get("X-Forwarded-For", websocket.client.host)
2. 改造ConnectionManager为IP键值存储
方案1:同IP仅允许单个连接(旧连接会被新连接覆盖/断开)
改造后的连接管理器代码:
from typing import Dict, WebSocket from fastapi import WebSocket, WebSocketDisconnect class ConnectionManager: def __init__(self): # 键为客户端IP,值为对应WebSocket实例 self.active_connections: Dict[str, WebSocket] = {} async def connect(self, websocket: WebSocket, client_ip: str): await websocket.accept() # 可选:同IP新连接接入时先断开旧连接 if client_ip in self.active_connections: await self.active_connections[client_ip].close() self.active_connections[client_ip] = websocket def disconnect(self, client_ip: str): if client_ip in self.active_connections: del self.active_connections[client_ip] async def send_personal_message(self, message: str, client_ip: str): if client_ip in self.active_connections: await self.active_connections[client_ip].send_text(message) async def broadcast(self, message: str): for connection in self.active_connections.values(): await connection.send_text(message) manager = ConnectionManager()
改造后的WebSocket路由代码(新增异常处理和连接保持逻辑,避免连接立即断开):
@router.websocket("/abcd") async def websocket_endpoint(websocket: WebSocket): client_ip = websocket.client.host # 反向代理场景替换为下方代码 # client_ip = websocket.headers.get("X-Forwarded-For", websocket.client.host) try: await manager.connect(websocket, client_ip) # 保持连接监听客户端消息,可在此处添加业务逻辑 while True: receive_data = await websocket.receive_text() # 业务逻辑处理 except WebSocketDisconnect: manager.disconnect(client_ip)
方案2:同IP允许多个连接
改造后的连接管理器代码:
from typing import Dict, List, WebSocket from fastapi import WebSocket, WebSocketDisconnect class ConnectionManager: def __init__(self): # 键为客户端IP,值为对应WebSocket实例列表 self.active_connections: Dict[str, List[WebSocket]] = {} async def connect(self, websocket: WebSocket, client_ip: str): await websocket.accept() if client_ip not in self.active_connections: self.active_connections[client_ip] = [] self.active_connections[client_ip].append(websocket) def disconnect(self, websocket: WebSocket, client_ip: str): if client_ip in self.active_connections: if websocket in self.active_connections[client_ip]: self.active_connections[client_ip].remove(websocket) # 该IP无活跃连接时删除键释放空间 if not self.active_connections[client_ip]: del self.active_connections[client_ip] async def send_personal_message(self, message: str, client_ip: str): if client_ip in self.active_connections: for conn in self.active_connections[client_ip]: await conn.send_text(message) async def broadcast(self, message: str): for conns in self.active_connections.values(): for conn in conns: await conn.send_text(message) manager = ConnectionManager()
对应改造后的路由代码:
@router.websocket("/abcd") async def websocket_endpoint(websocket: WebSocket): client_ip = websocket.client.host try: await manager.connect(websocket, client_ip) while True: receive_data = await websocket.receive_text() # 业务逻辑处理 except WebSocketDisconnect: manager.disconnect(websocket, client_ip)
内容的提问来源于stack exchange,提问作者Phong Phạm
相关产品推荐
相关产品推荐

