Files
xiaoxia-saas/packages/application/auth/login_use_case.py
T
灵应 e25fd86171
Deploy / Staging E2E Tests (push) Has been skipped
Deploy / Build Production Runtime Images (push) Has been skipped
Deploy / Deploy Production (push) Has been skipped
Deploy / Production Browser E2E (push) Has been skipped
Deploy / Deploy Staging (push) Failing after 137h58m6s
CI/CD Pipeline / Frontend Lint (push) Failing after 137h58m12s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 137h58m12s
style: 后端代码isort import排序
2026-07-03 18:56:16 +08:00

311 lines
10 KiB
Python
Executable File

"""
用户登录 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.application.auth.jwt_service import jwt_service
from packages.application.auth.password_hasher import 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()
# jwt_service proxy already imported at module level (line 15)
self.jwt_service = jwt_service
self._jwt_secret_key = jwt_service.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
使用 Redis 中的反向索引 (refresh_token -> session_id) 快速查找 session。
反向索引在 save_session 时创建,确保了 O(1) 的查找复杂度。
Args:
refresh_token: 刷新令牌
Returns:
Session 数据字典,包含 session_id, user_id 等信息;如果不存在返回 None
"""
return self.session_store.get_session_by_refresh_token(refresh_token)
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)}"