fix(P1-2): implement token refresh use case with session verification

This commit is contained in:
2026-06-26 17:52:34 +08:00
parent 73043ed72b
commit b326e4a249
+80 -21
View File
@@ -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:
"""登出请求"""