Files
xiaoxia-saas/packages/application/auth/login_use_case.py
T
xiaoxia 5f3688e002
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
Tests / lint (pull_request) Failing after 1965h33m12s
Tests / test (pull_request) Failing after 1965h34m13s
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 1965h35m13s
Deploy / Production Browser E2E (push) Failing after 1965h38m7s
Deploy / Deploy Production (push) Failing after 1965h38m9s
Deploy / Deploy Staging (push) Failing after 1965h38m13s
Deploy / Build Production Runtime Images (push) Failing after 1966h9m29s
fix(P1-2): implement token refresh use case with session verification
2026-06-26 17:52:34 +08:00

315 lines
10 KiB
Python

"""
用户登录 Use Case
"""
import hashlib
import secrets
from datetime import datetime, timedelta, timezone
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
def _is_legacy_sha256_hash(password_hash: str) -> bool:
return len(password_hash) == LEGACY_SHA256_HEX_LENGTH and all(
char in "0123456789abcdef" for char in password_hash.lower()
)
def _legacy_sha256(password: str) -> str:
return hashlib.sha256(password.encode()).hexdigest()
class LoginRequest:
"""登录请求"""
def __init__(
self,
email: str,
password: str,
device_info: Optional[str] = None,
ip_address: Optional[str] = None,
):
self.email = email.strip().lower()
self.password = password
self.device_info = device_info or "Unknown"
self.ip_address = ip_address or "unknown"
class LoginResponse:
"""登录响应"""
def __init__(
self,
access_token: str,
refresh_token: str,
user_id: str,
email: str,
username: str,
display_name: str,
expires_in: int,
):
self.access_token = access_token
self.refresh_token = refresh_token
self.user_id = user_id
self.email = email
self.username = username
self.display_name = display_name
self.expires_in = expires_in
class LoginUseCase:
"""用户登录用例"""
def __init__(self, user_repository, session_store=None, jwt_secret_key: str | None = None):
self.user_repository = user_repository
self.session_store = session_store or get_session_store()
self.jwt_secret_key = jwt_secret_key or jwt_service.config.SECRET_KEY
def execute(self, request: LoginRequest) -> tuple[Optional[LoginResponse], Optional[str]]:
"""
执行登录
Args:
request: 登录请求
Returns:
(登录响应, 错误信息)
"""
try:
# 1. 验证输入
if not request.email:
return None, "Email is required"
if not request.password:
return None, "Password is required"
# 2. 查找用户
user = self.user_repository.find_by_email(request.email)
if not user:
return None, "Invalid email or password"
# 3. 验证密码
password_is_valid = password_hasher.verify_password(request.password, user.password_hash)
if not password_is_valid and _is_legacy_sha256_hash(user.password_hash):
password_is_valid = secrets.compare_digest(_legacy_sha256(request.password), user.password_hash)
if password_is_valid:
user.password_hash = password_hasher.hash_password(request.password)
self.user_repository.save(user)
if not password_is_valid:
return None, "Invalid email or password"
# 4. 创建 session 并生成 refresh_token
session_id = secrets.token_urlsafe(16)
refresh_token = secrets.token_urlsafe(32)
# 5. 生成 JWT token
now = datetime.now(timezone.utc)
access_token_payload = {
"sub": user.id,
"sid": session_id,
"type": "user_auth",
"iat": now,
"exp": now + timedelta(minutes=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES),
}
access_token = pyjwt.encode(
access_token_payload,
self.jwt_secret_key,
algorithm=jwt_service.config.ALGORITHM,
)
self.session_store.save_session(
session_id=session_id,
user_id=user.id,
refresh_token=refresh_token,
device_info=request.device_info,
ip_address=request.ip_address,
expires_in_seconds=30 * 24 * 3600, # 30 天
)
# 6. 更新最后登录信息
user.last_login_at = datetime.now(timezone.utc)
user.last_login_ip = request.ip_address
self.user_repository.save(user)
# 7. 返回响应
return (
LoginResponse(
access_token=access_token,
refresh_token=refresh_token,
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
expires_in=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
),
None,
)
except Exception as e:
return None, f"Login failed: {str(e)}"
class RefreshTokenRequest:
"""刷新令牌请求"""
def __init__(self, refresh_token: str):
self.refresh_token = refresh_token
class RefreshTokenUseCase:
"""刷新令牌用例"""
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]]:
"""
执行令牌刷新
Args:
request: 刷新请求
Returns:
(登录响应, 错误信息)
"""
try:
if not request.refresh_token:
return None, "Refresh token is required"
# 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"
session = sessions
session_id = session.get("session_id")
user_id = session.get("user_id")
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:
"""登出请求"""
def __init__(
self,
user_id: str,
session_id: Optional[str] = None,
logout_all_devices: bool = False,
):
self.user_id = user_id
self.session_id = session_id
self.logout_all_devices = logout_all_devices
class LogoutUseCase:
"""用户登出用例"""
def __init__(self, session_store=None):
self.session_store = session_store or get_session_store()
def execute(self, request: LogoutRequest) -> tuple[bool, Optional[str]]:
"""
执行登出
Args:
request: 登出请求
Returns:
(是否成功, 错误信息)
"""
try:
if request.logout_all_devices:
# 删除所有设备的 session
count = self.session_store.delete_all_user_sessions(request.user_id)
return True, None
else:
# 删除当前 session
if not request.session_id:
return False, "Session ID is required"
success = self.session_store.delete_session(request.session_id)
if success:
return True, None
else:
return False, "Session not found"
except Exception as e:
return False, f"Logout failed: {str(e)}"