chore: 修复 P2/P3 技术债务 #15
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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"])
|
||||
Reference in New Issue
Block a user