如何在Flask中使用flask-sock按用户/会话存储WebSocket数据
Flask-Sock实现用户级WebSocket数据隔离方案
问题背景
开发Flask Web应用时使用flask-sock实现WebSocket通信,需要将send_output_with_context中接收的WebSocket数据传入index路由的filtered_chat_memory中。当前用全局列表terminal_output存储数据,多用户同时使用时会出现数据混淆,且flask-sock没有类似flask-socketio的会话管理机制,需实现WebSocket数据的用户级隔离,且不想切换到flask-socketio。
当前简化代码
terminal_output = [] # 当前使用全局变量 def send_output_with_context(channel, ws): """Handles WebSocket communication and stores output per user""" while True: if channel.recv_ready(): output = channel.recv(1024).decode('utf-8', errors='ignore') try: ws.send(output) # Send output to the WebSocket client terminal_output.append({'role': 'assistant', 'content': output}) except: break @sock.route('/ws/ssh') def ssh_websocket(ws): """WebSocket endpoint for SSH connection""" while True: data = ws.receive() if data: send_output_with_context(channel, ws) @main_bp.route("/", methods=['POST']) @login_required def index(): """Main route where I want to access WebSocket data""" filtered_chat_memory = terminal_output + [msg for msg in chat_memory.get_history()] return jsonify({'filtered_chat_memory': filtered_chat_memory})
解决方案
方案1:基于用户ID的线程安全全局字典
用全局字典存储每个用户的WebSocket输出,键为用户ID,值为对应的数据列表,同时用线程锁保证多线程环境下的操作原子性。
修改后代码示例:
from threading import Lock from flask_login import current_user # 全局字典:键为用户ID,值为该用户的输出列表 user_terminal_output = {} # 线程锁,保证字典操作安全 output_lock = Lock() def send_output_with_context(channel, ws, user_id): """Handles WebSocket communication and stores output per user""" while True: if channel.recv_ready(): output = channel.recv(1024).decode('utf-8', errors='ignore') try: ws.send(output) # 加锁操作字典,避免多线程冲突 with output_lock: if user_id not in user_terminal_output: user_terminal_output[user_id] = [] user_terminal_output[user_id].append({'role': 'assistant', 'content': output}) except: # 连接断开时清理该用户数据(可选) with output_lock: user_terminal_output.pop(user_id, None) break @sock.route('/ws/ssh') def ssh_websocket(ws): """WebSocket endpoint for SSH connection""" user_id = current_user.id while True: data = ws.receive() if data: send_output_with_context(channel, ws, user_id) @main_bp.route("/", methods=['POST']) @login_required def index(): """Main route where I want to access WebSocket data""" user_id = current_user.id with output_lock: # 获取当前用户的输出,无数据则返回空列表 user_output = user_terminal_output.get(user_id, []) filtered_chat_memory = user_output + [msg for msg in chat_memory.get_history()] return jsonify({'filtered_chat_memory': filtered_chat_memory})
方案2:结合Flask Session与分布式存储(如Redis)
如果应用是多进程/分布式部署,全局字典无法跨进程共享,可使用Redis等分布式存储,结合Flask Session标识用户会话。
代码示例:
import redis import json from flask import session from flask_login import current_user # 初始化Redis连接 r = redis.Redis(host='localhost', port=6379, db=0) def send_output_with_context(channel, ws): """Handles WebSocket communication and stores output per user""" # 用用户ID作为Redis键的前缀,保证唯一性 redis_key = f"terminal_output:{current_user.id}" while True: if channel.recv_ready(): output = channel.recv(1024).decode('utf-8', errors='ignore') try: ws.send(output) # 将数据追加到Redis列表中 r.rpush(redis_key, json.dumps({'role': 'assistant', 'content': output})) except: # 连接断开时清理数据(可选) r.delete(redis_key) break @sock.route('/ws/ssh') def ssh_websocket(ws): """WebSocket endpoint for SSH connection""" while True: data = ws.receive() if data: send_output_with_context(channel, ws) @main_bp.route("/", methods=['POST']) @login_required def index(): """Main route where I want to access WebSocket data""" redis_key = f"terminal_output:{current_user.id}" # 从Redis获取当前用户的所有输出 user_output = [] for item in r.lrange(redis_key, 0, -1): user_output.append(json.loads(item)) filtered_chat_memory = user_output + [msg for msg in chat_memory.get_history()] return jsonify({'filtered_chat_memory': filtered_chat_memory})
方案3:将数据绑定到WebSocket连接对象
flask-sock的ws对象为每个连接独有,可直接在其上附加属性存储用户数据,同时维护用户ID到连接的映射,方便路由中获取数据。
代码示例:
from threading import Lock from flask_login import current_user # 存储用户ID到WebSocket连接的映射 user_ws_map = {} map_lock = Lock() def send_output_with_context(channel, ws): """Handles WebSocket communication and stores output per user""" # 初始化连接的输出列表 if not hasattr(ws, 'output_list'): ws.output_list = [] while True: if channel.recv_ready(): output = channel.recv(1024).decode('utf-8', errors='ignore') try: ws.send(output) ws.output_list.append({'role': 'assistant', 'content': output}) except: # 连接断开时从映射中移除 user_id = getattr(ws, 'user_id', None) if user_id: with map_lock: user_ws_map.pop(user_id, None) break @sock.route('/ws/ssh') def ssh_websocket(ws): """WebSocket endpoint for SSH connection""" user_id = current_user.id ws.user_id = user_id # 将连接存入映射 with map_lock: user_ws_map[user_id] = ws while True: data = ws.receive() if data: send_output_with_context(channel, ws) @main_bp.route("/", methods=['POST']) @login_required def index(): """Main route where I want to access WebSocket data""" user_id = current_user.id with map_lock: ws = user_ws_map.get(user_id) # 获取当前用户的输出,无连接则返回空列表 user_output = ws.output_list if ws else [] filtered_chat_memory = user_output + [msg for msg in chat_memory.get_history()] return jsonify({'filtered_chat_memory': filtered_chat_memory})
方案选择建议
- 单进程部署:优先选方案1,实现简单且性能足够;
- 多进程/分布式部署:选方案2,保证数据跨进程共享;
- 需要直接操作WebSocket连接场景:选方案3,注意连接断开时的清理,避免内存泄漏。
内容的提问来源于stack exchange,提问作者Ken 1997
相关产品推荐
相关产品推荐

