Files
xiaoxia-saas/packages/application/auth/login_use_case.py
T
CI Bot 9c6c477f55
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 2m22s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 2m24s
CI/CD Pipeline / Integration Tests (pull_request) Failing after 37s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1561h35m7s
CI/CD Pipeline / Build Production Runtime Images (pull_request) Failing after 1561h35m8s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1561h35m7s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Failing after 1561h35m9s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1561h35m7s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1562h6m43s
fix(backend): Phase 1 后端代码清理与修复
P0 关键修复:
- P0-1: 注册接口添加 RateLimitMiddleware 限流保护
- P0-3: /metrics 端点添加 JWT 认证(移除匿名访问)
- P0-4: 修复 Celery 任务名冲突(generation_task vs generate_video)
- P1-5: JWT logout token 黑名单机制

P1 修复:
- P1-1: forgot_password 硬编码 localhost → 使用 settings.APP_BASE_URL
- P1-2: generation.py 直接创建 DB 连接 → 使用依赖注入
- P1-6: Image.open() 未关闭 → 统一使用 with 语句
- P1-7: 订阅续费事务修复

P2 代码质量:
- P2-1: 修复 EditingMode 枚举重复定义 → 统一引用 shared 包
- P2-2: 修复 SMTP_FRON_NAME → SMTP_FROM_NAME 拼写
- P2-3: UserModel subscription_quota 类型统一为 float
- P2-4: .env.production DATABASE_MAX_OVERFLOW 30 → 10
- 清理 15 处 except:pass(保留 2 处有注释说明的)
- 禁用 SVG 上传(XSS 风险)
- 删除 decode_token_unsafe() 不安全函数
- 简化 /ready 端点
- 删除 8 处死代码、10 个空文件/模块
- 合并 3 对 100% 重复函数
- 对齐 6 个废弃环境变量

v2 修复(代码审查后):
- 修复密码重置路由路径: /password/forgot → /forgot-password,
  /password/reset → /reset-password(与前端 API 对齐)
- 合并 _check_project_access: asset_libraries.py 和 edit_plans.py
  中的重复函数统一到 _helpers.py(含空字符串守卫 + 中文错误信息)
- 顺手修复: HTTPException 统一从 fastapi 导入(替换 starlette 导入)
- OSS_ENDPOINT 拼写修复拆分为单独 PR,本 PR 不包含
2026-07-13 13:50:52 +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
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)}"