feat(auth): add Redis session store for refresh tokens

- Implement SessionStore class with Redis backend
- Support save/get/delete session and refresh_token
- Support user multi-device sessions management
- Add last_active tracking and force logout all devices
- Add 10 comprehensive unit tests with Mock Redis (all passed)
- Install redis dependency

Phase 4 Task 3/68 completed
This commit is contained in:
Xiaoxia AI
2026-06-17 01:15:08 +08:00
parent 1d8d7bddb9
commit 86ea3c7eef
4 changed files with 587 additions and 0 deletions
+4
View File
@@ -6,6 +6,7 @@ from packages.domain.auth.password_hasher import (
password_hasher,
password_validator,
)
from packages.domain.auth.session_store import SessionStore, RedisConfig, session_store
__all__ = [
"JWTService",
@@ -16,4 +17,7 @@ __all__ = [
"PasswordValidator",
"password_hasher",
"password_validator",
"SessionStore",
"RedisConfig",
"session_store",
]
+294
View File
@@ -0,0 +1,294 @@
"""
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()
BIN
View File
Binary file not shown.
+289
View File
@@ -0,0 +1,289 @@
"""
Redis Session 存储测试
"""
import pytest
import json
from datetime import datetime
from unittest.mock import Mock, MagicMock
from packages.domain.auth.session_store import SessionStore
class TestSessionStore:
"""Session 存储测试"""
@pytest.fixture
def mock_redis(self):
"""创建 Mock Redis 客户端"""
redis_mock = Mock()
redis_mock.data = {} # 模拟内存存储
redis_mock.expires = {} # 模拟过期时间
redis_mock.sets = {} # 模拟集合
def setex(key, seconds, value):
redis_mock.data[key] = value
redis_mock.expires[key] = seconds
return True
def get(key):
return redis_mock.data.get(key)
def delete(key):
if key in redis_mock.data:
del redis_mock.data[key]
return 1
return 0
def exists(key):
return 1 if key in redis_mock.data else 0
def ttl(key):
return redis_mock.expires.get(key, -1)
def sadd(key, *values):
if key not in redis_mock.sets:
redis_mock.sets[key] = set()
redis_mock.sets[key].update(values)
return len(values)
def smembers(key):
return redis_mock.sets.get(key, set())
def srem(key, *values):
if key in redis_mock.sets:
redis_mock.sets[key].discard(*values)
return len(values)
return 0
def expire(key, seconds):
redis_mock.expires[key] = seconds
return True
redis_mock.setex = setex
redis_mock.get = get
redis_mock.delete = delete
redis_mock.exists = exists
redis_mock.ttl = ttl
redis_mock.sadd = sadd
redis_mock.smembers = smembers
redis_mock.srem = srem
redis_mock.expire = expire
return redis_mock
@pytest.fixture
def session_store(self, mock_redis):
"""创建 Session 存储实例"""
return SessionStore(redis_client=mock_redis)
def test_save_session(self, session_store, mock_redis):
"""测试保存 Session"""
result = session_store.save_session(
session_id="session-123",
user_id="user-456",
refresh_token="refresh-token-abc",
device_info="Chrome/Windows",
ip_address="192.168.1.1",
expires_in_seconds=3600,
)
assert result is True
# 验证数据已保存
session_key = "session:session-123"
assert session_key in mock_redis.data
session_data = json.loads(mock_redis.data[session_key])
assert session_data["session_id"] == "session-123"
assert session_data["user_id"] == "user-456"
assert session_data["device_info"] == "Chrome/Windows"
assert session_data["ip_address"] == "192.168.1.1"
# 验证 refresh_token 已保存
refresh_token_key = "refresh_token:session-123"
assert mock_redis.data[refresh_token_key] == "refresh-token-abc"
# 验证用户 Session 集合已更新
user_sessions_key = "user_sessions:user-456"
assert "session-123" in mock_redis.sets[user_sessions_key]
def test_get_session(self, session_store, mock_redis):
"""测试获取 Session"""
# 先保存
session_store.save_session(
session_id="session-123",
user_id="user-456",
refresh_token="token",
device_info="Chrome",
ip_address="127.0.0.1",
)
# 获取
session = session_store.get_session("session-123")
assert session is not None
assert session["session_id"] == "session-123"
assert session["user_id"] == "user-456"
assert session["device_info"] == "Chrome"
def test_get_nonexistent_session(self, session_store):
"""测试获取不存在的 Session"""
session = session_store.get_session("nonexistent")
assert session is None
def test_get_refresh_token(self, session_store):
"""测试获取 refresh_token"""
session_store.save_session(
session_id="session-123",
user_id="user-456",
refresh_token="my-refresh-token",
device_info="Chrome",
ip_address="127.0.0.1",
)
token = session_store.get_refresh_token("session-123")
assert token == "my-refresh-token"
def test_update_last_active(self, session_store, mock_redis):
"""测试更新最后活跃时间"""
session_store.save_session(
session_id="session-123",
user_id="user-456",
refresh_token="token",
device_info="Chrome",
ip_address="127.0.0.1",
)
# 获取原始时间
session1 = session_store.get_session("session-123")
original_time = session1["last_active_at"]
# 更新
result = session_store.update_last_active("session-123")
assert result is True
# 验证时间已更新
session2 = session_store.get_session("session-123")
assert session2["last_active_at"] >= original_time
def test_delete_session(self, session_store, mock_redis):
"""测试删除 Session"""
session_store.save_session(
session_id="session-123",
user_id="user-456",
refresh_token="token",
device_info="Chrome",
ip_address="127.0.0.1",
)
# 删除
result = session_store.delete_session("session-123")
assert result is True
# 验证已删除
session = session_store.get_session("session-123")
assert session is None
token = session_store.get_refresh_token("session-123")
assert token is None
# 验证从用户集合中移除
user_sessions_key = "user_sessions:user-456"
assert "session-123" not in mock_redis.sets.get(user_sessions_key, set())
def test_get_user_sessions(self, session_store):
"""测试获取用户的所有 Session"""
# 创建多个 Session
session_store.save_session(
session_id="session-1",
user_id="user-456",
refresh_token="token1",
device_info="Chrome",
ip_address="192.168.1.1",
)
session_store.save_session(
session_id="session-2",
user_id="user-456",
refresh_token="token2",
device_info="Firefox",
ip_address="192.168.1.2",
)
# 获取
sessions = session_store.get_user_sessions("user-456")
assert len(sessions) == 2
session_ids = [s["session_id"] for s in sessions]
assert "session-1" in session_ids
assert "session-2" in session_ids
def test_delete_all_user_sessions(self, session_store):
"""测试删除用户的所有 Session"""
# 创建多个 Session
session_store.save_session(
session_id="session-1",
user_id="user-456",
refresh_token="token1",
device_info="Chrome",
ip_address="127.0.0.1",
)
session_store.save_session(
session_id="session-2",
user_id="user-456",
refresh_token="token2",
device_info="Firefox",
ip_address="127.0.0.1",
)
# 删除所有
count = session_store.delete_all_user_sessions("user-456")
assert count == 2
# 验证已删除
sessions = session_store.get_user_sessions("user-456")
assert len(sessions) == 0
def test_session_exists(self, session_store):
"""测试检查 Session 是否存在"""
assert session_store.session_exists("nonexistent") is False
session_store.save_session(
session_id="session-123",
user_id="user-456",
refresh_token="token",
device_info="Chrome",
ip_address="127.0.0.1",
)
assert session_store.session_exists("session-123") is True
def test_multiple_users(self, session_store):
"""测试多用户隔离"""
# 用户 1 的 Session
session_store.save_session(
session_id="session-user1",
user_id="user-1",
refresh_token="token1",
device_info="Chrome",
ip_address="127.0.0.1",
)
# 用户 2 的 Session
session_store.save_session(
session_id="session-user2",
user_id="user-2",
refresh_token="token2",
device_info="Firefox",
ip_address="127.0.0.1",
)
# 验证隔离
user1_sessions = session_store.get_user_sessions("user-1")
assert len(user1_sessions) == 1
assert user1_sessions[0]["session_id"] == "session-user1"
user2_sessions = session_store.get_user_sessions("user-2")
assert len(user2_sessions) == 1
assert user2_sessions[0]["session_id"] == "session-user2"