""" Redis Session 存储 用于存储 refresh_token 和 Session 信息 """ import json from datetime import datetime, timedelta, timezone from typing import Optional 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 NoopSessionStore: def save_session(self, **kwargs) -> bool: return False def get_session(self, session_id: str) -> Optional[dict]: return None def get_session_by_refresh_token(self, refresh_token: str) -> Optional[dict]: return None def get_refresh_token(self, session_id: str) -> Optional[str]: return None def update_last_active(self, session_id: str) -> bool: return False def delete_session(self, session_id: str) -> bool: return False def get_user_sessions(self, user_id: str) -> list[dict]: return [] def delete_all_user_sessions(self, user_id: str) -> int: return 0 def session_exists(self, session_id: str) -> bool: return False 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 _refresh_token_to_session_key(self, refresh_token: str) -> str: """生成 refresh_token -> session_id 的反向映射 key""" return f"refresh_token_map:{refresh_token}" 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.now(timezone.utc) 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) # 保存 refresh_token -> session_id 的反向映射 refresh_token_map_key = self._refresh_token_to_session_key(refresh_token) self.redis.setex(refresh_token_map_key, expires_in_seconds, session_id) # 添加到用户的 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_session_by_refresh_token(self, refresh_token: str) -> Optional[dict]: """ 通过 refresh_token 获取 Session Args: refresh_token: 刷新令牌 Returns: Session 数据,如果不存在返回 None """ try: # 先通过反向映射找到 session_id refresh_token_map_key = self._refresh_token_to_session_key(refresh_token) session_id = self.redis.get(refresh_token_map_key) if not session_id: return None # 再获取完整的 session 数据 return self.get_session(session_id) except Exception as e: print(f"Failed to get session by refresh_token: {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.now(timezone.utc).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) refresh_token = self.redis.get(refresh_token_key) self.redis.delete(refresh_token_key) # 删除反向映射 if refresh_token: refresh_token_map_key = self._refresh_token_to_session_key(refresh_token) self.redis.delete(refresh_token_map_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 = None def get_session_store(enabled: bool = True): global _session_store if not enabled: return NoopSessionStore() if _session_store is None: _session_store = SessionStore() return _session_store