如何让单个Redis客户端等待其余客户端响应?集群WebSocket连接数统计
这问题我在实际生产环境中处理过类似的场景,给你分享几个落地性强的实现思路,分两种不同的场景来适配你的需求:
方案一:实时全量收集(适合对数据实时性要求高的场景)
这个方案的核心是利用Redis的Pub/Sub实现节点间的实时通信,收到健康检查请求的节点主动触发全量收集:
- 节点注册与心跳:每个WebSocket+HTTP服务器启动时,生成唯一的节点ID(比如容器ID、UUID),将其加入Redis的Set集合
ws_active_nodes,同时设置一个心跳键ws_heartbeat:{node_id}并设置过期时间(比如15秒),每隔10秒刷新一次心跳。这样可以随时知道当前活跃的节点总数。 - 触发收集流程:当某个节点收到
GET /health请求时:- 生成唯一的请求ID(比如
health_req:{uuid}),创建对应的Redis Hash键health_collect_data:{req_id}并设置5秒过期(防止残留垃圾数据)。 - 通过Redis的Pub/Sub频道
health_collect_cmd发布消息,内容携带这个请求ID。 - 等待所有活跃节点上报数据:可以通过轮询Hash的字段数,或者开启Redis的键空间通知(开启
notify-keyspace-events KEA)监听Hash的新增事件,这样不用轮询更高效。
- 生成唯一的请求ID(比如
- 节点响应收集:所有节点都订阅
health_collect_cmd频道,收到消息后:- 统计自己当前的WebSocket连接数。
- 将自己的节点ID和连接数存入对应的
health_collect_data:{req_id}Hash中。
- 汇总与返回:发起请求的节点等到所有活跃节点都上报(或等待超时,比如3秒),将Hash中的所有值求和得到总连接数,返回给请求方,最后删除临时Hash键。
方案二:定时上报+实时汇总(适合实现简单、对延迟容忍的场景)
如果你的健康检查不需要绝对实时的连接数,这个方案更简单,不需要实时通信:
- 定时上报连接数:每个WebSocket服务器每隔固定时间(比如10秒),将自己的连接数写入Redis的Hash键
ws_connection_counts,键为节点ID,值为连接数,同时刷新自己的心跳键。 - 健康检查时汇总:当某个节点收到
GET /health请求时:- 先清理超时节点:扫描
ws_active_nodes中的节点,删除心跳过期的节点,并从ws_connection_counts中移除对应数据。 - 对
ws_connection_counts中的所有值求和,得到总连接数并返回。
- 先清理超时节点:扫描
关键注意事项
- 超时处理:不管用哪个方案,发起汇总的节点一定要设置超时时间,避免某个节点挂掉导致请求一直阻塞(比如最多等待3秒,超时后用已收集到的数据返回,同时可以在响应中标记部分节点未上报)。
- 节点下线清理:一定要通过心跳机制清理失效节点,防止无效数据影响统计结果。
- Redis性能:如果节点数量较多(比如上百个),方案一的Pub/Sub可能会有短暂的消息风暴,但Redis的Pub/Sub性能足够支撑;方案二的定时上报要注意不要集中在同一时间点上报,可以给每个节点加随机延迟,避免Redis压力集中。
简单代码示例(方案一伪代码)
发起请求节点的处理逻辑
import redis import uuid import time r = redis.Redis(host="your_redis_host") NODE_ID = "unique-node-id-xxx" def get_current_ws_connections(): # 这里替换成你获取当前节点WebSocket连接数的逻辑 return len(active_ws_connections) def handle_health(): req_id = f"health_req:{uuid.uuid4().hex}" temp_hash_key = f"health_collect_data:{req_id}" # 设置临时键过期时间 r.expire(temp_hash_key, 5) # 获取当前活跃节点数 active_nodes = r.smembers("ws_active_nodes") expected_count = len(active_nodes) if active_nodes else 0 # 发布收集命令 r.publish("health_collect_cmd", req_id) # 等待上报,最多3秒超时 collected_count = 0 start_time = time.time() while collected_count < expected_count and time.time() - start_time < 3: collected_count = r.hlen(temp_hash_key) time.sleep(0.1) # 汇总数据 total_connections = sum(int(v) for v in r.hvals(temp_hash_key)) # 清理临时数据 r.delete(temp_hash_key) return { "status": "healthy", "total_ws_connections": total_connections, "reported_nodes": collected_count, "total_nodes": expected_count }
其他节点的监听逻辑
def listen_for_collect_cmds(): pubsub = r.pubsub() pubsub.subscribe("health_collect_cmd") for msg in pubsub.listen(): if msg["type"] == "message": req_id = msg["data"].decode() temp_hash_key = f"health_collect_data:{req_id}" # 上报自己的连接数 r.hset(temp_hash_key, NODE_ID, get_current_ws_connections())
内容的提问来源于stack exchange,提问作者sp00m
相关产品推荐
相关产品推荐

