Files
xiaoxia-saas/packages/adapters/redis/session_store.py
T
2026-06-21 06:52:19 +08:00

292 lines
8.0 KiB
Python

"""
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 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.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)
# 添加到用户的 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.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)
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 = None
def get_session_store() -> SessionStore:
global _session_store
if _session_store is None:
_session_store = SessionStore()
return _session_store