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:
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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()
|
||||
Binary file not shown.
@@ -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"
|
||||
Reference in New Issue
Block a user