chore: 修复 P2/P3 技术债务 #15

Merged
xiaoxia merged 12 commits from fix/p2p3-tech-debt-v2 into main 2026-06-26 18:41:32 +08:00
12 changed files with 777 additions and 86 deletions
+37
View File
@@ -94,3 +94,40 @@ jobs:
echo "✅ Build completed successfully!"
echo "Branch: ${GITHUB_REF_NAME}"
echo "Commit: ${GITHUB_SHA}"
# P3-4 Fix: 添加前端 Lint 检查 job
frontend-lint:
name: Frontend Lint
runs-on: ubuntu-latest
container: node:20
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: apps/web/package-lock.json
- name: Install dependencies
working-directory: apps/web
run: npm ci
- name: Run ESLint
working-directory: apps/web
run: npx eslint src --ext .ts,.tsx --max-warnings 0
- name: Run TypeScript type check
working-directory: apps/web
run: npx tsc --noEmit
- name: Run Prettier check
working-directory: apps/web
run: npx prettier --check "src/**/*.{ts,tsx,css,md}"
- name: Run Vitest tests
working-directory: apps/web
run: npx vitest run --coverage
+4 -13
View File
@@ -7,6 +7,7 @@ from app.schemas.project import (
ListProjectsResponse,
ProjectResponse,
)
from app.api.routes.permissions import require_workspace_member
from fastapi import APIRouter, Depends, HTTPException, status
from packages.application import (
@@ -20,16 +21,6 @@ from packages.ports.workspace_member_repository import WorkspaceMemberRepository
router = APIRouter()
def _require_workspace_member(
workspace_id: str,
authenticated_user: AuthenticatedUser,
workspace_member_repository: WorkspaceMemberRepository,
) -> None:
member = workspace_member_repository.find_by_workspace_and_user(workspace_id, authenticated_user.user.id)
if member is None:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Workspace access denied")
def _to_project_response(item) -> ProjectResponse:
return ProjectResponse(
id=item.id,
@@ -50,7 +41,7 @@ def get_project(
project = use_case.execute(project_id)
if project is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
_require_workspace_member(project.workspace_id, authenticated_user, workspace_member_repository)
require_workspace_member(project.workspace_id, authenticated_user, workspace_member_repository)
return _to_project_response(project)
@@ -61,7 +52,7 @@ def list_projects(
project_repository: Any = Depends(get_project_repository),
workspace_member_repository: WorkspaceMemberRepository = Depends(get_workspace_member_repository),
) -> ListProjectsResponse:
_require_workspace_member(workspace_id, authenticated_user, workspace_member_repository)
require_workspace_member(workspace_id, authenticated_user, workspace_member_repository)
use_case = ListProjectsUseCase(project_repository)
projects = use_case.execute(workspace_id)
return ListProjectsResponse(items=[_to_project_response(item) for item in projects])
@@ -74,7 +65,7 @@ def create_project(
project_repository: Any = Depends(get_project_repository),
workspace_member_repository: WorkspaceMemberRepository = Depends(get_workspace_member_repository),
) -> ProjectResponse:
_require_workspace_member(request.workspace_id, authenticated_user, workspace_member_repository)
require_workspace_member(request.workspace_id, authenticated_user, workspace_member_repository)
use_case = CreateProjectUseCase(project_repository)
project = use_case.execute(
CreateProjectCommand(
+45 -14
View File
@@ -2,6 +2,7 @@ from typing import Any
from uuid import uuid4
from app.auth import AuthenticatedUser, get_current_user
from app.api.routes.permissions import require_workspace_member
from app.config import get_settings
from app.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service
@@ -25,15 +26,38 @@ from packages.ports.workspace_member_repository import WorkspaceMemberRepository
router = APIRouter()
# 允许上传的文件 MIME 类型
ALLOWED_MIME_TYPES = frozenset({
# 视频
"video/mp4", "video/mpeg", "video/quicktime", "video/x-msvideo",
"video/webm", "video/x-matroska", "video/3gpp",
# 音频
"audio/mpeg", "audio/wav", "audio/ogg", "audio/flac", "audio/aac",
"audio/mp3", "audio/x-m4a", "audio/webm",
# 图片
"image/jpeg", "image/png", "image/gif", "image/webp", "image/bmp",
"image/svg+xml", "image/tiff",
})
def _require_workspace_member(
workspace_id: str,
authenticated_user: AuthenticatedUser,
workspace_member_repository: WorkspaceMemberRepository,
) -> None:
member = workspace_member_repository.find_by_workspace_and_user(workspace_id, authenticated_user.user.id)
if member is None:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Workspace access denied")
def _validate_mime_type(content_type: str | None) -> str:
"""验证并返回标准化的 MIME 类型,如果无效则抛出异常。"""
if not content_type:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Content-Type header is required",
)
# 处理带参数的类型,如 "video/mp4; charset=utf-8"
base_type = content_type.split(";")[0].strip().lower()
if base_type not in ALLOWED_MIME_TYPES:
raise HTTPException(
status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
detail=f"File type '{base_type}' is not supported. Allowed types: video, audio, and image files.",
)
return base_type
def _require_project_and_library(
@@ -90,7 +114,10 @@ async def prepare_direct_upload(
detail=f"File exceeds upload limit ({settings.OSS_DIRECT_UPLOAD_MAX_MB}MB)",
)
_require_workspace_member(request.workspace_id, authenticated_user, workspace_member_repository)
# P2-5: 服务端验证 MIME 类型
validated_content_type = _validate_mime_type(request.content_type)
require_workspace_member(request.workspace_id, authenticated_user, workspace_member_repository)
_require_project_and_library(
request.workspace_id,
request.project_id,
@@ -105,7 +132,7 @@ async def prepare_direct_upload(
try:
payload = storage_service.create_direct_upload_post(
storage_key=storage_key,
content_type=request.content_type or "application/octet-stream",
content_type=validated_content_type,
max_size_bytes=max_size_bytes,
expires_seconds=settings.OSS_DIRECT_UPLOAD_EXPIRE_SECONDS,
)
@@ -133,7 +160,7 @@ async def complete_direct_upload(
storage_service: OSSStorageService = Depends(get_storage_service),
) -> DirectUploadCompleteResponse:
"""确认浏览器直传完成并创建导入任务。"""
_require_workspace_member(request.workspace_id, authenticated_user, workspace_member_repository)
require_workspace_member(request.workspace_id, authenticated_user, workspace_member_repository)
_require_project_and_library(
request.workspace_id,
request.project_id,
@@ -171,16 +198,20 @@ async def upload_asset(
storage_service: OSSStorageService = Depends(get_storage_service),
) -> UploadAssetResponse:
"""上传素材文件并触发导入流水线。"""
_require_workspace_member(workspace_id, authenticated_user, workspace_member_repository)
require_workspace_member(workspace_id, authenticated_user, workspace_member_repository)
_require_project_and_library(workspace_id, project_id, library_id, project_repository, asset_library_repository)
# P2-5: 服务端验证 MIME 类型
validated_content_type = _validate_mime_type(file.content_type)
file_id = uuid4().hex[:8]
storage_key = f"uploads/{file_id}/{file.filename}"
safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown"
storage_key = f"uploads/{file_id}/{safe_filename}"
file_url = storage_service.upload_file(
file.file,
storage_key,
content_type=file.content_type or "application/octet-stream",
content_type=validated_content_type,
)
job = _submit_ingest_job(
+9 -31
View File
@@ -1,34 +1,12 @@
import os
from typing import Optional
"""
数据库连接管理
from pydantic_settings import BaseSettings, SettingsConfigDict
P3-1 Fix: 统一使用 config.py 中的 Settings 消除配置重复定义
"""
from app.config import settings
from packages.adapters.sqlalchemy_impl.session import build_session_factory
class DatabaseSettings(BaseSettings):
database_url: str = "postgresql+psycopg://postgres:postgres@postgres:5432/xiaoxia_saas"
pool_size: int = 20
max_overflow: int = 40
pool_timeout: int = 30
pool_recycle: int = 3600
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
case_sensitive=False,
extra="ignore",
)
_settings: Optional[DatabaseSettings] = None
def get_database_settings() -> DatabaseSettings:
global _settings
if _settings is None:
env = os.getenv("APP_ENV", "development")
env_file = f".env.{env}" if env != "development" else ".env"
if os.path.exists(env_file):
_settings = DatabaseSettings(_env_file=env_file)
else:
_settings = DatabaseSettings()
return _settings
# P3-1 Fix: 直接使用 settings 中的配置,而不是重复定义
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
+7 -4
View File
@@ -1,3 +1,5 @@
from typing import Generator
import redis
from app.config import settings
from fastapi import Depends
@@ -40,7 +42,8 @@ from packages.ports.workspace_repository import WorkspaceRepository
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
def get_db_session():
def get_db_session() -> Generator[Session, None, None]:
"""获取数据库会话"""
session: Session = _SessionLocal()
try:
yield session
@@ -120,13 +123,13 @@ def get_workspace_invitation_repository(
return SQLAlchemyWorkspaceInvitationRepository(session)
def get_auth_session_store():
if not settings.ENABLE_REDIS_SESSIONS:
def get_auth_session_store() -> NoopSessionStore | SessionStore:
if not settings.ENABLE_REDIS_SESSION:
return NoopSessionStore()
return SessionStore(redis_client=redis.from_url(settings.REDIS_URL, decode_responses=True))
def get_auth_email_service():
def get_auth_email_service() -> NoopEmailService:
if not settings.ENABLE_EMAIL_DELIVERY:
return NoopEmailService()
return get_email_service(
+23 -10
View File
@@ -61,18 +61,27 @@ class VideoProcessor:
# 确保输出目录存在
os.makedirs(os.path.dirname(output_path), exist_ok=True)
# P2-6 Fix: 使用 try-finally 确保临时文件清理
concat_file = None
try:
# 创建临时文件列表
concat_file = os.path.join(self.temp_dir, f"concat_{os.getpid()}.txt")
with open(concat_file, "w") as f:
for path in input_paths:
# FFmpeg concat demuxer 格式
f.write(f"file '{os.path.abspath(path)}'\n")
# 使用 NamedTemporaryFile 确保临时文件正确清理
concat_file = tempfile.NamedTemporaryFile(
mode="w",
suffix=".txt",
prefix="ffmpeg_concat_",
dir=self.temp_dir,
delete=True,
)
for path in input_paths:
# FFmpeg concat demuxer 格式
concat_file.write(f"file '{os.path.abspath(path)}'\n")
concat_file.flush()
concat_file_path = concat_file.name
# 使用 FFmpeg 拼接视频
width, height = resolution
(
ffmpeg.input(concat_file, format="concat", safe=0)
ffmpeg.input(concat_file_path, format="concat", safe=0)
.output(
output_path,
vcodec="libx264",
@@ -86,9 +95,6 @@ class VideoProcessor:
.run(capture_stdout=True, capture_stderr=True)
)
# 清理临时文件
os.remove(concat_file)
# 获取视频元数据
probe = ffmpeg.probe(output_path)
video_info = next(s for s in probe["streams"] if s["codec_type"] == "video")
@@ -120,6 +126,13 @@ class VideoProcessor:
except ffmpeg.Error as e:
stderr = e.stderr.decode() if e.stderr else ""
raise RuntimeError(f"FFmpeg error: {stderr}") from e
finally:
# P2-6 Fix: 确保临时文件在所有情况下都被清理
if concat_file is not None:
try:
concat_file.close()
except Exception:
pass # 忽略关闭时的错误
def generate_thumbnail(
self,
+60 -12
View File
@@ -1,16 +1,64 @@
FROM python:3.12-slim
# =============================================================================
# 小虾 SaaS API 多阶段构建
# =============================================================================
# Stage 1: Builder - 安装依赖
FROM python:3.12-slim AS builder
WORKDIR /app
ENV PYTHONPATH=/app
ENV PIP_INDEX_URL=https://mirrors.aliyun.com/pypi/simple/
ENV PIP_TRUSTED_HOST=mirrors.aliyun.com
COPY requirements.txt ./
RUN python -m pip install --upgrade pip setuptools wheel && \
pip install --no-cache-dir --default-timeout=120 --retries 10 -r requirements.txt
COPY apps/api /app/apps/api
COPY packages /app/packages
COPY scripts /app/scripts
COPY alembic /app/alembic
COPY alembic.ini /app/alembic.ini
# 安装构建依赖
RUN apt-get update && apt-get install -y --no-install-recommends \
gcc \
libpq-dev \
&& rm -rf /var/lib/apt/lists/*
# 复制依赖文件并安装
COPY requirements.txt .
RUN pip install --no-cache-dir --user --upgrade pip setuptools wheel && \
pip install --no-cache-dir --user -r requirements.txt
# =============================================================================
# Stage 2: Runtime - 运行镜像
FROM python:3.12-slim AS runtime
# 安全:使用非 root 用户
RUN groupadd --gid 1000 appgroup && \
useradd --uid 1000 --gid appgroup --shell /bin/bash --create-home appuser
WORKDIR /app
# 设置环境变量
ENV PYTHONPATH=/app \
PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_NO_CACHE_DIR=1 \
PIP_INDEX_URL=https://mirrors.aliyun.com/pypi/simple/ \
PIP_TRUSTED_HOST=mirrors.aliyun.com
# 安装运行时依赖(排除构建工具)
RUN apt-get update && apt-get install -y --no-install-recommends \
libpq5 \
&& rm -rf /var/lib/apt/lists/*
# 从 builder 阶段复制已安装的包
COPY --from=builder /root/.local /home/appuser/.local
COPY --from=builder /app /app
# 设置 PATH 包含用户本地 bin
ENV PATH=/home/appuser/.local/bin:$PATH
# 复制应用代码
COPY --chown=appuser:appgroup apps/api /app/apps/api
COPY --chown=appuser:appgroup packages /app/packages
COPY --chown=appuser:appgroup scripts /app/scripts
COPY --chown=appuser:appgroup alembic /app/alembic
COPY --chown=appuser:appgroup alembic.ini /app/alembic.ini
# 切换到非 root 用户
USER appuser
WORKDIR /app/apps/api
EXPOSE 8000
CMD ["uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8000"]
@@ -1,11 +1,16 @@
from __future__ import annotations
from sqlalchemy.orm import Session
from typing import List, Optional, TYPE_CHECKING
from sqlalchemy.orm import Session, joinedload
from packages.adapters.sqlalchemy_impl.models import WorkspaceMemberModel
from packages.domain.entities import WorkspaceMember
from packages.ports.workspace_member_repository import WorkspaceMemberRepository
if TYPE_CHECKING:
from packages.adapters.sqlalchemy_impl.models import UserModel
class SQLAlchemyWorkspaceMemberRepository(WorkspaceMemberRepository):
def __init__(self, session: Session):
@@ -56,9 +61,15 @@ class SQLAlchemyWorkspaceMemberRepository(WorkspaceMemberRepository):
return [self._to_entity(model) for model in models if model is not None]
def find_by_workspace(self, workspace_id: str) -> list[WorkspaceMember]:
"""查找 workspace 的所有成员,使用 JOIN 预加载用户数据避免 N+1 查询问题。
P2-3 Fix: 使用 joinedload 预加载关联的用户数据
这样在访问 member.user 时不会触发额外的数据库查询
"""
models = (
self.session.query(WorkspaceMemberModel)
.filter(WorkspaceMemberModel.workspace_id == workspace_id)
.options(joinedload(WorkspaceMemberModel.user)) # P2-3: JOIN 预加载用户数据
.order_by(WorkspaceMemberModel.joined_at.asc())
.all()
)
+19 -1
View File
@@ -1,4 +1,4 @@
"""认证相关 Use Cases"""
"""认证相关 Use Cases 和委托处理器"""
from packages.application.auth.login_use_case import (
LoginRequest,
@@ -22,8 +22,19 @@ from packages.application.auth.register_user_use_case import (
VerifyEmailRequest,
VerifyEmailUseCase,
)
from packages.application.auth.jwt_handler import (
JWTHandler,
configure_jwt_handler,
get_jwt_handler,
)
from packages.application.auth.password_handler import (
PasswordHandler,
configure_password_handler,
get_password_handler,
)
__all__ = [
# Use Cases
"RegisterUserUseCase",
"RegisterUserRequest",
"RegisterUserResponse",
@@ -40,4 +51,11 @@ __all__ = [
"RequestPasswordResetRequest",
"ResetPasswordUseCase",
"ResetPasswordRequest",
# Handlers (委托层,隔离 Domain 层对 jwt/bcrypt 的依赖)
"JWTHandler",
"configure_jwt_handler",
"get_jwt_handler",
"PasswordHandler",
"configure_password_handler",
"get_password_handler",
]
+143
View File
@@ -0,0 +1,143 @@
"""
JWT 处理器委托层
此模块作为 Domain (packages.domain.auth.jwt_service) Application 层之间的委托层
隔离 Domain 层对 jwt 库的直接依赖
使用方式
from packages.application.auth.jwt_handler import JWTHandler, get_jwt_handler
jwt_handler = JWTHandler(secret_key="your-secret-key")
token = jwt_handler.create_access_token(user_id="user123", workspace_id="ws456", role="admin")
payload = jwt_handler.verify_access_token(token)
"""
from datetime import datetime, timedelta
from typing import Any, Dict, Optional
from packages.domain.auth.jwt_service import JWTConfig, JWTService, TokenType
class JWTHandler:
"""
JWT 处理器委托类
委托给 packages.domain.auth.jwt_service.JWTService 进行实际的 JWT 操作
此层仅负责配置和封装不直接依赖 jwt
"""
def __init__(self, secret_key: str, algorithm: str = "HS256", access_token_expire_minutes: int = 30):
"""
初始化 JWT 处理器
Args:
secret_key: JWT 签名密钥必须从环境变量或配置注入
algorithm: 加密算法默认 HS256
access_token_expire_minutes: Access Token 过期时间分钟
"""
config = JWTConfig(
secret_key=secret_key,
algorithm=algorithm,
access_token_expire_minutes=access_token_expire_minutes,
)
self._service = JWTService(config)
def create_access_token(
self,
user_id: str,
workspace_id: str,
role: str,
additional_claims: Optional[Dict[str, Any]] = None,
) -> str:
"""
创建 access_token
Args:
user_id: 用户 ID
workspace_id: 工作空间 ID
role: 用户角色
additional_claims: 额外的声明信息
Returns:
JWT Token 字符串
"""
return self._service.create_access_token(
user_id=user_id,
workspace_id=workspace_id,
role=role,
additional_claims=additional_claims,
)
def verify_access_token(self, token: str) -> Dict[str, Any]:
"""
验证 access_token
Args:
token: JWT Token 字符串
Returns:
Token payload
Raises:
ExpiredSignatureError: Token 已过期
ValueError: Token 类型不是 access
"""
return self._service.verify_access_token(token)
def verify_token(self, token: str) -> Dict[str, Any]:
"""
验证任意 Token
Args:
token: JWT Token 字符串
Returns:
Token payload
"""
return self._service.verify_token(token)
# 默认处理器实例(需要通过 configure_jwt_handler 配置)
_default_handler: Optional[JWTHandler] = None
def configure_jwt_handler(
secret_key: str,
algorithm: str = "HS256",
access_token_expire_minutes: int = 30,
) -> JWTHandler:
"""
配置全局 JWT 处理器
Args:
secret_key: JWT 签名密钥
algorithm: 加密算法
access_token_expire_minutes: Access Token 过期时间分钟
Returns:
配置好的 JWTHandler 实例
"""
global _default_handler
_default_handler = JWTHandler(
secret_key=secret_key,
algorithm=algorithm,
access_token_expire_minutes=access_token_expire_minutes,
)
return _default_handler
def get_jwt_handler() -> JWTHandler:
"""
获取全局 JWT 处理器
Returns:
JWTHandler 实例
Raises:
RuntimeError: 如果尚未配置 JWT 处理器
"""
if _default_handler is None:
raise RuntimeError(
"JWT handler not configured. Call configure_jwt_handler() first."
)
return _default_handler
@@ -0,0 +1,126 @@
"""
密码处理器委托层
此模块作为 Domain (packages.domain.auth.password_hasher) Application 层之间的委托层
隔离 Domain 层对 bcrypt 库的直接依赖
使用方式
from packages.application.auth.password_handler import PasswordHandler, get_password_handler
password_handler = PasswordHandler()
hashed = password_handler.hash_password("my_secure_password")
is_valid = password_handler.verify_password("my_secure_password", hashed)
"""
from typing import Optional, Tuple
from packages.domain.auth.password_hasher import PasswordHasher, PasswordValidator
class PasswordHandler:
"""
密码处理器委托类
委托给 packages.domain.auth.password_hasher 进行实际的密码哈希操作
此层仅负责配置和封装不直接依赖 bcrypt
"""
def __init__(self, rounds: int = 12):
"""
初始化密码处理器
Args:
rounds: bcrypt cost factor默认 12推荐范围 10-14
"""
self._hasher = PasswordHasher(rounds=rounds)
self._validator = PasswordValidator(
min_length=8,
require_uppercase=True,
require_lowercase=True,
require_digit=True,
require_special=False,
)
def hash_password(self, password: str) -> str:
"""
哈希密码
Args:
password: 明文密码
Returns:
bcrypt 哈希字符串
Raises:
ValueError: 密码为空
"""
return self._hasher.hash_password(password)
def verify_password(self, password: str, hashed_password: str) -> bool:
"""
验证密码
Args:
password: 明文密码
hashed_password: 存储的哈希密码
Returns:
True 如果密码正确否则 False
"""
return self._hasher.verify_password(password, hashed_password)
def needs_rehash(self, hashed_password: str) -> bool:
"""
检查哈希是否需要重新计算
Args:
hashed_password: 存储的哈希密码
Returns:
True 如果需要重新哈希
"""
return self._hasher.needs_rehash(hashed_password)
def validate_strength(self, password: str) -> Tuple[bool, Optional[str]]:
"""
验证密码强度
Args:
password: 明文密码
Returns:
(是否有效, 错误信息)
"""
return self._validator.validate(password)
# 默认处理器实例
_default_handler: Optional[PasswordHandler] = None
def configure_password_handler(rounds: int = 12) -> PasswordHandler:
"""
配置全局密码处理器
Args:
rounds: bcrypt cost factor
Returns:
配置好的 PasswordHandler 实例
"""
global _default_handler
_default_handler = PasswordHandler(rounds=rounds)
return _default_handler
def get_password_handler() -> PasswordHandler:
"""
获取全局密码处理器
Returns:
PasswordHandler 实例
"""
global _default_handler
if _default_handler is None:
_default_handler = PasswordHandler()
return _default_handler
+292
View File
@@ -0,0 +1,292 @@
"""
认证集成测试
测试完整的认证流程包括注册登录令牌刷新登出等
"""
import pytest
from fastapi.testclient import TestClient
from apps.api.main import app
client = TestClient(app)
class TestUserRegistration:
"""用户注册集成测试"""
def test_register_with_valid_data(self):
"""测试使用有效数据进行注册"""
response = client.post(
"/api/v1/auth/register",
json={
"email": "newuser@example.com",
"password": "SecurePass123",
"username": "newuser",
"display_name": "New User",
},
)
assert response.status_code == 201
data = response.json()
assert data["email"] == "newuser@example.com"
assert data["username"] == "newuser"
assert data["display_name"] == "New User"
assert "user_id" in data
def test_register_with_invalid_email(self):
"""测试使用无效邮箱进行注册"""
response = client.post(
"/api/v1/auth/register",
json={
"email": "invalid-email",
"password": "SecurePass123",
"username": "testuser",
},
)
assert response.status_code == 422 # Validation error
def test_register_with_weak_password(self):
"""测试使用弱密码进行注册"""
response = client.post(
"/api/v1/auth/register",
json={
"email": "weak@example.com",
"password": "123", # Too short and simple
"username": "weakuser",
},
)
# Should fail validation or business logic
assert response.status_code in [400, 422]
def test_register_duplicate_email(self):
"""测试重复邮箱注册"""
# First registration
client.post(
"/api/v1/auth/register",
json={
"email": "duplicate@example.com",
"password": "SecurePass123",
"username": "user1",
"display_name": "User 1",
},
)
# Second registration with same email
response = client.post(
"/api/v1/auth/register",
json={
"email": "duplicate@example.com",
"password": "SecurePass123",
"username": "user2",
"display_name": "User 2",
},
)
assert response.status_code == 400
assert "already" in response.json()["detail"].lower() or "exists" in response.json()["detail"].lower()
class TestUserLogin:
"""用户登录集成测试"""
def setup_method(self):
"""每个测试前的准备:注册用户"""
client.post(
"/api/v1/auth/register",
json={
"email": "loginuser@example.com",
"password": "SecurePass123",
"username": "loginuser",
"display_name": "Login User",
},
)
def test_login_with_correct_credentials(self):
"""测试使用正确凭据登录"""
response = client.post(
"/api/v1/auth/login",
json={
"email": "loginuser@example.com",
"password": "SecurePass123",
},
)
assert response.status_code == 200
data = response.json()
assert "access_token" in data
assert "refresh_token" in data
assert data["token_type"] == "bearer"
assert data["email"] == "loginuser@example.com"
def test_login_with_wrong_password(self):
"""测试使用错误密码登录"""
response = client.post(
"/api/v1/auth/login",
json={
"email": "loginuser@example.com",
"password": "WrongPassword123",
},
)
assert response.status_code == 401
assert "error" in response.json() or "detail" in response.json()
def test_login_with_nonexistent_email(self):
"""测试使用不存在的邮箱登录"""
response = client.post(
"/api/v1/auth/login",
json={
"email": "nonexistent@example.com",
"password": "AnyPassword123",
},
)
assert response.status_code == 401
def test_login_case_insensitive_email(self):
"""测试邮箱大小写不敏感登录"""
response = client.post(
"/api/v1/auth/login",
json={
"email": "LOGINUSER@EXAMPLE.COM", # Uppercase email
"password": "SecurePass123",
},
)
# Should still work because email is normalized
assert response.status_code == 200
class TestTokenRefresh:
"""令牌刷新集成测试"""
def setup_method(self):
"""每个测试前的准备:注册并登录获取令牌"""
client.post(
"/api/v1/auth/register",
json={
"email": "refresh@example.com",
"password": "SecurePass123",
"username": "refreshuser",
"display_name": "Refresh User",
},
)
response = client.post(
"/api/v1/auth/login",
json={
"email": "refresh@example.com",
"password": "SecurePass123",
},
)
self.refresh_token = response.json().get("refresh_token")
self.access_token = response.json().get("access_token")
def test_refresh_token_success(self):
"""测试成功刷新令牌"""
if not self.refresh_token:
pytest.skip("Refresh token not available")
response = client.post(
"/api/v1/auth/refresh",
json={"refresh_token": self.refresh_token},
)
# If refresh endpoint exists
if response.status_code != 404:
assert response.status_code == 200
data = response.json()
assert "access_token" in data
class TestCurrentUser:
"""当前用户信息集成测试"""
def setup_method(self):
"""每个测试前的准备:注册并登录获取令牌"""
client.post(
"/api/v1/auth/register",
json={
"email": "me@example.com",
"password": "SecurePass123",
"username": "meuser",
"display_name": "Me User",
},
)
response = client.post(
"/api/v1/auth/login",
json={
"email": "me@example.com",
"password": "SecurePass123",
},
)
self.token = response.json()["access_token"]
self.headers = {"Authorization": f"Bearer {self.token}"}
def test_get_current_user_success(self):
"""测试获取当前用户信息成功"""
response = client.get("/api/v1/auth/me", headers=self.headers)
assert response.status_code == 200
data = response.json()
assert data["email"] == "me@example.com"
assert data["username"] == "meuser"
assert "user_id" in data
def test_get_current_user_without_token(self):
"""测试无令牌获取当前用户信息"""
response = client.get("/api/v1/auth/me")
assert response.status_code == 403
def test_get_current_user_with_invalid_token(self):
"""测试使用无效令牌获取当前用户信息"""
response = client.get(
"/api/v1/auth/me",
headers={"Authorization": "Bearer invalid-token"},
)
assert response.status_code == 401
class TestPasswordReset:
"""密码重置集成测试"""
def test_request_password_reset_success(self):
"""测试请求密码重置成功"""
# Register user first
client.post(
"/api/v1/auth/register",
json={
"email": "reset@example.com",
"password": "SecurePass123",
"username": "resetuser",
},
)
response = client.post(
"/api/v1/auth/password/forgot",
json={"email": "reset@example.com"},
)
# Should return 202 Accepted (even if email not sent)
assert response.status_code == 202
def test_request_password_reset_nonexistent_user(self):
"""测试请求不存在的用户密码重置"""
response = client.post(
"/api/v1/auth/password/forgot",
json={"email": "nonexistent@example.com"},
)
# Should still return 202 for security (don't reveal if email exists)
assert response.status_code == 202
if __name__ == "__main__":
pytest.main([__file__, "-v"])