292 lines
8.0 KiB
Python
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
|