fix(P1-2): implement token refresh use case with session verification
This commit is contained in:
@@ -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:
|
||||
"""登出请求"""
|
||||
|
||||
Reference in New Issue
Block a user