""" Redis Session 存储 用于存储 refresh_token 和 Session 信息 """ from typing import Optional from datetime import datetime, timedelta import json import redis from redis import Redis class RedisConfig: """Redis 配置""" HOST: str = "localhost" PORT: int = 6379 DB: int = 0 PASSWORD: Optional[str] = None DECODE_RESPONSES: bool = True class SessionStore: """Session 存储服务""" def __init__(self, redis_client: Optional[Redis] = None, config: Optional[RedisConfig] = None): """ 初始化 Session 存储 Args: redis_client: Redis 客户端(可选,用于注入) config: Redis 配置(可选) """ if redis_client: self.redis = redis_client else: cfg = config or RedisConfig() self.redis = redis.Redis( host=cfg.HOST, port=cfg.PORT, db=cfg.DB, password=cfg.PASSWORD, decode_responses=cfg.DECODE_RESPONSES, ) def _session_key(self, session_id: str) -> str: """生成 Session key""" return f"session:{session_id}" def _refresh_token_key(self, session_id: str) -> str: """生成 refresh_token key""" return f"refresh_token:{session_id}" def _user_sessions_key(self, user_id: str) -> str: """生成用户所有 Session 的 key""" return f"user_sessions:{user_id}" def save_session( self, session_id: str, user_id: str, refresh_token: str, device_info: str, ip_address: str, expires_in_seconds: int = 30 * 24 * 60 * 60, # 30 天 ) -> bool: """ 保存 Session Args: session_id: Session ID user_id: 用户 ID refresh_token: 刷新令牌 device_info: 设备信息 ip_address: IP 地址 expires_in_seconds: 过期时间(秒) Returns: 是否保存成功 """ try: now = datetime.utcnow() expires_at = now + timedelta(seconds=expires_in_seconds) session_data = { "session_id": session_id, "user_id": user_id, "device_info": device_info, "ip_address": ip_address, "created_at": now.isoformat(), "last_active_at": now.isoformat(), "expires_at": expires_at.isoformat(), } # 保存 Session 数据 session_key = self._session_key(session_id) self.redis.setex( session_key, expires_in_seconds, json.dumps(session_data) ) # 保存 refresh_token 映射 refresh_token_key = self._refresh_token_key(session_id) self.redis.setex( refresh_token_key, expires_in_seconds, refresh_token ) # 添加到用户的 Session 集合 user_sessions_key = self._user_sessions_key(user_id) self.redis.sadd(user_sessions_key, session_id) self.redis.expire(user_sessions_key, expires_in_seconds) return True except Exception as e: print(f"Failed to save session: {e}") return False def get_session(self, session_id: str) -> Optional[dict]: """ 获取 Session Args: session_id: Session ID Returns: Session 数据,如果不存在返回 None """ try: session_key = self._session_key(session_id) data = self.redis.get(session_key) if data: return json.loads(data) return None except Exception as e: print(f"Failed to get session: {e}") return None def get_refresh_token(self, session_id: str) -> Optional[str]: """ 获取 refresh_token Args: session_id: Session ID Returns: refresh_token,如果不存在返回 None """ try: refresh_token_key = self._refresh_token_key(session_id) return self.redis.get(refresh_token_key) except Exception as e: print(f"Failed to get refresh_token: {e}") return None def update_last_active(self, session_id: str) -> bool: """ 更新 Session 最后活跃时间 Args: session_id: Session ID Returns: 是否更新成功 """ try: session = self.get_session(session_id) if not session: return False session["last_active_at"] = datetime.utcnow().isoformat() session_key = self._session_key(session_id) ttl = self.redis.ttl(session_key) if ttl > 0: self.redis.setex( session_key, ttl, json.dumps(session) ) return True return False except Exception as e: print(f"Failed to update last active: {e}") return False def delete_session(self, session_id: str) -> bool: """ 删除 Session(登出) Args: session_id: Session ID Returns: 是否删除成功 """ try: session = self.get_session(session_id) if not session: return False user_id = session["user_id"] # 删除 Session 数据 session_key = self._session_key(session_id) self.redis.delete(session_key) # 删除 refresh_token refresh_token_key = self._refresh_token_key(session_id) self.redis.delete(refresh_token_key) # 从用户 Session 集合中移除 user_sessions_key = self._user_sessions_key(user_id) self.redis.srem(user_sessions_key, session_id) return True except Exception as e: print(f"Failed to delete session: {e}") return False def get_user_sessions(self, user_id: str) -> list[dict]: """ 获取用户的所有活跃 Session Args: user_id: 用户 ID Returns: Session 列表 """ try: user_sessions_key = self._user_sessions_key(user_id) session_ids = self.redis.smembers(user_sessions_key) sessions = [] for session_id in session_ids: session = self.get_session(session_id) if session: sessions.append(session) return sessions except Exception as e: print(f"Failed to get user sessions: {e}") return [] def delete_all_user_sessions(self, user_id: str) -> int: """ 删除用户的所有 Session(强制登出所有设备) Args: user_id: 用户 ID Returns: 删除的 Session 数量 """ try: sessions = self.get_user_sessions(user_id) count = 0 for session in sessions: if self.delete_session(session["session_id"]): count += 1 # 清空用户 Session 集合 user_sessions_key = self._user_sessions_key(user_id) self.redis.delete(user_sessions_key) return count except Exception as e: print(f"Failed to delete all user sessions: {e}") return 0 def session_exists(self, session_id: str) -> bool: """ 检查 Session 是否存在 Args: session_id: Session ID Returns: 是否存在 """ try: session_key = self._session_key(session_id) return self.redis.exists(session_key) > 0 except Exception: return False # 全局实例(生产环境应该从配置读取) session_store = SessionStore()