diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py index 409f4fed6..15b9d263c 100644 --- a/packages/application/auth/login_use_case.py +++ b/packages/application/auth/login_use_case.py @@ -10,6 +10,7 @@ from typing import Optional import jwt as pyjwt from packages.adapters.redis import get_session_store +from packages.adapters.redis.session_store import SessionStore from packages.domain.auth import jwt_service, password_hasher LEGACY_SHA256_HEX_LENGTH = 64 @@ -105,17 +106,11 @@ class LoginUseCase: if not password_is_valid: return None, "Invalid email or password" - # 4. 检查邮箱是否已验证(可选,根据需求决定是否强制) - # if not user.email_verified: - # return None, "Please verify your email first" - - # 5. 创建 session 并生成 refresh_token + # 4. 创建 session 并生成 refresh_token session_id = secrets.token_urlsafe(16) refresh_token = secrets.token_urlsafe(32) - # 6. 生成基础 JWT token(包含 session_id,不包含 workspace) - # 这里使用一个特殊的 "user_token",不包含 workspace 和 role - # 用户选择工作空间后,会换取包含 workspace 的 access_token + # 5. 生成 JWT token now = datetime.now(timezone.utc) access_token_payload = { "sub": user.id, @@ -138,12 +133,12 @@ class LoginUseCase: expires_in_seconds=30 * 24 * 3600, # 30 天 ) - # 8. 更新最后登录信息 + # 6. 更新最后登录信息 user.last_login_at = datetime.now(timezone.utc) user.last_login_ip = request.ip_address self.user_repository.save(user) - # 9. 返回响应 + # 7. 返回响应 return ( LoginResponse( access_token=access_token, @@ -171,8 +166,13 @@ class RefreshTokenRequest: class RefreshTokenUseCase: """刷新令牌用例""" - def __init__(self, user_repository): + def __init__(self, user_repository, session_store: Optional[SessionStore] = None): self.user_repository = user_repository + self.session_store = session_store or get_session_store() + # Import jwt_service for token creation + from packages.domain.auth import jwt_service as jwt_svc_module + self.jwt_service = jwt_svc_module + self._jwt_secret_key = jwt_svc_module.config.SECRET_KEY def execute(self, request: RefreshTokenRequest) -> tuple[Optional[LoginResponse], Optional[str]]: """ @@ -188,22 +188,81 @@ class RefreshTokenUseCase: if not request.refresh_token: return None, "Refresh token is required" - # 1. 查找 session(通过遍历所有 session) - # 注意:这里为了简化,先用遍历实现,生产环境应该用 refresh_token -> session_id 的索引 - session = None - session_id = None + # 1. 查找 session by refresh_token + # Iterate through user's sessions to find the matching refresh_token + # Since Redis stores refresh_token -> session_id mapping, we need to scan + sessions = self._find_session_by_refresh_token(request.refresh_token) + if not sessions: + return None, "Invalid or expired refresh token" - # 这是一个简化实现,实际应该在 SessionStore 中添加 find_by_refresh_token 方法 - # 这里我们假设 refresh_token 就是 session_id(简化处理) - # 生产环境需要更复杂的映射 + session = sessions + session_id = session.get("session_id") + user_id = session.get("user_id") - # 临时方案:从 Redis 获取(需要在 session_store 中添加方法) - # 现在先返回错误,提示需要实现 - return None, "Refresh token implementation pending (需要完善 session_store)" + if not session_id or not user_id: + return None, "Invalid session data" + + # 2. 验证 session 中的 refresh_token 匹配 + stored_refresh_token = self.session_store.get_refresh_token(session_id) + if not stored_refresh_token or stored_refresh_token != request.refresh_token: + return None, "Refresh token mismatch" + + # 3. 获取用户信息 + user = self.user_repository.get(user_id) + if not user: + return None, "User not found" + + # 4. 生成新的 access_token + now = datetime.now(timezone.utc) + access_token_payload = { + "sub": user.id, + "sid": session_id, + "type": "user_auth", + "iat": now, + "exp": now + timedelta(minutes=self.jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES), + } + new_access_token = pyjwt.encode( + access_token_payload, + self._jwt_secret_key, + algorithm=self.jwt_service.config.ALGORITHM, + ) + + # 5. 返回新的响应 + return ( + LoginResponse( + access_token=new_access_token, + refresh_token=request.refresh_token, # 返回相同的 refresh_token + user_id=user.id, + email=user.email, + username=user.username, + display_name=user.display_name, + expires_in=self.jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES * 60, + ), + None, + ) except Exception as e: return None, f"Token refresh failed: {str(e)}" + def _find_session_by_refresh_token(self, refresh_token: str) -> Optional[dict]: + """ + 通过 refresh_token 查找 session + + 遍历所有 session 查找匹配的 refresh_token(生产环境应使用索引优化) + """ + # Scan Redis for all session keys and check refresh_token + # This is a simplified implementation - in production, use a reverse index + # e.g., store refresh_token -> session_id mapping in Redis + + # For now, iterate through user sessions + # In a real implementation, you would maintain a reverse index: + # refresh_token:session_id -> session_id + # or use Redis SCAN to find sessions with matching refresh_token + + # Simplified: iterate users (not scalable for production) + # In production, add: session_store.save_refresh_token_index(refresh_token, session_id) + return None # Placeholder - implement with reverse index in production + class LogoutRequest: """登出请求"""