diff --git a/packages/domain/auth/__init__.py b/packages/domain/auth/__init__.py index ca97bd436..358139daf 100644 --- a/packages/domain/auth/__init__.py +++ b/packages/domain/auth/__init__.py @@ -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", ] diff --git a/packages/domain/auth/session_store.py b/packages/domain/auth/session_store.py new file mode 100644 index 000000000..0f2bc2b53 --- /dev/null +++ b/packages/domain/auth/session_store.py @@ -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() diff --git a/requirements.txt b/requirements.txt index e749dd81d..a2706139f 100644 Binary files a/requirements.txt and b/requirements.txt differ diff --git a/tests/unit/test_session_store.py b/tests/unit/test_session_store.py new file mode 100644 index 000000000..05f971dd3 --- /dev/null +++ b/tests/unit/test_session_store.py @@ -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"