diff --git a/.gitea/workflows/ci-pipeline.yml b/.gitea/workflows/ci-pipeline.yml index 04902aef0..de0a655b5 100755 --- a/.gitea/workflows/ci-pipeline.yml +++ b/.gitea/workflows/ci-pipeline.yml @@ -1187,6 +1187,8 @@ jobs: COSYVOICE_API_KEY: ${{ secrets.COSYVOICE_API_KEY }} DASHSCOPE_API_KEY: ${{ secrets.DASHSCOPE_API_KEY }} MEDIAKIT_API_KEY: ${{ secrets.MEDIAKIT_API_KEY }} + WECHAT_APP_ID: ${{ secrets.WECHAT_APP_ID }} + WECHAT_APP_SECRET: ${{ secrets.WECHAT_APP_SECRET }} run: | set -eu echo "Rendering .env from template + secrets..." diff --git a/alembic/versions/066_asset_client_upload_id_ingest_job_asset_id.py b/alembic/versions/066_asset_client_upload_id_ingest_job_asset_id.py new file mode 100644 index 000000000..6156bbfc1 --- /dev/null +++ b/alembic/versions/066_asset_client_upload_id_ingest_job_asset_id.py @@ -0,0 +1,34 @@ +"""add client_upload_id to assets and asset_id to ingest_jobs + +Issue #1714:上传 complete 幂等 + worker 转码回写关联。 +- assets.client_upload_id:客户端幂等 token(complete 去重) +- ingest_jobs.asset_id:complete 阶段创建的占位 asset id(worker 回写关联, + 防止 HEVC 转码改写 storage_key 后找不到占位而兜底新建 READY 记录) + +Revision ID: 066_upload_idempotency +Revises: 065_dup_record_sim_match +Create Date: 2026-09-05 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "066_upload_idempotency" +down_revision = "065_dup_record_sim_match" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column("assets", sa.Column("client_upload_id", sa.String(64), nullable=True)) + op.create_index("ix_assets_client_upload_id", "assets", ["client_upload_id"]) + op.add_column("ingest_jobs", sa.Column("asset_id", sa.String(36), nullable=False, server_default="")) + op.create_index("ix_ingest_jobs_asset_id", "ingest_jobs", ["asset_id"]) + + +def downgrade() -> None: + op.drop_index("ix_ingest_jobs_asset_id", table_name="ingest_jobs") + op.drop_column("ingest_jobs", "asset_id") + op.drop_index("ix_assets_client_upload_id", table_name="assets") + op.drop_column("assets", "client_upload_id") diff --git a/alembic/versions/067_celery_task_id_revoke.py b/alembic/versions/067_celery_task_id_revoke.py new file mode 100644 index 000000000..bd2b68a40 --- /dev/null +++ b/alembic/versions/067_celery_task_id_revoke.py @@ -0,0 +1,35 @@ +"""add celery_task_id to generation_tasks and ingest_jobs + +Issue #1714:孤儿恢复/超时清理撤销队列消息。 +- generation_tasks.celery_task_id:入队时记录的 Celery 消息 ID,清理时 revoke +- ingest_jobs.celery_task_id:同上(素材转码任务) + +Revision ID: 067_celery_task_id +Revises: 066_upload_idempotency +Create Date: 2026-09-05 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "067_celery_task_id" +down_revision = "066_upload_idempotency" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column( + "generation_tasks", + sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""), + ) + op.add_column( + "ingest_jobs", + sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""), + ) + + +def downgrade() -> None: + op.drop_column("ingest_jobs", "celery_task_id") + op.drop_column("generation_tasks", "celery_task_id") diff --git a/alembic/versions/068_user_profile_completed.py b/alembic/versions/068_user_profile_completed.py new file mode 100644 index 000000000..61d9f6a90 --- /dev/null +++ b/alembic/versions/068_user_profile_completed.py @@ -0,0 +1,26 @@ +"""add profile_completed to users + +Issue #1718:微信新用户首次登录需设置昵称(PATCH /auth/me)。 +- users.profile_completed:资料是否已完善;存量行默认 True(不触发引导), + 微信新建用户在应用层置 False。 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "068_user_profile_completed" +down_revision = "067_celery_task_id" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column( + "users", + sa.Column("profile_completed", sa.Boolean(), nullable=False, server_default=sa.text("true")), + ) + + +def downgrade() -> None: + op.drop_column("users", "profile_completed") diff --git a/apps/api/app/api/routes/auth.py b/apps/api/app/api/routes/auth.py index 38438f931..44cdaa6b8 100755 --- a/apps/api/app/api/routes/auth.py +++ b/apps/api/app/api/routes/auth.py @@ -1,4 +1,5 @@ """ +from __future__ import annotations Canonical authentication API routes. The route layer is intentionally thin: repository construction lives in @@ -13,9 +14,9 @@ import jwt from app.auth import AuthenticatedUser, blacklist_token, get_current_user from app.config import settings from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository -from fastapi import APIRouter, Depends, Header, HTTPException, status +from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer -from pydantic import BaseModel, EmailStr +from pydantic import BaseModel, EmailStr, field_validator from packages.adapters.redis import NoopSessionStore from packages.adapters.smtp import NoopEmailService @@ -84,6 +85,23 @@ class CurrentUserResponse(BaseModel): phone: str = "" phone_verified: bool = False binding_complete: bool = False + wechat_bound: bool = False + profile_completed: bool = True + + +class UserProfileResponse(BaseModel): + """用户资料负载(PATCH /me、绑定/解绑接口复用;字段与 GET /auth/me 一致,前端 normalizeUser 直接消费)""" + + user_id: str + email: str + username: str + display_name: str + email_verified: bool + phone: str = "" + phone_verified: bool = False + binding_complete: bool = False + wechat_bound: bool = False + profile_completed: bool = True class PasswordResetRequestModel(BaseModel): @@ -272,9 +290,52 @@ async def get_current_user_info( phone=user.phone or "", phone_verified=user.phone_verified, binding_complete=binding_complete, + wechat_bound=bool(user.wechat_openid), + profile_completed=user.profile_completed, ) +class UpdateProfileRequest(BaseModel): + """更新个人资料请求(当前仅支持昵称)""" + + display_name: str + + @field_validator("display_name") + @classmethod + def _validate_display_name(cls, v: str) -> str: + name = (v or "").strip() + if not name: + raise ValueError("昵称不能为空白") + if len(name) > 20: + raise ValueError("昵称长度需在 1-20 个字符之间") + return name + + +class UpdateProfileResponse(BaseModel): + """更新资料响应:前端 normalizeUser(response.user) 直接消费""" + + user: UserProfileResponse + + +@router.patch("/me", response_model=UpdateProfileResponse) +async def update_current_user_profile( + request: UpdateProfileRequest, + current_user: AuthenticatedUser = Depends(get_current_user), + user_repository: UserRepository = Depends(get_user_repository), +) -> UpdateProfileResponse: + """更新当前登录用户昵称(微信新用户首次设置昵称后置 profile_completed=True)。""" + user = current_user.user + user.display_name = request.display_name # 已 strip(validator) + if not user.profile_completed: + user.profile_completed = True + user_repository.save(user) + + logger.info("[资料更新] 用户 %s 更新昵称,profile_completed=%s", user.id, user.profile_completed) + # 重新读取,确保返回的是持久化后的最新状态 + fresh = user_repository.find_by_id(user.id) or user + return UpdateProfileResponse(user=_user_profile(fresh)) + + class _NoopSessionStore(NoopSessionStore): pass @@ -426,6 +487,7 @@ async def get_wechat_auth_url() -> WechatAuthUrlResponse: @router.post("/wechat/callback", response_model=WechatLoginResponse) async def wechat_callback( request: WechatCallbackRequest, + http_request: Request, user_repository: UserRepository = Depends(get_user_repository), ) -> WechatLoginResponse: """微信登录回调处理""" @@ -433,11 +495,30 @@ async def wechat_callback( from packages.application.auth.wechat_sync_use_case import WechatSyncRequest as SyncRequest from packages.application.auth.wechat_sync_use_case import WechatSyncUseCase + # 回调可观测性:记录 UA(区分微信内置浏览器 MicroMessenger)与 state, + # 便于排查"停留 open.weixin.qq.com / 回调失败"类问题(#1718) + user_agent = http_request.headers.get("User-Agent", "") + is_wechat_browser = "MicroMessenger" in user_agent + logger.info( + "[微信回调] 收到回调: state=%s code_len=%d UA=%r 微信内置浏览器=%s", + (request.state or "")[:8], + len(request.code or ""), + user_agent[:200], + is_wechat_browser, + ) + # 1. 用 code 换微信用户信息 oauth_service = get_wechat_oauth_service() wechat_user, err = oauth_service.handle_callback(request.code, request.state) if err: + # state 校验失败 / 微信 errcode 等错误原文已在 service 内 log,这里带上 UA 上下文 + logger.warning("[微信回调] 处理失败: err=%s 微信内置浏览器=%s", err, is_wechat_browser) raise HTTPException(status_code=400, detail=err) + logger.info( + "[微信回调] state 校验通过,微信用户信息获取成功: openid=%s unionid=%s", + wechat_user.openid[:8] if wechat_user.openid else "", + bool(wechat_user.unionid), + ) # 2. 同步登录/注册(复用 wechat-sync 逻辑) use_case = WechatSyncUseCase(user_repository=user_repository) @@ -472,6 +553,109 @@ async def wechat_callback( ) +# ==================== 微信账号绑定/解绑(已登录用户) ==================== + + +class WechatBindUrlResponse(BaseModel): + auth_url: str + state: str + + +class WechatBindCompleteRequest(BaseModel): + code: str + state: str = "" + + +class WechatBindCompleteResponse(BaseModel): + success: bool + user: UserProfileResponse + + +class WechatUnbindResponse(BaseModel): + success: bool + + +def _user_profile(user) -> UserProfileResponse: + binding_complete = bool( + user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email + ) + return UserProfileResponse( + user_id=user.id, + email=user.email, + username=user.username, + display_name=user.display_name, + email_verified=user.email_verified, + phone=user.phone or "", + phone_verified=user.phone_verified, + binding_complete=binding_complete, + wechat_bound=bool(user.wechat_openid), + profile_completed=user.profile_completed, + ) + + +@router.get("/wechat/bind/url", response_model=WechatBindUrlResponse) +async def get_wechat_bind_url( + current_user: AuthenticatedUser = Depends(get_current_user), +) -> WechatBindUrlResponse: + """获取微信绑定授权链接(已登录用户场景)。state 经 Redis 存储做 CSRF 校验。""" + from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service + + oauth_service = get_wechat_oauth_service() + auth_url, state = oauth_service.generate_auth_url() + logger.info("[微信绑定] 用户 %s 请求绑定授权链接", current_user.user.id) + return WechatBindUrlResponse(auth_url=auth_url, state=state) + + +@router.post("/wechat/bind", response_model=WechatBindCompleteResponse) +async def wechat_bind( + request: WechatBindCompleteRequest, + current_user: AuthenticatedUser = Depends(get_current_user), + user_repository: UserRepository = Depends(get_user_repository), +) -> WechatBindCompleteResponse: + """微信绑定完成:扫码回调后用 code 换 openid,绑定到当前登录账号(不创建新用户)。""" + from packages.application.auth.wechat_bind_use_case import WechatBindRequest, WechatBindUseCase + from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service + + oauth_service = get_wechat_oauth_service() + wechat_user, err = oauth_service.handle_callback(request.code, request.state) + if err: + logger.warning("[微信绑定] 用户 %s 换取微信信息失败: %s", current_user.user.id, err) + raise HTTPException(status_code=400, detail=err) + + use_case = WechatBindUseCase(user_repository=user_repository) + result, error, http_status = use_case.bind( + WechatBindRequest( + user_id=current_user.user.id, + openid=wechat_user.openid, + unionid=wechat_user.unionid or "", + ) + ) + if error: + logger.warning("[微信绑定] 用户 %s 绑定失败: %s", current_user.user.id, error) + raise HTTPException(status_code=http_status, detail=error) + + logger.info("[微信绑定] 用户 %s 绑定成功 openid=%s", current_user.user.id, wechat_user.openid[:8]) + return WechatBindCompleteResponse(success=True, user=_user_profile(result.user)) + + +@router.delete("/wechat/bind", response_model=WechatUnbindResponse) +async def wechat_unbind( + current_user: AuthenticatedUser = Depends(get_current_user), + user_repository: UserRepository = Depends(get_user_repository), +) -> WechatUnbindResponse: + """解绑微信:需账号仍有其他登录方式(密码/手机/真实邮箱),否则拒绝。""" + from packages.application.auth.wechat_bind_use_case import WechatUnbindUseCase + + use_case = WechatUnbindUseCase(user_repository=user_repository) + result, error, http_status = use_case.unbind(current_user.user.id) + if error: + logger.warning("[微信解绑] 用户 %s 解绑失败: %s", current_user.user.id, error) + raise HTTPException(status_code=http_status, detail=error) + + logger.info("[微信解绑] 用户 %s 解绑成功", current_user.user.id) + return WechatUnbindResponse(success=True) + + # ==================== 验证码 & 绑定 ==================== diff --git a/apps/api/app/api/routes/chunked_upload.py b/apps/api/app/api/routes/chunked_upload.py index 1db5a1647..6d25a05fb 100644 --- a/apps/api/app/api/routes/chunked_upload.py +++ b/apps/api/app/api/routes/chunked_upload.py @@ -14,6 +14,7 @@ from typing import Any from uuid import uuid4 from app.api.routes._helpers import require_project_and_library +from app.api.routes.upload import _persist_celery_task_id from app.auth import AuthenticatedUser, get_current_user from app.core.celery_app import celery_app from app.core.storage import OSSStorageService, get_storage_service @@ -381,7 +382,8 @@ async def complete_chunked_upload( file_hash=request.file_hash, ) ) - celery_app.send_task("worker.ingest_asset", args=[job.id]) + celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id]) + _persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", "")) # Update metadata status meta["status"] = "completed" diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index c40fdf343..e8d0b8a53 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -14,6 +14,7 @@ from app.core.task_enqueue import ( USER_PENDING_LIMIT, GlobalQueueFull, UserPendingLimitExceeded, + build_rate_limit_detail, safe_enqueue_generation_task, ) from app.dependencies import ( @@ -23,6 +24,7 @@ from app.dependencies import ( get_generation_task_repository, ) from app.schemas.generation_task import ( + BatchPreviewGenerationTaskResponse, CreatePreviewGenerationTaskRequest, PreviewGenerationTaskResponse, ) @@ -193,11 +195,19 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG if started_at and completed_at: generate_duration = (completed_at - started_at).total_seconds() + title_cfg = getattr(task, "title_config", None) + title_cfg = title_cfg if isinstance(title_cfg, dict) else {} + extra_meta = getattr(task, "extra_meta", None) + extra_meta = extra_meta if isinstance(extra_meta, dict) else {} + voice_library_id = getattr(task, "voice_library_id", "") or "" + if not isinstance(voice_library_id, str): + voice_library_id = str(voice_library_id) if voice_library_id else "" return PreviewGenerationTaskResponse( task_id=task.id, status=task.status.value if hasattr(task.status, "value") else str(task.status), progress=float(task.progress or 0.0), is_preview=bool(getattr(task, "is_preview", True)), + variant_index=int(extra_meta.get("variant_index", 0) or 0), resolution=getattr(task, "resolution", "") or "", video_url=video_url, duration=duration, @@ -206,6 +216,8 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG transition_count=transition_count, material_usage=material_usage, error_message=task.error_message or "", + title_text=str(title_cfg.get("text", "") or ""), + voice_library_id=voice_library_id, created_at=task.created_at, started_at=started_at, finished_at=completed_at, @@ -213,50 +225,100 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG ) -@router.post("/preview", response_model=PreviewGenerationTaskResponse, status_code=201) +def _resolve_preview_edit_plan_id( + *, + request: CreatePreviewGenerationTaskRequest, + task, + db: Session, + user_id: str, +) -> str: + """确定任务关联的编辑计划ID:优先前端传入,否则按 template_id+user 兜底查找。""" + if task.source_edit_plan_id: + return task.source_edit_plan_id + if not request.template_id: + return "" + try: + from packages.adapters.sqlalchemy_impl.edit_plan_repository import ( + SQLAlchemyEditPlanRepository, + ) + + _plan_repo = SQLAlchemyEditPlanRepository(db) + _plans = _plan_repo.list_by_template(request.template_id, limit=20) + for _p in _plans: + if (_p.created_by_user_id or "") == user_id: + logger.info( + "[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s", + task.id, + _p.id, + ) + return _p.id + except Exception: + logger.warning( + "[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s", + task.id, + exc_info=True, + ) + return "" + + +def _variant_value(values: list[str], index: int, fallback: str = "") -> str: + """从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。""" + if not values: + return fallback + if len(values) == 1: + return values[0] + return values[index] if index < len(values) else fallback + + +@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201) def create_preview_generation_task( request: CreatePreviewGenerationTaskRequest, authenticated_user: AuthenticatedUser = Depends(get_current_user), generation_task_repository=Depends(get_generation_task_repository), db: Session = Depends(get_db_session), asset_repo=Depends(get_asset_repository), -) -> PreviewGenerationTaskResponse: - """创建预览生成任务。 +) -> BatchPreviewGenerationTaskResponse: + """创建预览生成任务(支持批量)。 - 预览渲染品质与正式生成一致(1080p, CRF 23, medium preset),确认生成时可直接复用预览产物。 - - Args: - request: 预览任务创建请求(template_id + asset_ids 等) + preview_count=1 时行为与旧版完全一致(创建 1 个任务); + preview_count=N 时一次创建 N 个独立变体任务: + - 每个变体克隆独立编辑计划(独立 clips、独立随机素材起点),N 个预览内容互不相同 + - 每个变体拥有独立 task_id / 状态 / 预览视频 URL,前端按 task_id 分别轮询 + - 标题样式(font/color/position 等)全局共用;标题文字/配音/封面可按变体独立 + (titles[] / voice_library_ids[] / cover_urls[],长度1=共用,长度N=独立) Returns: - 201 + 预览任务详情 + 201 + 变体任务数组 {items: [...], total: N} """ user_id = authenticated_user.user.id + count = max(1, request.preview_count) logger.info( "[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d, preview_count=%d", user_id, request.template_id, len(request.asset_ids), - request.preview_count, + count, ) - # 预检查队列限流 + # 预检查队列限流(按变体总数计) try: user_pending = generation_task_repository.count_pending_by_user(user_id) global_pending = generation_task_repository.count_pending_total() - if user_pending + 1 > USER_PENDING_LIMIT: - raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending + 1, limit=USER_PENDING_LIMIT) - if global_pending + 1 > GLOBAL_PENDING_LIMIT: - raise GlobalQueueFull(pending_count=global_pending + 1, limit=GLOBAL_PENDING_LIMIT) + if user_pending + count > USER_PENDING_LIMIT: + raise UserPendingLimitExceeded( + user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT + ) + if global_pending + count > GLOBAL_PENDING_LIMIT: + raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT) except UserPendingLimitExceeded as e: raise HTTPException( status_code=429, - detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交", + detail=build_rate_limit_detail(e, generation_task_repository, scope="user"), ) from e except GlobalQueueFull as e: raise HTTPException( status_code=503, - detail="系统繁忙,请稍后再试", + detail=build_rate_limit_detail(e, generation_task_repository, scope="global"), ) from e # 确定视频比例:优先前端传入,否则从模板 mode 推断 @@ -273,14 +335,11 @@ def create_preview_generation_task( w, h = int(parts[0]), int(parts[1]) base = 1920 if w < h: - # 竖屏 output_width = round(base * w / h) output_height = base else: - # 横屏 output_width = base output_height = round(base * h / w) - # 对齐到偶数 output_width = output_width - output_width % 2 output_height = output_height - output_height % 2 except (ValueError, ZeroDivisionError): @@ -289,42 +348,71 @@ def create_preview_generation_task( logger.info( "[预览生成] 分辨率: video_ratio=%s → %s (%dx%d)", - video_ratio, resolution, output_width, output_height, + video_ratio, + resolution, + output_width, + output_height, ) - # 从模板读取 editing_mode / mode 作为 strategy_id(渲染 pipeline 的 mode 参数) strategy_id = _resolve_strategy_id_from_template(request.template_id, db, user_id) - - title_config = request.title_config or {} + base_title_config = request.title_config or {} use_case = CreateGenerationTaskUseCase(generation_task_repository) + # ── 预创建第一个任务,仅用于解析源编辑计划(不落库为最终任务)── + # 先创建一个临时任务拿到 task 对象上下文,实际 N 个任务在循环中统一创建; + # 为保持与旧版一致的源 plan 解析逻辑,先创建任务0、解析源 plan, + # 再预克隆 N 个变体 plan,最后重建任务关联。 + # 简化实现:直接创建全部任务,plan 关联在创建后、入队前完成。 + + created_tasks: list = [] + variant_plan_ids: list[str] = [] # 每个变体最终关联的 plan_id(按变体顺序) + try: - task = use_case.execute( - CreateGenerationTaskCommand( - project_id="", - asset_library_id="", - strategy_id=strategy_id, - voice_library_id=request.voice_library_id, - template_id=request.template_id, - asset_ids=list(request.asset_ids), - title_ids=list(request.title_ids), - voice_ids=list(request.voice_ids), - created_by_user_id=user_id, - source_edit_plan_id=request.source_edit_plan_id, - asset_select_mode="", - batch_id="", - video_title=request.video_title, - resolution=resolution, - bgm_config=request.bgm_config or {}, - auto_retry_enabled=False, - auto_retry_max=0, - is_preview=True, - title_config=title_config, - output_width=output_width, - output_height=output_height, + for variant_index in range(count): + # 变体独立标题文字:titles[] 覆盖 title_config.text + variant_title_text = _variant_value(request.titles, variant_index, "") + variant_title_config = dict(base_title_config) + if variant_title_text.strip(): + variant_title_config["text"] = variant_title_text.strip() + + # 变体独立配音 + variant_voice_library_id = _variant_value( + request.voice_library_ids, variant_index, request.voice_library_id ) - ) + + task = use_case.execute( + CreateGenerationTaskCommand( + project_id="", + asset_library_id="", + strategy_id=strategy_id, + voice_library_id=variant_voice_library_id, + template_id=request.template_id, + asset_ids=list(request.asset_ids), + title_ids=list(request.title_ids), + voice_ids=list(request.voice_ids), + created_by_user_id=user_id, + source_edit_plan_id=request.source_edit_plan_id, + asset_select_mode="", + batch_id="", + video_title=request.video_title, + resolution=resolution, + bgm_config=request.bgm_config or {}, + auto_retry_enabled=False, + auto_retry_max=0, + is_preview=True, + title_config=variant_title_config, + output_width=output_width, + output_height=output_height, + ) + ) + task.extra_meta["variant_index"] = variant_index + + # 解析源编辑计划(前端传入或按模板兜底查找) + source_plan_id = _resolve_preview_edit_plan_id(request=request, task=task, db=db, user_id=user_id) + task.source_edit_plan_id = source_plan_id + generation_task_repository.update(task) + created_tasks.append(task) except ValueError as e: logger.warning("[预览生成] 创建失败: %s", e) raise HTTPException(status_code=400, detail=str(e)) from e @@ -332,93 +420,128 @@ def create_preview_generation_task( logger.error("[预览生成] 创建失败: %s", e, exc_info=True) raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e - # 关联编辑计划:如果前端未传 source_edit_plan_id,通过 template_id + user_id 查找 - if not task.source_edit_plan_id and request.template_id: - try: - from packages.adapters.sqlalchemy_impl.edit_plan_repository import ( - SQLAlchemyEditPlanRepository, - ) - - _plan_repo = SQLAlchemyEditPlanRepository(db) - _plans = _plan_repo.list_by_template(request.template_id, limit=20) - for _p in _plans: - if (_p.created_by_user_id or "") == user_id: - task.source_edit_plan_id = _p.id - generation_task_repository.update(task) - logger.info( - "[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s", - task.id, - _p.id, - ) - break - except Exception: - logger.warning( - "[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s", - task.id, - exc_info=True, - ) - - # 每条预览都关联独立克隆 plan:多预览前端为 N 次并发调用,若共用同一 plan - # 则 N 条预览片段完全相同;克隆时片段起点按持久化历史区间重算(含受控复用), - # 保证各预览版本内容不同 - if task.source_edit_plan_id: + # ── 克隆独立变体 plan:N 个预览全部克隆(预览不污染源 plan)── + # 源 plan 不存在(无编辑历史)时各任务走自身随机选片流程,不克隆。 + source_plan_id = created_tasks[0].source_edit_plan_id if created_tasks else "" + if source_plan_id: try: from app.services.edit_plan_service import EditPlanService _plan_svc = EditPlanService(db) - _preview_plan = _plan_svc.clone_plan_for_variant( - task.source_edit_plan_id, - created_by_user_id=user_id, - name_suffix="预览变体", - ) - task.source_edit_plan_id = _preview_plan.id - generation_task_repository.update(task) - logger.info( - "[预览生成] 预览关联独立克隆 plan: task_id=%s clone_plan_id=%s", - task.id, - _preview_plan.id, - ) - except Exception as clone_err: - # 不退回共用原 plan(否则多条预览内容相同,违反去重诉求): - # 标记任务失败并中断,前端可重新发起预览 - logger.error( - "[预览生成] 克隆预览变体 plan 失败,任务标记失败: task_id=%s error=%s", - task.id, - clone_err, - exc_info=True, - ) - _mark_task_failed(generation_task_repository, task, "预览变体计划创建失败") + for variant_index in range(count): + last_err: Exception | None = None + variant_plan = None + for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动 + try: + variant_plan = _plan_svc.clone_plan_for_variant( + source_plan_id, + created_by_user_id=user_id, + name_suffix=f"预览变体{variant_index + 1}" if count > 1 else "预览变体", + ) + break + except Exception as clone_err: # noqa: PERF203 + last_err = clone_err + logger.warning( + "[预览生成] 克隆变体 plan 失败(尝试%d/2): variant=%d error=%s", + _attempt + 1, + variant_index, + clone_err, + exc_info=True, + ) + if variant_plan is None: + logger.error( + "[预览生成] 克隆预览变体 plan 重试仍失败: variant=%d source=%s", + variant_index, + source_plan_id, + exc_info=last_err, + ) + # 标记已创建任务失败 + for t in created_tasks: + _mark_task_failed(generation_task_repository, t, "预览变体计划创建失败") + raise HTTPException( + status_code=500, + detail="创建预览任务失败:无法生成独立剪辑计划,请重试", + ) from last_err + variant_plan_ids.append(variant_plan.id) + except HTTPException: + raise + except Exception as e: + logger.error("[预览生成] 克隆变体 plan 异常: %s", e, exc_info=True) + for t in created_tasks: + _mark_task_failed(generation_task_repository, t, "预览变体计划创建失败") raise HTTPException( status_code=500, detail="创建预览任务失败:无法生成独立剪辑计划,请重试", - ) from clone_err + ) from e - # 入队执行;若入队失败则标记任务为 failed 避免僵尸数据 - try: - if not safe_enqueue_generation_task( - task, - generation_task_repository, - user_id=user_id, - log_prefix="[预览生成]", - log_task_status=True, - ): - logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id) - _mark_task_failed(generation_task_repository, task, "任务入队失败") - raise HTTPException(status_code=500, detail="任务入队失败,请稍后重试") - except UserPendingLimitExceeded as e: - _mark_task_failed(generation_task_repository, task, "待处理任务超限") - raise HTTPException( - status_code=429, - detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交", - ) from None - except GlobalQueueFull: - _mark_task_failed(generation_task_repository, task, "系统队列已满") + # 关联变体 plan 并回写标题配置 + for variant_index, task in enumerate(created_tasks): + if variant_plan_ids: + task.source_edit_plan_id = variant_plan_ids[variant_index] + generation_task_repository.update(task) + # 回写变体标题到 plan config(worker 渲染时从 plan 读取 title 配置) + if task.source_edit_plan_id and (task.title_config or {}).get("text", "").strip(): + try: + from app.api.routes.generation_tasks import _writeback_edit_plan_config + + _writeback_edit_plan_config( + plan_id=task.source_edit_plan_id, + task_id=task.id, + title_config=task.title_config, + db=db, + ) + except Exception: + logger.warning( + "[预览生成] 回写标题配置失败(不影响主流程): task_id=%s", + task.id, + exc_info=True, + ) + + # ── 入队 ── + responses: list[PreviewGenerationTaskResponse] = [] + rate_limit_exc: Exception | None = None # 记录首个限流异常,全部失败时返回结构化提示 + for variant_index, task in enumerate(created_tasks): + try: + enqueued = safe_enqueue_generation_task( + task, + generation_task_repository, + user_id=user_id, + log_prefix=f"[预览生成][变体{variant_index + 1}]", + log_task_status=True, + ) + if not enqueued: + logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id) + _mark_task_failed(generation_task_repository, task, "任务入队失败") + except UserPendingLimitExceeded as e: + _mark_task_failed(generation_task_repository, task, "待处理任务超限") + rate_limit_exc = rate_limit_exc or e + except GlobalQueueFull as e: + _mark_task_failed(generation_task_repository, task, "系统队列已满") + rate_limit_exc = rate_limit_exc or e + except Exception: + logger.exception("[预览生成] 入队异常: task_id=%s", task.id) + _mark_task_failed(generation_task_repository, task, "任务入队异常") + # enqueue 会原地更新 task 状态/进度,直接用 task 构造响应 + responses.append(_to_preview_response(task)) + + # 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败") + if all(r.status == "failed" for r in responses) and rate_limit_exc is not None: + if isinstance(rate_limit_exc, UserPendingLimitExceeded): + raise HTTPException( + status_code=429, + detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="user"), + ) raise HTTPException( status_code=503, - detail="系统繁忙,请稍后再试", - ) from None + detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="global"), + ) - return _to_preview_response(task) + logger.info( + "[预览生成] 创建完成: %d 个变体任务, task_ids=%s", + len(responses), + [r.task_id for r in responses], + ) + return BatchPreviewGenerationTaskResponse(items=responses, total=len(responses)) @router.get("/preview/{task_id}", response_model=PreviewGenerationTaskResponse) diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index d5fbad1f1..b0e01892f 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -10,6 +10,7 @@ from app.core.task_enqueue import ( USER_PENDING_LIMIT, GlobalQueueFull, UserPendingLimitExceeded, + build_rate_limit_detail, safe_enqueue_generation_task, ) from app.dependencies import ( @@ -47,6 +48,15 @@ logger = logging.getLogger(__name__) router = APIRouter() +def _variant_value(values: list[str], index: int, fallback: str = "") -> str: + """从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。""" + if not values: + return fallback + if len(values) == 1: + return values[0] + return values[index] if index < len(values) else fallback + + def _to_generation_task_response(task) -> GenerationTaskResponse: return GenerationTaskResponse( id=task.id, @@ -408,12 +418,12 @@ def create_generation_task( except UserPendingLimitExceeded as e: raise HTTPException( status_code=429, - detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交", + detail=build_rate_limit_detail(e, generation_task_repository, scope="user"), ) from e except GlobalQueueFull as e: raise HTTPException( status_code=503, - detail="系统繁忙,请稍后再试", + detail=build_rate_limit_detail(e, generation_task_repository, scope="global"), ) from e # 画中画已下线:strategy_id 中的 pip/voice_pip 统一映射为 one_take @@ -472,12 +482,21 @@ def create_generation_task( if task_index > 0 and variant_plan_ids: effective_plan_id = variant_plan_ids[task_index - 1] + # 变体级独立配置:titles[]/voice_library_ids[]/cover_urls[] + # 长度1=所有变体共用,长度=count=每个变体独立,空数组=回退单值字段 + variant_title_text = _variant_value(request.titles, task_index, "") + variant_title_config = dict(request.title_config or {}) + if variant_title_text.strip(): + variant_title_config["text"] = variant_title_text.strip() + variant_voice_library_id = _variant_value(request.voice_library_ids, task_index, request.voice_library_id) + variant_cover_url = _variant_value(request.cover_urls, task_index, request.cover_url) + task = use_case.execute( CreateGenerationTaskCommand( project_id=project_id, asset_library_id=asset_library_id, strategy_id=effective_strategy_id, - voice_library_id=request.voice_library_id, + voice_library_id=variant_voice_library_id, template_id=request.template_id, asset_ids=resolved_asset_ids, title_ids=request.title_ids, @@ -495,10 +514,12 @@ def create_generation_task( source_task_id=request.source_task_id, output_width=request.output_width, output_height=request.output_height, - cover_url=request.cover_url, - title_config=request.title_config or {}, + cover_url=variant_cover_url, + title_config=variant_title_config, ) ) + # 变体序号写入 extra_meta(响应/排查时可辨识) + task.extra_meta["variant_index"] = task_index try: # 兜底关联编辑计划:前端未传 source_edit_plan_id 时, # 通过 template_id + user_id 在 DB 层直接查找最新的 plan。 @@ -533,13 +554,13 @@ def create_generation_task( # 回写 plan.config:必须在 enqueue 之前执行, # 确保 worker 读取 plan 时 config 中已包含 generation_task_id。 - # 只在首个任务时回写一次,避免批量生成时循环覆盖。 + # 批量场景下每个变体关联独立 plan,需各自回写自己的变体标题配置。 _effective_plan_id = task.source_edit_plan_id - if _effective_plan_id and len(created_tasks) == 0: + if _effective_plan_id: _writeback_edit_plan_config( plan_id=_effective_plan_id, task_id=task.id, - title_config=request.title_config, + title_config=variant_title_config, db=db, ) @@ -559,7 +580,7 @@ def create_generation_task( if not created_tasks: raise HTTPException( status_code=429, - detail="您的待处理任务过多,请等待完成后再提交", + detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"), ) from _e break except GlobalQueueFull as _e: @@ -567,7 +588,7 @@ def create_generation_task( if not created_tasks: raise HTTPException( status_code=503, - detail="系统繁忙,请稍后再试", + detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"), ) from _e break except HTTPException: @@ -693,15 +714,15 @@ def confirm_generation( log_task_status=True, ): logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id) - except UserPendingLimitExceeded: + except UserPendingLimitExceeded as _e: raise HTTPException( status_code=429, - detail="您的待处理任务过多,请等待完成后再提交", + detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"), ) from None - except GlobalQueueFull: + except GlobalQueueFull as _e: raise HTTPException( status_code=503, - detail="系统繁忙,请稍后再试", + detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"), ) from None return BatchGenerationTaskResponse( @@ -783,12 +804,24 @@ def retry_generation_task( if user_pending >= USER_PENDING_LIMIT: raise HTTPException( status_code=429, - detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交", + detail=build_rate_limit_detail( + UserPendingLimitExceeded( + user_id=user_id, + pending_count=user_pending, + limit=USER_PENDING_LIMIT, + ), + generation_task_repository, + scope="user", + ), ) if global_pending >= GLOBAL_PENDING_LIMIT: raise HTTPException( status_code=503, - detail="系统繁忙,请稍后再试", + detail=build_rate_limit_detail( + GlobalQueueFull(pending_count=global_pending, limit=GLOBAL_PENDING_LIMIT), + generation_task_repository, + scope="global", + ), ) use_case = CreateGenerationTaskUseCase(generation_task_repository) @@ -823,15 +856,15 @@ def retry_generation_task( log_task_status=True, ): logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id) - except UserPendingLimitExceeded: + except UserPendingLimitExceeded as _e: raise HTTPException( status_code=429, - detail="您的待处理任务过多,请等待完成后再提交", + detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"), ) from None - except GlobalQueueFull: + except GlobalQueueFull as _e: raise HTTPException( status_code=503, - detail="系统繁忙,请稍后再试", + detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"), ) from None return _to_generation_task_response(retried) diff --git a/apps/api/app/api/routes/ingest_jobs.py b/apps/api/app/api/routes/ingest_jobs.py index 791be8964..ecb6f6a30 100644 --- a/apps/api/app/api/routes/ingest_jobs.py +++ b/apps/api/app/api/routes/ingest_jobs.py @@ -43,7 +43,13 @@ def submit_ingest_job( ) ) - celery_app.send_task("worker.ingest_asset", args=[job.id]) + celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id]) + if getattr(celery_result, "id", ""): + try: + job.celery_task_id = celery_result.id + ingest_job_repository.update(job) + except Exception: # noqa: BLE001 + pass return IngestJobResponse( id=job.id, diff --git a/apps/api/app/api/routes/task_center.py b/apps/api/app/api/routes/task_center.py index 63a75a77a..cc5188457 100755 --- a/apps/api/app/api/routes/task_center.py +++ b/apps/api/app/api/routes/task_center.py @@ -375,7 +375,13 @@ def retry_project_task( storage_key=job.storage_key, ) ) - celery_app.send_task("worker.ingest_asset", args=[retried.id]) + celery_result = celery_app.send_task("worker.ingest_asset", args=[retried.id]) + if getattr(celery_result, "id", ""): + try: + retried.celery_task_id = celery_result.id + ingest_job_repository.update(retried) + except Exception: # noqa: BLE001 + pass return ProjectTaskResponse( id=f"ingest:{retried.id}", task_type="ingest", diff --git a/apps/api/app/api/routes/upload.py b/apps/api/app/api/routes/upload.py index 58269c35c..08aed4e7b 100644 --- a/apps/api/app/api/routes/upload.py +++ b/apps/api/app/api/routes/upload.py @@ -85,12 +85,22 @@ def _infer_mime_type_from_storage_key(storage_key: str) -> str: """从 storage_key 推断 MIME 类型(与 worker 端保持一致)。""" lower_filename = storage_key.rsplit("/", 1)[-1].lower() _MIME_MAP = { - ".mov": "video/quicktime", ".mp4": "video/mp4", ".avi": "video/x-msvideo", - ".mkv": "video/x-matroska", ".webm": "video/webm", - ".png": "image/png", ".gif": "image/gif", ".bmp": "image/bmp", - ".svg": "image/svg+xml", ".jpg": "image/jpeg", ".jpeg": "image/jpeg", - ".mp3": "audio/mpeg", ".wav": "audio/wav", ".ogg": "audio/ogg", - ".flac": "audio/flac", ".m4a": "audio/x-m4a", + ".mov": "video/quicktime", + ".mp4": "video/mp4", + ".avi": "video/x-msvideo", + ".mkv": "video/x-matroska", + ".webm": "video/webm", + ".png": "image/png", + ".gif": "image/gif", + ".bmp": "image/bmp", + ".svg": "image/svg+xml", + ".jpg": "image/jpeg", + ".jpeg": "image/jpeg", + ".mp3": "audio/mpeg", + ".wav": "audio/wav", + ".ogg": "audio/ogg", + ".flac": "audio/flac", + ".m4a": "audio/x-m4a", } for ext, mime in _MIME_MAP.items(): if lower_filename.endswith(ext): @@ -98,10 +108,137 @@ def _infer_mime_type_from_storage_key(storage_key: str) -> str: return "video/mp4" # default +# 兜底去重:无 file_hash / client_upload_id 且大小已知时,同库同名同大小近期活动记录视为重复 +FALLBACK_DEDUP_WINDOW_MINUTES = 30 + + +def _find_duplicate_asset( + asset_repository: Any, + *, + library_id: str, + file_hash: str, + client_upload_id: str, + filename: str, + file_size: int = 0, +) -> Any: + """complete/上传幂等去重,按优先级查找已存在的素材。 + + 1. client_upload_id(客户端幂等 token,同一次上传的重试保持一致) + 2. file_hash(内容哈希,不同上传只要内容相同即去重) + 3. 兜底(严格模式,宁可漏判不可误杀):file_hash 与 client_upload_id + 均缺失、且 file_size > 0 时,同库 + 同文件名 + **同大小** 且 30 分钟内 + 仍处 uploading/processing 的记录才判重。 + - file_hash 非空时跳过兜底(hash 已代表内容;同名但内容全新的视频 + 如 iPhone 的 IMG_xxxx.MOV 绝不能被同名占位误杀) + - file_size=0(未知)时不允许仅凭同名 + processing 判重,直接放行 + + 全部为鸭子类型调用:旧仓储无对应方法时静默跳过,不破坏既有实现。 + """ + if client_upload_id: + find = getattr(asset_repository, "find_by_library_and_client_upload_id", None) + if callable(find): + existing = find(library_id=library_id, client_upload_id=client_upload_id) + if existing is not None: + logger.info( + "素材幂等命中(client_upload_id): library=%s token=%s asset=%s", + library_id, + client_upload_id, + getattr(existing, "id", "?"), + ) + return existing + if file_hash: + existing = asset_repository.find_by_library_and_file_hash( + library_id=library_id, + file_hash=file_hash, + ) + if existing is not None: + logger.info( + "素材去重命中(file_hash): library=%s hash=%s asset=%s", + library_id, + file_hash, + existing.id, + ) + return existing + # 同名兜底去重(最后防线,严格模式): + # - 仅当 file_hash / client_upload_id 均缺失时启用(hash 能代表内容时不靠同名猜) + # - file_size 必须 > 0 且与记录大小严格一致;大小未知(0)直接放行 + # - 只命中近期 UPLOADING/PROCESSING 活动记录(READY 历史素材不拦) + if filename and not file_hash and not client_upload_id and file_size and file_size > 0: + find_recent = getattr(asset_repository, "find_recent_active_by_library_and_name", None) + if callable(find_recent): + existing = find_recent( + library_id=library_id, + name=filename, + within_minutes=FALLBACK_DEDUP_WINDOW_MINUTES, + file_size=file_size, + ) + if existing is not None: + logger.info( + "素材幂等兜底命中(近期同名同大小活动记录): library=%s name=%s asset=%s status=%s size=%s", + library_id, + filename, + getattr(existing, "id", "?"), + getattr(existing, "status", None), + file_size, + ) + return existing + elif filename and not file_hash and not client_upload_id and not file_size: + logger.debug( + "同名兜底去重跳过(file_size 未知,宁可放行不可误杀): library=%s name=%s", + library_id, + filename, + ) + return None + + def _create_pending_asset( - asset_repository, project_id, library_id, storage_key, filename, mime_type, user_id, file_hash="" + asset_repository, + project_id, + library_id, + storage_key, + filename, + mime_type, + user_id, + file_hash="", + client_upload_id="", + file_size: int = 0, ): - """立即创建一条 PROCESSING 状态的 Asset 记录,使前端能马上看到新素材。""" + """立即创建或复用一条 PROCESSING 状态的 Asset 记录。 + + find-or-create:prepare 阶段已按 file_hash/client_upload_id 预建的占位记录 + 会被 find_by_library_and_file_hash/find_by_library_and_client_upload_id 命中, + 直接复用并补齐字段(避免 pre-create + complete 重复建两条)。 + """ + # 1. 按 client_upload_id / file_hash 查找现有记录 + existing = None + if client_upload_id: + find_by_cuid = getattr(asset_repository, "find_by_library_and_client_upload_id", None) + if callable(find_by_cuid): + existing = find_by_cuid(library_id=library_id, client_upload_id=client_upload_id) + if existing is None and file_hash: + existing = asset_repository.find_by_library_and_file_hash(library_id=library_id, file_hash=file_hash) + if existing is not None: + # 补齐字段(幂等:避免重复建记录,前端已拿到 asset_id) + changed = False + if file_hash and not existing.file_hash: + existing.file_hash = file_hash + changed = True + if client_upload_id and not existing.client_upload_id: + existing.client_upload_id = client_upload_id + changed = True + if file_size and not existing.file_size: + existing.file_size = file_size + changed = True + if existing.status not in (AssetStatus.PROCESSING, AssetStatus.UPLOADING): + existing.status = AssetStatus.PROCESSING + changed = True + if changed: + try: + asset_repository.update(existing) + except Exception: # noqa: BLE001 — 字段补齐失败不阻塞主流程 + pass + return existing + asset = Asset.create( project_id=project_id, library_id=library_id, @@ -111,16 +248,30 @@ def _create_pending_asset( status=AssetStatus.PROCESSING, uploaded_by_user_id=user_id, file_hash=file_hash, + client_upload_id=client_upload_id, + file_size=file_size, ) return asset_repository.create(asset) +def _persist_celery_task_id(repo: Any, job: Any, celery_task_id: str) -> None: + """记录 celery 消息 ID 到任务行,供孤儿清理时 revoke/清除队列消息(#1714)。""" + if not celery_task_id: + return + try: + job.celery_task_id = celery_task_id + repo.update(job) + except Exception: # noqa: BLE001 — 记录失败不影响主流程(执行前状态守卫兜底) + pass + + def _submit_ingest_job( project_id: str, library_id: str, storage_key: str, ingest_job_repository: Any, file_hash: str = "", + asset_id: str = "", ) -> Any: use_case = SubmitIngestJobUseCase(ingest_job_repository) job = use_case.execute( @@ -129,9 +280,11 @@ def _submit_ingest_job( library_id=library_id, storage_key=storage_key, file_hash=file_hash, + asset_id=asset_id, ) ) - celery_app.send_task("worker.ingest_asset", args=[job.id]) + celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id]) + _persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", "")) return job @@ -141,9 +294,15 @@ async def prepare_direct_upload( authenticated_user: AuthenticatedUser = Depends(get_current_user), project_repository: Any = Depends(get_project_repository), asset_library_repository: Any = Depends(get_asset_library_repository), + asset_repository: Any = Depends(get_asset_repository), storage_service: OSSStorageService = Depends(get_storage_service), ) -> DirectUploadPrepareResponse: - """创建浏览器直传 OSS 的短期表单签名。""" + """创建浏览器直传 OSS 的短期表单签名,并在签名前按 file_hash/client_upload_id 去重。 + + 命中去重:直接返回 duplicated=True + skip_transfer=True(前端跳过 OSS 直传), + 未命中:正常签名 OSS 并立即预建一条 PROCESSING 状态的 asset 记录占住 + file_hash 闸门,响应带 asset_id 供前端/后续 complete 关联。 + """ settings = get_settings() max_size_bytes = settings.OSS_DIRECT_UPLOAD_MAX_MB * 1024 * 1024 if request.file_size > max_size_bytes: @@ -162,8 +321,39 @@ async def prepare_direct_upload( asset_library_repository, ) - file_id = uuid4().hex[:8] safe_filename = request.filename.replace("/", "_").replace("\\", "_") + + # ── prepare 阶段去重:OSS 签名之前先查已存在素材 ── + if request.file_hash or request.client_upload_id: + existing = _find_duplicate_asset( + asset_repository, + library_id=request.library_id, + file_hash=request.file_hash, + client_upload_id=request.client_upload_id, + filename=request.filename, + file_size=request.file_size, + ) + if existing is not None: + logger.info( + "prepare 命中去重: library=%s hash=%s cuid=%s existing_asset=%s", + request.library_id, + request.file_hash, + request.client_upload_id, + existing.id, + ) + return DirectUploadPrepareResponse( + upload_url="", + method="", + storage_key=existing.storage_key, + expires_at="", + fields={}, + max_size_bytes=0, + duplicated=True, + skip_transfer=True, + asset_id=existing.id, + ) + + file_id = uuid4().hex[:8] storage_key = f"uploads/{file_id}/{safe_filename}" try: payload = storage_service.create_direct_upload_post( @@ -182,6 +372,27 @@ async def prepare_direct_upload( detail=f"Failed to prepare upload: {type(error).__name__}", ) from error + # ── 预建 asset 占位:占住 file_hash/client_upload_id 闸门,避免并发重复上传 ── + pending_asset_id = "" + if request.file_hash or request.client_upload_id: + try: + pending = _create_pending_asset( + asset_repository=asset_repository, + project_id=request.project_id, + library_id=request.library_id, + storage_key=storage_key, + filename=safe_filename, + mime_type=validated_content_type, + user_id=authenticated_user.user.id, + file_hash=request.file_hash, + client_upload_id=request.client_upload_id, + file_size=request.file_size, + ) + pending_asset_id = pending.id + except Exception as error: + # 预建失败不阻塞签名:complete 仍可按 OSS 文件 + hash 兜底去重 + logger.warning("预建 asset 占位失败,降级走 old flow: %s", error) + return DirectUploadPrepareResponse( upload_url=str(payload["url"]), method=str(payload["method"]), @@ -189,6 +400,9 @@ async def prepare_direct_upload( expires_at=str(payload["expires_at"]), fields={str(key): str(value) for key, value in dict(payload["fields"]).items()}, max_size_bytes=max_size_bytes, + duplicated=False, + skip_transfer=False, + asset_id=pending_asset_id, ) @@ -202,7 +416,7 @@ async def complete_direct_upload( asset_repository: Any = Depends(get_asset_repository), storage_service: OSSStorageService = Depends(get_storage_service), ) -> DirectUploadCompleteResponse: - """确认浏览器直传完成并创建导入任务。""" + """确认浏览器直传完成并创建导入任务(幂等:重复 complete 返回同一素材)。""" require_project_and_library( request.project_id, request.library_id, @@ -212,6 +426,29 @@ async def complete_direct_upload( normalized_key = storage_service._normalize_storage_key(request.storage_key) if not normalized_key.startswith("uploads/"): raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid upload key") + + filename = normalized_key.rsplit("/", 1)[-1] + + # ── 幂等去重(放在 OSS 检查之前):complete 超时后前端重试时, + # 第一次 complete 可能已建好占位记录,此时即使 OSS 检查失败也必须返回 + # 已存在记录,绝不能再建第二条。─ + existing = _find_duplicate_asset( + asset_repository, + library_id=request.library_id, + file_hash=request.file_hash, + client_upload_id=request.client_upload_id, + filename=filename, + file_size=request.file_size, + ) + if existing is not None: + return DirectUploadCompleteResponse( + storage_key=existing.storage_key, + ingest_job_id="", + duplicated=True, + asset_id=existing.id, + url=storage_service.get_url(existing.storage_key), + ) + try: file_exists = storage_service.file_exists(normalized_key) except Exception as error: @@ -223,29 +460,7 @@ async def complete_direct_upload( if not file_exists: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Uploaded file not found") - # ── 素材去重检测:同素材库 + 同 file_hash 视为重复 ── - if request.file_hash: - existing = asset_repository.find_by_library_and_file_hash( - library_id=request.library_id, - file_hash=request.file_hash, - ) - if existing is not None: - logger.info( - "素材去重命中: library=%s hash=%s existing_asset=%s", - request.library_id, - request.file_hash, - existing.id, - ) - return DirectUploadCompleteResponse( - storage_key=normalized_key, - ingest_job_id="", - duplicated=True, - asset_id=existing.id, - url=storage_service.get_url(normalized_key), - ) - # 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材 - filename = normalized_key.rsplit("/", 1)[-1] mime_type = _infer_mime_type_from_storage_key(normalized_key) pending_asset = _create_pending_asset( asset_repository=asset_repository, @@ -256,6 +471,8 @@ async def complete_direct_upload( mime_type=mime_type, user_id=authenticated_user.user.id, file_hash=request.file_hash, + client_upload_id=request.client_upload_id, + file_size=request.file_size, ) job = _submit_ingest_job( @@ -264,6 +481,7 @@ async def complete_direct_upload( storage_key=normalized_key, ingest_job_repository=ingest_job_repository, file_hash=request.file_hash, + asset_id=pending_asset.id, ) return DirectUploadCompleteResponse( storage_key=normalized_key, @@ -283,7 +501,8 @@ async def upload_asset( project_id: str = Form(..., min_length=1, description="项目 ID"), library_id: str = Form(..., min_length=1, description="素材库 ID"), file: UploadFile = File(..., description="要上传的文件(视频、音频、图片等)"), - file_hash: str = Form(default="", description="文件 MD5 哈希,用于去重检测"), + file_hash: str = Form(default="", description="文件哈希,用于去重检测"), + client_upload_id: str = Form(default="", description="客户端幂等 token(同一次上传的重试保持一致)"), authenticated_user: AuthenticatedUser = Depends(get_current_user), ingest_job_repository: Any = Depends(get_ingest_job_repository), project_repository: Any = Depends(get_project_repository), @@ -294,32 +513,31 @@ async def upload_asset( """上传素材文件并触发导入流水线。""" require_project_and_library(project_id, library_id, project_repository, asset_library_repository) - # ── 素材去重检测:上传前检查同素材库 + 同 file_hash ── - if file_hash: - existing = asset_repository.find_by_library_and_file_hash( - library_id=library_id, - file_hash=file_hash, - ) - if existing is not None: - logger.info( - "素材去重命中(multipart): library=%s hash=%s existing_asset=%s", - library_id, - file_hash, - existing.id, - ) - return UploadAssetResponse( - storage_key=existing.storage_key, - ingest_job_id="", - url="", - duplicated=True, - asset_id=existing.id, - ) - - # P2-5: 服务端验证 MIME 类型 + # P2-5: 服务端验证 MIME 类型(先验证,再幂等去重,避免非法类型绕过) validated_content_type = _validate_mime_type(file.content_type) - file_id = uuid4().hex[:8] safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown" + + # ── 幂等去重:client_upload_id → file_hash → 近期活动同名记录兜底 ── + # 放在 OSS 上传之前:重复提交直接返回,不占 OSS 流量、不建新记录。 + existing = _find_duplicate_asset( + asset_repository, + library_id=library_id, + file_hash=file_hash, + client_upload_id=client_upload_id, + filename=safe_filename, + file_size=0, + ) + if existing is not None: + return UploadAssetResponse( + storage_key=existing.storage_key, + ingest_job_id="", + url="", + duplicated=True, + asset_id=existing.id, + ) + + file_id = uuid4().hex[:8] storage_key = f"uploads/{file_id}/{safe_filename}" try: @@ -348,6 +566,7 @@ async def upload_asset( mime_type=validated_content_type, user_id=authenticated_user.user.id, file_hash=file_hash, + client_upload_id=client_upload_id, ) job = _submit_ingest_job( @@ -356,6 +575,7 @@ async def upload_asset( storage_key=storage_key, ingest_job_repository=ingest_job_repository, file_hash=file_hash, + asset_id=pending_asset.id, ) return UploadAssetResponse( diff --git a/apps/api/app/api/routes/videos.py b/apps/api/app/api/routes/videos.py index 77be5a3ff..20e623f43 100644 --- a/apps/api/app/api/routes/videos.py +++ b/apps/api/app/api/routes/videos.py @@ -252,6 +252,10 @@ class RecomputeDedupRequest(BaseModel): None, description="指定视频 ID 列表。为空则对当前用户所有缺少查重数据的视频重新计算。", ) + force: bool = Field( + False, + description="强制重算:即使视频已有查重数据也重新入队(#1702 查重算法升级后用于存量视频重算)。", + ) class RecomputeDedupResponse(BaseModel): @@ -291,15 +295,15 @@ def recompute_dedup( skipped = 0 for video in target_videos: - # 已有完整查重数据的跳过 - if video.duplicate_rate is not None and video.video_fingerprint: + # 已有完整查重数据的跳过(force=True 时强制重算,#1702 算法升级后存量视频需要重算指纹/分片) + if not request.force and video.duplicate_rate is not None and video.video_fingerprint: skipped += 1 continue # 触发异步查重任务 celery_app.send_task("worker.check_duplicate", args=[video.id]) enqueued += 1 - logger.info("Enqueued re-dedup for video %s (user=%s)", video.id, user_id) + logger.info("Enqueued re-dedup for video %s (user=%s, force=%s)", video.id, user_id, request.force) return RecomputeDedupResponse( enqueued=enqueued, diff --git a/apps/api/app/core/celery_app.py b/apps/api/app/core/celery_app.py index 52b515335..3d7d7bb2a 100644 --- a/apps/api/app/core/celery_app.py +++ b/apps/api/app/core/celery_app.py @@ -5,3 +5,11 @@ settings = get_settings() celery_app = Celery("xiaoxia-saas-api") celery_app.conf.broker_url = settings.CELERY_BROKER_URL celery_app.conf.result_backend = settings.CELERY_RESULT_BACKEND + +# #1714 队列隔离:视频生成走 generation 队列,素材转码走 transcode 队列 +try: + from packages.shared.celery_queues import apply_queue_settings + + apply_queue_settings(celery_app) +except Exception: # noqa: BLE001 — 队列配置失败不阻断 API 启动 + pass diff --git a/apps/api/app/core/task_enqueue.py b/apps/api/app/core/task_enqueue.py index 31545f177..c3b11329c 100755 --- a/apps/api/app/core/task_enqueue.py +++ b/apps/api/app/core/task_enqueue.py @@ -8,27 +8,145 @@ logger = logging.getLogger(__name__) # ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ── USER_PENDING_LIMIT = 3 # 单用户 pending 上限 GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限 +WORKER_CONCURRENCY = 4 # worker 渲染并发数(infra/docker/compose.yml WORKER_CONCURRENCY 默认值) + +# 限流错误码:前端据此区分"排队等待"与"创建失败" +ERROR_CODE_USER_QUEUE_FULL = "USER_QUEUE_FULL" # 429:用户自己的任务排队中 +ERROR_CODE_SYSTEM_QUEUE_FULL = "SYSTEM_QUEUE_FULL" # 503:系统整体繁忙 class UserPendingLimitExceeded(Exception): """用户 pending 任务数超限,返回 429。""" - def __init__(self, user_id: str, pending_count: int, limit: int): + def __init__( + self, + user_id: str, + pending_count: int, + limit: int, + *, + running_count: int = 0, + requested_count: int = 1, + queue_ahead: int = 0, + estimated_wait_seconds: int = 0, + ): self.user_id = user_id self.pending_count = pending_count self.limit = limit + # 排队上下文(用于 429 结构化提示,前端展示"排队中"而非"创建失败") + self.running_count = running_count + self.requested_count = requested_count + self.queue_ahead = queue_ahead + self.estimated_wait_seconds = estimated_wait_seconds super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}") class GlobalQueueFull(Exception): """全局限流,返回 503。""" - def __init__(self, pending_count: int, limit: int): + def __init__( + self, + pending_count: int, + limit: int, + *, + running_count: int = 0, + queue_ahead: int = 0, + estimated_wait_seconds: int = 0, + ): self.pending_count = pending_count self.limit = limit + self.running_count = running_count + self.queue_ahead = queue_ahead + self.estimated_wait_seconds = estimated_wait_seconds super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}") +def _estimate_wait_seconds(queue_ahead: int, generation_task_repository: Any) -> int: + """根据排队任务数 + worker 并发数 + 历史平均任务耗时估算等待秒数。 + + 估算公式:ceil(排队任务数 / 并发数) × 平均单任务耗时。 + 拿不到历史数据时仓储层返回默认 120 秒。 + """ + import math + + if queue_ahead <= 0: + return 0 + try: + estimator = getattr(generation_task_repository, "estimate_avg_duration_seconds", None) + avg_seconds = estimator() if estimator is not None else 120.0 + except Exception: + avg_seconds = 120.0 + return int(math.ceil(queue_ahead / WORKER_CONCURRENCY) * avg_seconds) + + +def build_rate_limit_detail( + exc: Exception, + generation_task_repository: Any, + *, + scope: str = "user", +) -> dict: + """构造结构化限流响应体(HTTPException 的 detail)。 + + 前端按 detail.code 判断场景: + - USER_QUEUE_FULL (429):用户自己的任务在排队,应提示"等待/继续排队",不是创建失败 + - SYSTEM_QUEUE_FULL (503):系统繁忙,稍后重试 + + detail 字段: + - code: 错误码 + - message: 可读中文提示(可直接展示) + - queued_count: 当前排队(pending)任务数 + - running_count: 当前渲染中(running)任务数 + - queue_ahead: 前方排队任务数(预计等待批次依据) + - estimated_wait_seconds: 预计等待秒数 + - limit: 对应限流上限 + """ + if scope == "user" and isinstance(exc, UserPendingLimitExceeded): + running = exc.running_count + if not running: + try: + counter = getattr(generation_task_repository, "count_running_by_user", None) + running = counter(exc.user_id) if counter is not None else 0 + except Exception: + running = 0 + queue_ahead = exc.queue_ahead or max(exc.pending_count, 0) + wait = exc.estimated_wait_seconds or _estimate_wait_seconds(queue_ahead, generation_task_repository) + wait_minutes = max(1, round(wait / 60)) + message = ( + f"您有 {exc.pending_count} 个任务正在排队、{running} 个正在渲染," + f"同一时间最多提交 {exc.limit} 个任务。请等待约 {wait_minutes} 分钟后再提交" + ) + return { + "code": ERROR_CODE_USER_QUEUE_FULL, + "message": message, + "queued_count": exc.pending_count, + "running_count": running, + "queue_ahead": queue_ahead, + "estimated_wait_seconds": wait, + "limit": exc.limit, + } + + # 全局繁忙 + pending = getattr(exc, "pending_count", 0) + running = getattr(exc, "running_count", 0) + if not running: + try: + counter = getattr(generation_task_repository, "count_running_total", None) + running = counter() if counter is not None else 0 + except Exception: + running = 0 + queue_ahead = getattr(exc, "queue_ahead", 0) or pending + wait = getattr(exc, "estimated_wait_seconds", 0) or _estimate_wait_seconds(queue_ahead, generation_task_repository) + wait_minutes = max(1, round(wait / 60)) + return { + "code": ERROR_CODE_SYSTEM_QUEUE_FULL, + "message": f"系统繁忙:当前 {pending} 个任务排队中、{running} 个渲染中,预计等待约 {wait_minutes} 分钟,请稍后再试", + "queued_count": pending, + "running_count": running, + "queue_ahead": queue_ahead, + "estimated_wait_seconds": wait, + "limit": getattr(exc, "limit", GLOBAL_PENDING_LIMIT), + } + + def check_queue_limits( user_id: str, generation_task_repository: Any, @@ -161,7 +279,17 @@ def safe_enqueue_generation_task( # ── 发送 Celery 任务 ── try: - celery_app.send_task("worker.generate_video", args=[task.id]) + celery_result = celery_app.send_task("worker.generate_video", args=[task.id]) + # 记录 celery 消息 ID:孤儿清理/超时作废时据此 revoke + 清除队列消息(#1714) + celery_task_id = getattr(celery_result, "id", "") + if celery_task_id: + try: + task.celery_task_id = celery_task_id + generation_task_repository.update(task) + except Exception as persist_err: # noqa: BLE001 + logger.warning( + "%s 持久化 celery_task_id 失败(不影响主流程): task_id=%s err=%s", log_prefix, task.id, persist_err + ) except Exception as e: logger.error( "%s 入队失败,标记为失败: task_id=%s error=%s", diff --git a/apps/api/app/schemas/generation_task.py b/apps/api/app/schemas/generation_task.py index bead763e3..15f55d895 100755 --- a/apps/api/app/schemas/generation_task.py +++ b/apps/api/app/schemas/generation_task.py @@ -25,6 +25,12 @@ class CreateGenerationTaskRequest(BaseModel): asset_library_id: str = "" strategy_id: str = "" voice_library_id: str = "" + # ── 多变体独立配音(批量生成)── + # 长度 1 = 所有变体共用;长度 = count = 每个变体独立配音;空数组 = 回退 voice_library_id + voice_library_ids: list[str] = Field( + default_factory=list, + description="各变体独立配音素材库ID数组:长度1=共用,长度=count=独立。为空时回退 voice_library_id", + ) created_by_user_id: str = "" # ── 模板模式新增字段 ── template_id: str = "" @@ -75,6 +81,27 @@ class CreateGenerationTaskRequest(BaseModel): output_width: int = Field(default=1280, description="输出视频宽度") output_height: int = Field(default=720, description="输出视频高度") cover_url: str = Field(default="", description="封面图片 URL") + # ── 多变体独立封面(批量生成)── + # 长度 1 = 所有变体共用;长度 = count = 每个变体独立封面;空数组 = 回退 cover_url + cover_urls: list[str] = Field( + default_factory=list, + description="各变体独立封面URL数组:长度1=共用,长度=count=独立。为空时回退 cover_url", + ) + # ── 多变体独立标题文字(批量生成)── + # 长度 1 = 所有变体共用;长度 = count = 每个变体独立标题文字;空数组 = 使用 title_config.text + titles: list[str] = Field( + default_factory=list, + description="各变体独立标题文字数组:长度1=共用,长度=count=独立。为空时使用 title_config.text", + ) + + @model_validator(mode="after") + def _check_variant_arrays(self) -> "CreateGenerationTaskRequest": + """变体数组字段长度校验:空数组(回退单值)、长度 1(共用)、或长度 = count(独立)。""" + for name in ("voice_library_ids", "cover_urls", "titles"): + arr = getattr(self, name) + if arr and len(arr) != 1 and len(arr) != self.count: + raise ValueError(f"{name} 长度必须为 1(共用)或 {self.count}(与 count 一致),当前为 {len(arr)}") + return self @model_validator(mode="after") def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest": @@ -185,8 +212,33 @@ class CreatePreviewGenerationTaskRequest(BaseModel): ) title_config: dict = Field( default_factory=dict, - description="标题配置(可选),渲染时烧录到预览视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow", + description="标题配置(可选),渲染时烧录到预览视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow。N个变体时样式全局共用", ) + # ── 多变体独立配置(preview_count > 1)── + # 长度 1 = 所有变体共用;长度 = preview_count = 每个变体独立;空数组 = 回退单值字段 + titles: list[str] = Field( + default_factory=list, + description="各变体独立标题文字数组:长度1=共用,长度=preview_count=独立。为空时使用 title_config.text", + ) + voice_library_ids: list[str] = Field( + default_factory=list, + description="各变体独立配音素材库ID数组:长度1=共用,长度=preview_count=独立。为空时回退 voice_library_id", + ) + cover_urls: list[str] = Field( + default_factory=list, + description="各变体独立封面URL数组:长度1=共用,长度=preview_count=独立(预览阶段通常为空)", + ) + + @model_validator(mode="after") + def _check_variant_arrays(self) -> "CreatePreviewGenerationTaskRequest": + """变体数组字段长度校验:空数组(回退单值)、长度 1(共用)、或长度 = preview_count(独立)。""" + for name in ("titles", "voice_library_ids", "cover_urls"): + arr = getattr(self, name) + if arr and len(arr) != 1 and len(arr) != self.preview_count: + raise ValueError( + f"{name} 长度必须为 1(共用)或 {self.preview_count}(与 preview_count 一致),当前为 {len(arr)}" + ) + return self @model_validator(mode="after") def _check_template_id(self) -> "CreatePreviewGenerationTaskRequest": @@ -202,7 +254,7 @@ class CreatePreviewGenerationTaskRequest(BaseModel): class PreviewGenerationTaskResponse(BaseModel): - """预览生成任务响应。 + """单个预览变体任务响应。 包含任务状态、进度、分辨率、生成结果 URL 等关键字段。 """ @@ -211,6 +263,7 @@ class PreviewGenerationTaskResponse(BaseModel): status: str progress: float is_preview: bool = True + variant_index: int = 0 resolution: str = "" video_url: str = "" duration: float = 0.0 @@ -219,7 +272,21 @@ class PreviewGenerationTaskResponse(BaseModel): transition_count: int = 0 material_usage: dict = Field(default_factory=dict) error_message: str = "" + title_text: str = "" + voice_library_id: str = "" created_at: datetime | None = None started_at: datetime | None = None finished_at: datetime | None = None generate_duration: float = 0.0 + + +class BatchPreviewGenerationTaskResponse(BaseModel): + """批量预览任务响应:preview_count=N 时返回 N 个独立变体任务。 + + - items: 变体任务数组,按 variant_index 顺序排列,每个含独立 task_id/状态/预览视频URL + - total: 变体总数(= preview_count) + - 前端按 items[i].task_id 分别轮询 GET /preview/{task_id} 获取进度与结果 + """ + + items: list[PreviewGenerationTaskResponse] + total: int diff --git a/apps/api/app/schemas/upload.py b/apps/api/app/schemas/upload.py index c6d798288..54b6cba6f 100644 --- a/apps/api/app/schemas/upload.py +++ b/apps/api/app/schemas/upload.py @@ -16,6 +16,7 @@ class DirectUploadPrepareRequest(BaseModel): content_type: str = Field(default="application/octet-stream", min_length=1, max_length=100) file_size: int = Field(..., gt=0) file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测") + client_upload_id: str = Field(default="", max_length=64, description="客户端幂等 token(同一次上传的重试保持一致)") class DirectUploadPrepareResponse(BaseModel): @@ -25,20 +26,25 @@ class DirectUploadPrepareResponse(BaseModel): expires_at: str fields: dict[str, str] max_size_bytes: int + duplicated: bool = False + skip_transfer: bool = False + asset_id: str = "" class DirectUploadCompleteRequest(BaseModel): project_id: str = Field(..., min_length=1) library_id: str = Field(..., min_length=1) storage_key: str = Field(..., min_length=1, max_length=255) - file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测") + file_hash: str = Field(default="", max_length=64, description="文件哈希,用于去重检测") + client_upload_id: str = Field(default="", max_length=64, description="客户端幂等 token(同一次上传的重试保持一致)") + file_size: int = Field(default=0, ge=0, description="文件大小(字节),用于无 hash 时的兜底去重") class DirectUploadCompleteResponse(BaseModel): storage_key: str ingest_job_id: str - duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)") - asset_id: str = Field(default="", description="重复素材的 asset_id(duplicated=true 时返回)") + duplicated: bool = Field(default=False, description="是否为重复素材/重复 complete(命中幂等去重)") + asset_id: str = Field(default="", description="素材 asset_id(重复 complete 时返回已存在记录)") url: str = Field(default="", description="Public URL of uploaded file") @@ -46,5 +52,5 @@ class UploadAssetResponse(BaseModel): storage_key: str ingest_job_id: str url: str = Field(..., description="Public URL of uploaded file") - duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)") - asset_id: str = Field(default="", description="重复素材的 asset_id(duplicated=true 时返回)") + duplicated: bool = Field(default=False, description="是否为重复素材/重复提交(命中幂等去重)") + asset_id: str = Field(default="", description="素材 asset_id(重复提交时返回已存在记录)") diff --git a/apps/web/e2e/core-generation.spec.ts b/apps/web/e2e/core-generation.spec.ts index 1e215d69d..c282269f3 100755 --- a/apps/web/e2e/core-generation.spec.ts +++ b/apps/web/e2e/core-generation.spec.ts @@ -185,6 +185,13 @@ test.describe("Core generation flow", () => { await expect(page.locator(".xx-choice-item.selected")).toBeVisible() await page.getByRole("button", { name: "下一步" }).click() + // Step1 下一步弹出数量选择弹窗(Issue #1677 固定6步:模板→素材→配音→标题→确认生成→封面) + // 单视频流程:默认 1 个,点击「生成 1 个视频」进入步骤2 + await expect(page.getByRole("heading", { name: "要生成几个视频?" })).toBeVisible({ + timeout: 10_000, + }) + await page.getByRole("button", { name: "生成 1 个视频" }).click() + // Step 2: select material (card grid UI) await expect(page.getByRole("heading", { name: /选择素材/ })).toBeVisible() const librarySelect = page.locator("select").first() @@ -255,32 +262,19 @@ test.describe("Core generation flow", () => { expect(genData.items.length).toBeGreaterThan(0) expect(genData.items[0].id).toBeTruthy() - // Step 5: 确认生成页 — 任务创建成功后自动跳转,展示渲染进度 - await expect(page.getByRole("heading", { name: /确认生成/ })).toBeVisible({ - timeout: 15_000, + // 单视频(N=1):点击「确认生成视频」后跳 Step 5「确认生成」,展示实时渲染进度 + await expect(page.getByRole("heading", { name: "🎬 确认生成" })).toBeVisible({ + timeout: 30_000, }) - // Step 5 → Step 6:等待渲染终态 - // - 完成:页面出现「视频生成完成」,步骤5「下一步」按钮解锁,点击进入封面 - // - 失败:出现「生成失败」,停在确认生成页也算向导流程走通 - // - 超时未终态(测试环境 worker 可能不处理任务):进度仍在轮询,同样算走通 - const renderSucceeded = await page - .getByText("视频生成完成", { exact: false }) - .waitFor({ timeout: 180_000 }) - .then(() => true) - .catch(() => false) - if (renderSucceeded) { - // 渲染完成:手动点「下一步」进入封面步骤(渲染完不自动跳转) - await page.getByRole("button", { name: "下一步" }).click() - // Step 6: 封面(最后一步,无主按钮),仅验证页面渲染 - await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({ - timeout: 15_000, - }) - } else { - // 失败或超时:仍在确认生成页(进度展示或失败提示),向导流程已完整走通 - await expect(page.getByRole("heading", { name: /确认生成/ })).toBeVisible() - console.log("[E2E] 渲染任务失败或未在 180s 内完成,冒烟测试仍通过(已达确认生成页)") - } + // 等待渲染完成:进度卡变为「视频生成完成」(最长等待 3 分钟) + await expect(page.getByText("视频生成完成")).toBeVisible({ timeout: 180_000 }) + + // 全部完成后「下一步:选择封面」解锁,点击进入 Step 6 + await page.getByRole("button", { name: /下一步:选择封面/ }).click() + await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({ + timeout: 30_000, + }) } else { console.log(`[E2E] Generate API returned ${genResp.status()}, wizard flow test still passes`) // 创建失败时停留在标题页并展示错误提示 diff --git a/apps/web/public/vite.svg b/apps/web/public/vite.svg new file mode 100644 index 000000000..7f7809d44 --- /dev/null +++ b/apps/web/public/vite.svg @@ -0,0 +1,4 @@ + + + 🦐 + diff --git a/apps/web/src/api/assets/types.ts b/apps/web/src/api/assets/types.ts index 0efe111ab..77ec261fc 100644 --- a/apps/web/src/api/assets/types.ts +++ b/apps/web/src/api/assets/types.ts @@ -139,6 +139,17 @@ export interface DirectUploadPrepareResult { * 旧后端不返回该字段,前端降级为无预建卡片的原有行为。 */ asset_id?: string + /** + * 后端 file_hash 命中素材库已有相同文件时为 true,前端应跳过 transfer + complete 阶段 + * 直接按「去重命中」处理(不调 transfer、不调 complete、立即刷新素材列表)。 + * 旧后端不返回该字段,前端降级为走老流程。 + */ + duplicated?: boolean + /** + * 与 duplicated 语义一致:true 表示跳过传输,前端据此短路。 + * 两个字段是同一语义的别名(后端可能只返回其一),前端任意为 true 即视为命中去重。 + */ + skip_transfer?: boolean } /** 直传完成确认返回 */ diff --git a/apps/web/src/api/assets/upload.ts b/apps/web/src/api/assets/upload.ts index e396b5d77..490b08df7 100644 --- a/apps/web/src/api/assets/upload.ts +++ b/apps/web/src/api/assets/upload.ts @@ -4,6 +4,7 @@ import apiClient from "../client" import { getOrCreateDefaultProject } from "../projects" import type { DirectUploadPrepareResult, DirectUploadCompleteResult } from "./types" +import { computeFileHash, makeClientUploadId } from "./uploadDedup" /** 预签名直传准备 */ export const prepareDirectUpload = async (data: { @@ -12,8 +13,13 @@ export const prepareDirectUpload = async (data: { filename: string content_type: string file_size: number + /** 前端算好的文件内容哈希(SHA-256 hex),打开后端 file_hash 去重闸门 */ + file_hash?: string + /** 前端生成的上传幂等 token,同一次逻辑上传(含重试)保持不变 */ + client_upload_id?: string }): Promise => { - const response = await apiClient.post("/upload/direct/prepare", data) + // prepare 单独放宽到 30s(全局 axios 实例只有 10s,staging 抖动时易超时) + const response = await apiClient.post("/upload/direct/prepare", data, { timeout: 30_000 }) return response.data } @@ -22,8 +28,16 @@ export const completeDirectUpload = async (data: { project_id: string library_id: string storage_key: string + /** 前端算好的文件内容哈希(与 prepare 一致),后端按 hash 幂等去重 */ + file_hash?: string + /** 前端上传幂等 token(与 prepare 一致),同一次上传重发 complete 不重复建记录 */ + client_upload_id?: string + /** 文件字节数;后端同名兜底去重需用它做大小校验,缺失(=0)时同名记录一律不判重 */ + file_size?: number }): Promise => { - const response = await apiClient.post("/upload/direct/complete", data) + // complete 内含 OSS 存在性检查 + 建库 + 派单,放宽到 60s; + // 超时不代表失败(记录可能已建成),调用方禁止超时后盲目重传整个文件 + const response = await apiClient.post("/upload/direct/complete", data, { timeout: 60_000 }) return response.data } @@ -109,8 +123,20 @@ export interface DirectUploadHandle { export const prepareDirectUploadHandle = async (data: { file: File library_id: string + /** 前端算好的文件内容哈希(SHA-256 hex),prepare/complete 均携带 */ + fileHash?: string + /** 本次逻辑上传的幂等 token,prepare/complete 一致、重试复用 */ + clientUploadId?: string }): Promise => { - const project = await getOrCreateDefaultProject() + // 默认项目初始化失败(项目列表接口异常/自动创建失败)给出独立、明确的提示, + // 不与 prepare 的签名接口错误混在一起 + let project: Awaited> + try { + project = await getOrCreateDefaultProject() + } catch (err) { + const reason = err instanceof Error ? err.message : "网络异常" + throw new Error(`初始化默认项目失败,无法开始上传:${reason}`) + } const prepared = await prepareDirectUpload({ project_id: project.id, @@ -118,6 +144,8 @@ export const prepareDirectUploadHandle = async (data: { filename: data.file.name, content_type: data.file.type || "application/octet-stream", file_size: data.file.size, + file_hash: data.fileHash, + client_upload_id: data.clientUploadId, }) return { @@ -128,6 +156,10 @@ export const prepareDirectUploadHandle = async (data: { project_id: project.id, library_id: data.library_id, storage_key: prepared.storage_key, + file_hash: data.fileHash, + client_upload_id: data.clientUploadId, + // 透传文件字节数:后端同名兜底去重依赖大小校验,缺省会导致同名新视频被误判重复 + file_size: data.file.size, }), } } @@ -137,8 +169,30 @@ export const uploadAssetDirect = async (data: { file: File library_id: string onProgress?: (percent: number) => void + /** 文件内容哈希;未传时自动补算(配音/封面/克隆等非队列链路统一受益) */ + fileHash?: string + /** 幂等 token;未传时自动生成 */ + clientUploadId?: string }): Promise => { - const handle = await prepareDirectUploadHandle({ file: data.file, library_id: data.library_id }) + // 自动补算哈希与幂等 token:确保 file_hash 去重闸门对所有上传链路生效 + const fileHash = data.fileHash ?? (await computeFileHash(data.file)) + const clientUploadId = data.clientUploadId ?? makeClientUploadId() + const handle = await prepareDirectUploadHandle({ + file: data.file, + library_id: data.library_id, + fileHash, + clientUploadId, + }) + // prepare 阶段后端 file_hash 命中素材库已有相同文件:跳过 transfer + complete + if (handle.prepared.skip_transfer || handle.prepared.duplicated) { + return { + storage_key: handle.prepared.storage_key, + ingest_job_id: "", + url: "", + duplicated: true, + asset_id: handle.prepared.asset_id, + } + } await handle.transfer(data.onProgress) return handle.complete() } diff --git a/apps/web/src/api/assets/uploadDedup.ts b/apps/web/src/api/assets/uploadDedup.ts new file mode 100644 index 000000000..f30ad6472 --- /dev/null +++ b/apps/web/src/api/assets/uploadDedup.ts @@ -0,0 +1,143 @@ +/** + * 上传去重 / 幂等工具(Issue #1714) + * + * 背景:同一文件被反复入队、complete 超时后盲目重传,导致后端创建大量重复 + * PROCESSING 素材记录。本模块提供两类纯函数: + * + * 1. 文件指纹: + * - makeFileFingerprint():文件名+大小+lastModified,入队去重用(同步、零开销) + * - computeFileHash():SHA-256 内容哈希(小文件全量、大文件抽样头尾), + * prepare/complete 时发给后端打开 file_hash 去重闸门 + * 2. 队列去重:findDuplicateInQueue() 判断文件是否已在队列中 + * 3. 幂等 token:makeClientUploadId() 生成上传幂等 ID(每次"一次逻辑上传"一个, + * 重试复用同一 ID,重新入队才生成新 ID) + */ + +/** 全量哈希阈值:≤64MB 全量读入计算;超过即走头尾抽样,避免 100~256MB 视频被整文件读进内存卡死页面 */ +export const HASH_FULL_READ_LIMIT = 64 * 1024 * 1024 // 64MB +/** 抽样读取的头尾片段大小(各 16MB) */ +export const HASH_SAMPLE_CHUNK = 16 * 1024 * 1024 + +/** 计算指纹时,文件在队列中已存在的状态(已失败的可以重试,不算重复) */ +export type DedupExcludeStatus = "error" | "done" + +/** + * 文件入队指纹:同库 + 文件名 + 大小 + 修改时间。 + * 同一文件(File 对象由 重选或拖拽重复触发时三个字段均一致)稳定复现; + * 不同文件极小概率碰撞时可由后端 file_hash 内容去重兜底。 + */ +export function makeFileFingerprint(file: Pick): string { + return `${file.name}::${file.size}::${file.lastModified}` +} + +/** + * 在现有队列项中查找同一文件的在途记录。 + * 已失败(error)的项允许重试路径复用、已完成(done)的可跳过; + * 处于 preparing/uploading/ingesting 的在途项一律视为重复,禁止重复入队。 + * + * 返回命中的队列项 id(tempId),未命中返回 null。 + */ +export function findDuplicateInQueue( + queue: T[], + fileKey: string, + excludeStatuses: DedupExcludeStatus[] = [], +): T | null { + const exclude = new Set(excludeStatuses) + return queue.find((it) => it.fileKey === fileKey && !exclude.has(it.status)) ?? null +} + +/** 生成上传幂等 token:一次"逻辑上传"一个,重试复用、重新入队换新 */ +export function makeClientUploadId(): string { + const rand = + typeof crypto !== "undefined" && "randomUUID" in crypto + ? crypto.randomUUID() + : `${Date.now()}-${Math.random().toString(36).slice(2, 10)}-${Math.random() + .toString(36) + .slice(2, 10)}` + return `up_${Date.now().toString(36)}_${rand.replace(/-/g, "").slice(0, 16)}` +} + +/** 读取 Blob/File 片段为 ArrayBuffer:优先 Blob.arrayBuffer(),老环境回退 FileReader */ +function readAsArrayBuffer(blob: Blob): Promise { + if (typeof blob.arrayBuffer === "function") { + return blob.arrayBuffer() + } + return new Promise((resolve, reject) => { + const reader = new FileReader() + reader.onload = () => resolve(reader.result as ArrayBuffer) + reader.onerror = () => reject(reader.error ?? new Error("FileReader read failed")) + reader.readAsArrayBuffer(blob) + }) +} + +/** + * 把 buffer 复制到当前 JS realm 的 Uint8Array 再哈希。 + * jsdom/测试环境中 Blob.arrayBuffer() 可能返回另一 realm 的 ArrayBuffer, + * Node WebCrypto 的 WebIDL instanceof 校验会拒绝跨 realm 参数。 + */ +async function digestSha256(buffer: ArrayBuffer): Promise { + const subtle = + typeof globalThis !== "undefined" && globalThis.crypto ? globalThis.crypto.subtle : null + if (!subtle) throw new Error("crypto.subtle unavailable") + const local = new Uint8Array(buffer.byteLength) + local.set(new Uint8Array(buffer)) + return subtle.digest("SHA-256", local) +} + +function toHex(buffer: ArrayBuffer): string { + const bytes = new Uint8Array(buffer) + let hex = "" + for (let i = 0; i < bytes.length; i += 1) { + hex += bytes[i].toString(16).padStart(2, "0") + } + return hex +} + +/** + * 计算文件内容 SHA-256(hex,64 字符,与后端 file_hash 字段长度一致)。 + * - ≤64MB:全量哈希,内容一致必然一致 + * - >64MB:哈希「头部 16MB + 尾部 16MB + 文件大小」,视频素材体积大、 + * 头部含 moov 元数据、尾部含 mdat 结尾,抽样碰撞概率可忽略, + * 且避免 100~256MB 视频被整文件读进内存导致页面卡死/崩溃 + * + * 运行环境不支持 crypto.subtle(非安全上下文/老浏览器)时返回空字符串, + * 调用方据此降级为不传 hash(后端仍有幂等 token + 同文件名兜底去重)。 + */ +export async function computeFileHash(file: File): Promise { + try { + const subtle = + typeof globalThis !== "undefined" && + globalThis.crypto && + typeof globalThis.crypto.subtle?.digest === "function" + ? globalThis.crypto.subtle + : null + if (!subtle) return "" + + if (file.size <= HASH_FULL_READ_LIMIT) { + const data = await readAsArrayBuffer(file.slice(0, file.size)) + return toHex(await digestSha256(data)) + } + + // 大文件:头 8MB + 尾 8MB + 大小,拼成一段后哈希 + const head = await readAsArrayBuffer(file.slice(0, HASH_SAMPLE_CHUNK)) + const tail = + file.size > HASH_SAMPLE_CHUNK + ? await readAsArrayBuffer(file.slice(Math.max(0, file.size - HASH_SAMPLE_CHUNK), file.size)) + : new ArrayBuffer(0) + const merged = new Uint8Array(head.byteLength + tail.byteLength + 8) + merged.set(new Uint8Array(head), 0) + merged.set(new Uint8Array(tail), head.byteLength) + const sizeView = new DataView(merged.buffer, head.byteLength + tail.byteLength, 8) + // 文件大小以 64 位大端写入(BigInt 最稳;不支持 BigInt64 时手算高低位) + if (typeof sizeView.setBigUint64 === "function") { + sizeView.setBigUint64(0, BigInt(file.size), false) + } else { + sizeView.setUint32(0, Math.floor(file.size / 0x100000000), false) + sizeView.setUint32(4, file.size >>> 0, false) + } + return toHex(await digestSha256(merged.buffer)) + } catch (err) { + console.warn("[uploadDedup] 计算文件哈希失败,降级为不传 file_hash:", err) + return "" + } +} diff --git a/apps/web/src/api/auth/index.ts b/apps/web/src/api/auth/index.ts index 966c92c44..322dbe7c2 100644 --- a/apps/web/src/api/auth/index.ts +++ b/apps/web/src/api/auth/index.ts @@ -12,13 +12,18 @@ export type { UserResponse, WechatAuthUrlResponse, WechatCallbackResponse, + WechatBindUrlResponse, + WechatBindCompleteResponse, + WechatUnbindResponse, + UpdateProfileRequest, + UpdateProfileResponse, SendVerificationCodeRequest, BindContactRequest, BindContactResponse, } from "./types" // 用户工具函数 -export { normalizeUser } from "./user" +export { normalizeUser, updateProfile } from "./user" // 登录/注册/登出/刷新 export { login, refreshAccessToken, register, logout } from "./login" @@ -32,8 +37,14 @@ export { requestPasswordReset, resetPassword } from "./password" // 邮箱验证 export { verifyEmail } from "./email" -// 微信登录 -export { getWechatAuthUrl, wechatCallback } from "./wechat" +// 微信登录 / 绑定 +export { + getWechatAuthUrl, + wechatCallback, + getWechatBindUrl, + bindWechat, + unbindWechat, +} from "./wechat" // 联系方式 export { sendVerificationCode, bindContact } from "./contact" diff --git a/apps/web/src/api/auth/types.ts b/apps/web/src/api/auth/types.ts index 023e0fa32..0ee6e2999 100644 --- a/apps/web/src/api/auth/types.ts +++ b/apps/web/src/api/auth/types.ts @@ -34,6 +34,17 @@ export interface User { is_email_verified: boolean email_verified: boolean created_at?: string + /** 微信是否已绑定 */ + wechat_bound?: boolean + /** 微信昵称(绑定后展示) */ + wechat_nickname?: string + /** 头像 URL(微信头像等) */ + avatar_url?: string + /** 手机号 */ + phone?: string + phone_verified?: boolean + /** 资料是否完善(微信新用户首次登录为 false,需填昵称引导) */ + profile_completed?: boolean } export interface UserResponse { @@ -45,6 +56,12 @@ export interface UserResponse { is_email_verified?: boolean email_verified?: boolean created_at?: string + wechat_bound?: boolean + wechat_nickname?: string + avatar_url?: string + phone?: string + phone_verified?: boolean + profile_completed?: boolean } export interface WechatAuthUrlResponse { @@ -80,3 +97,30 @@ export interface BindContactResponse { success: boolean user: User } + +/** 更新个人资料请求 */ +export interface UpdateProfileRequest { + display_name?: string +} + +/** 更新个人资料响应(返回最新用户信息) */ +export interface UpdateProfileResponse { + user: UserResponse +} + +/** 微信绑定授权链接响应 */ +export interface WechatBindUrlResponse { + auth_url: string + state: string +} + +/** 微信绑定完成响应 */ +export interface WechatBindCompleteResponse { + success: boolean + user: UserResponse +} + +/** 微信解绑响应 */ +export interface WechatUnbindResponse { + success: boolean +} diff --git a/apps/web/src/api/auth/user.ts b/apps/web/src/api/auth/user.ts index 5c345e3b5..395cb2721 100644 --- a/apps/web/src/api/auth/user.ts +++ b/apps/web/src/api/auth/user.ts @@ -1,4 +1,5 @@ -import type { User, UserResponse } from "./types" +import apiClient from "../client" +import type { User, UserResponse, UpdateProfileRequest, UpdateProfileResponse } from "./types" /** * 规范化用户数据,兼容不同后端返回格式 @@ -16,5 +17,19 @@ export const normalizeUser = (data: UserResponse): User => { is_email_verified: emailVerified, email_verified: emailVerified, created_at: data.created_at, + wechat_bound: data.wechat_bound, + wechat_nickname: data.wechat_nickname, + avatar_url: data.avatar_url, + phone: data.phone, + phone_verified: data.phone_verified, + profile_completed: data.profile_completed, } } + +/** + * 更新个人资料(昵称等) + */ +export const updateProfile = async (data: UpdateProfileRequest): Promise => { + const response = await apiClient.patch("/auth/me", data) + return normalizeUser(response.data.user) +} diff --git a/apps/web/src/api/auth/wechat.ts b/apps/web/src/api/auth/wechat.ts index 1e41a8844..38da5f7ce 100644 --- a/apps/web/src/api/auth/wechat.ts +++ b/apps/web/src/api/auth/wechat.ts @@ -1,8 +1,14 @@ import apiClient from "../client" -import type { WechatAuthUrlResponse, WechatCallbackResponse } from "./types" +import type { + WechatAuthUrlResponse, + WechatCallbackResponse, + WechatBindUrlResponse, + WechatBindCompleteResponse, + WechatUnbindResponse, +} from "./types" /** - * 获取微信授权链接 + * 获取微信授权链接(登录场景) */ export const getWechatAuthUrl = async (): Promise => { const response = await apiClient.get("/auth/wechat/url") @@ -19,3 +25,30 @@ export const wechatCallback = async ( const response = await apiClient.post("/auth/wechat/callback", { code, state }) return response.data } + +/** + * 获取微信绑定授权链接(已登录用户绑定场景) + */ +export const getWechatBindUrl = async (): Promise => { + const response = await apiClient.get("/auth/wechat/bind/url") + return response.data +} + +/** + * 微信绑定完成(扫码回调后用 code 绑定到当前登录账号) + */ +export const bindWechat = async ( + code: string, + state: string, +): Promise => { + const response = await apiClient.post("/auth/wechat/bind", { code, state }) + return response.data +} + +/** + * 解绑微信 + */ +export const unbindWechat = async (): Promise => { + const response = await apiClient.delete("/auth/wechat/bind") + return response.data +} diff --git a/apps/web/src/api/auth/wxLogin.ts b/apps/web/src/api/auth/wxLogin.ts new file mode 100644 index 000000000..98706076f --- /dev/null +++ b/apps/web/src/api/auth/wxLogin.ts @@ -0,0 +1,112 @@ +/** + * 微信扫码登录 WxLogin JS-SDK 动态加载与授权参数解析 + * + * 微信官网嵌入式二维码方案:页面引入 https://res.wx.qq.com/connect/zh_CN/htmledition/js/wxLogin.js + * 后挂载全局 window.WxLogin,new WxLogin({...}) 会在指定容器内渲染二维码 iframe。 + * 本模块负责:动态加载该脚本(带超时/失败检测)、从后端返回的 auth_url 中解析 + * WxLogin 所需的 appid / redirect_uri / state。 + */ + +const WX_LOGIN_SRC = "https://res.wx.qq.com/connect/zh_CN/htmledition/js/wxLogin.js" +/** 脚本加载超时(毫秒):超时视为加载失败,调用方回退整页跳转 */ +const WX_LOGIN_LOAD_TIMEOUT = 8000 + +/** WxLogin 构造参数(微信官方字段,保持原名) */ +export interface WxLoginOptions { + /** 是否内嵌二维码(回调在 iframe 内完成) */ + self_redirect: boolean + /** 二维码容器元素 id */ + id: string + /** 微信开放平台 AppID */ + appid: string + /** 应用授权作用域,网站应用固定 snsapi_login */ + scope: "snsapi_login" + /** 回调地址(需与微信开放平台配置一致,WxLogin 内部会 encodeURIComponent) */ + redirect_uri: string + /** 防 CSRF 随机串,由后端 state store 生成并在回调时一次性消费 */ + state: string + /** 二维码样式:black / white */ + style?: "black" | "white" + /** 自定义样式链接(可选) */ + href?: string +} + +/** 微信脚本挂载到 window 上的全局构造函数类型 */ +export interface WxLoginConstructor { + new (options: WxLoginOptions): unknown +} + +declare global { + interface Window { + WxLogin?: WxLoginConstructor + } +} + +let loadPromise: Promise | null = null + +/** + * 动态加载微信 WxLogin JS(单例:并发调用复用同一个 promise)。 + * 加载失败或超时会 reject,调用方应回退到整页跳转授权方式。 + */ +export function loadWxLoginScript(): Promise { + if (window.WxLogin) return Promise.resolve(window.WxLogin) + if (loadPromise) return loadPromise + + loadPromise = new Promise((resolve, reject) => { + const script = document.createElement("script") + script.src = WX_LOGIN_SRC + script.async = true + script.onload = () => { + if (window.WxLogin) { + resolve(window.WxLogin) + } else { + loadPromise = null + reject(new Error("微信登录脚本加载完成但 WxLogin 未挂载")) + } + } + script.onerror = () => { + loadPromise = null + script.remove() + reject(new Error("微信登录脚本加载失败")) + } + document.head.appendChild(script) + + // 超时兜底:部分网络环境下脚本既不 onload 也不 onerror + window.setTimeout(() => { + if (window.WxLogin) { + resolve(window.WxLogin) + return + } + loadPromise = null + script.remove() + reject(new Error("微信登录脚本加载超时")) + }, WX_LOGIN_LOAD_TIMEOUT) + }) + + return loadPromise +} + +/** 从微信授权链接 query 中解析出的 WxLogin 所需参数 */ +export interface ParsedWxAuthParams { + appid: string + /** 已 URL 解码的回调地址(传给 WxLogin 时由其内部再次编码) */ + redirect_uri: string + state: string +} + +/** + * 从后端返回的微信授权链接(https://open.weixin.qq.com/connect/qrconnect?appid=...&redirect_uri=...&state=...) + * 中解析 appid / redirect_uri / state。解析失败时返回 null,由调用方回退整页跳转。 + */ +export function parseWxAuthUrl(authUrl: string, stateFallback?: string): ParsedWxAuthParams | null { + try { + const url = new URL(authUrl) + const appid = url.searchParams.get("appid") + const redirectUri = url.searchParams.get("redirect_uri") + const state = url.searchParams.get("state") || stateFallback || "" + if (!appid || !redirectUri || !state) return null + return { appid, redirect_uri: redirectUri, state } + } catch { + return null + } +} diff --git a/apps/web/src/api/errors.ts b/apps/web/src/api/errors.ts new file mode 100644 index 000000000..7eacd4bc9 --- /dev/null +++ b/apps/web/src/api/errors.ts @@ -0,0 +1,132 @@ +/** + * 统一错误信息提取 + * 把 axios 错误(后端 detail / FastAPI 校验错误 / HTTP 状态码)、XHR/OSS 错误、 + * 网络/超时错误、普通 Error 统一转成「可直接展示给用户」的中文信息。 + * + * 与 api/client.ts 响应拦截器的提示口径保持一致;拦截器负责全局 toast, + * 页面/队列卡片用本工具把真实原因展示在持久位置(回调页、失败卡片等)。 + */ +import type { AxiosError } from "axios" + +/** 后端错误响应体可能出现的字段(FastAPI:detail;历史接口:message/msg) */ +interface ErrorBody { + detail?: unknown + message?: unknown + msg?: unknown +} + +/** FastAPI 422 校验错误单项 */ +interface ValidationItem { + loc?: (string | number)[] + msg?: string +} + +/** 从后端响应体提取人类可读信息(detail 可能是字符串、对象、422 数组) */ +function extractBodyMessage(data: unknown): string { + if (!data || typeof data !== "object") return "" + const body = data as ErrorBody + + const walk = (val: unknown): string => { + if (typeof val === "string") return val + if (Array.isArray(val)) { + // FastAPI 422: [{loc, msg, type}, ...] → 取每条 msg 拼接 + const parts = val + .map((item) => { + if (typeof item === "string") return item + if (item && typeof item === "object") { + const v = item as ValidationItem + if (typeof v.msg === "string") { + const field = Array.isArray(v.loc) ? v.loc.filter((x) => x !== "body").join(".") : "" + return field ? `${field}: ${v.msg}` : v.msg + } + return walk(item) + } + return "" + }) + .filter(Boolean) + return parts.join(";") + } + if (val && typeof val === "object") { + const obj = val as Record + if (typeof obj.message === "string") return obj.message + if (typeof obj.msg === "string") return obj.msg + if (typeof obj.detail === "string") return obj.detail + if (obj.message && typeof obj.message === "object") return walk(obj.message) + if (obj.msg && typeof obj.msg === "object") return walk(obj.msg) + try { + return JSON.stringify(val) + } catch { + return "" + } + } + return "" + } + + return walk(body.detail) || walk(body.message) || walk(body.msg) +} + +/** 无响应体时按 HTTP 状态码给出兜底提示(与 client.ts 拦截器口径一致) */ +function statusFallback(status: number): string { + switch (status) { + case 400: + return "请求参数有误(HTTP 400)" + case 401: + return "登录状态已失效,请重新登录(HTTP 401)" + case 403: + return "没有权限执行该操作(HTTP 403)" + case 404: + return "请求的资源不存在(HTTP 404)" + case 409: + return "操作冲突,资源状态已变化(HTTP 409)" + case 413: + return "文件过大,请缩小后重试(HTTP 413)" + case 415: + return "不支持的文件格式(HTTP 415)" + case 429: + return "操作过于频繁,请稍后再试(HTTP 429)" + case 503: + return "服务暂不可用,请稍后再试(HTTP 503)" + default: + if (status >= 500) return `服务器繁忙,请稍后再试(HTTP ${status})` + return `请求失败(HTTP ${status})` + } +} + +/** + * 从任意抛出值提取可展示的错误信息。 + * @param fallback 全部提取失败时的兜底文案 + */ +export function getErrorMessage(err: unknown, fallback = "操作失败,请稍后重试"): string { + if (!err) return fallback + + // axios 错误(后端 JSON 响应 / HTTP 错误状态) + const ax = err as AxiosError + if (ax.isAxiosError || (typeof ax === "object" && "response" in (ax as object))) { + // 超时 + if (ax.code === "ECONNABORTED" || /timeout/i.test(ax.message || "")) { + return "请求超时,请检查网络后重试" + } + const resp = ax.response + if (resp) { + const bodyMsg = extractBodyMessage(resp.data) + if (bodyMsg) return bodyMsg + return statusFallback(resp.status) + } + // 请求已发出但无响应(断网/CORS/DNS) + if (ax.request) return "网络连接异常,请检查网络设置" + return ax.message || fallback + } + + if (err instanceof Error) { + // XHR 直传 OSS 失败等场景自带详细 message(含 HTTP 状态 + OSS Code/Message) + if (err.message) return err.message + } + if (typeof err === "string") return err + + return fallback +} + +/** client.ts 拦截器是否已对该错误弹过全局 toast(__msgShown 标记) */ +export function isErrorMsgShown(err: unknown): boolean { + return Boolean((err as { __msgShown?: boolean } | null)?.__msgShown) +} diff --git a/apps/web/src/api/generation/types.ts b/apps/web/src/api/generation/types.ts index 6b877681e..61ae5749c 100755 --- a/apps/web/src/api/generation/types.ts +++ b/apps/web/src/api/generation/types.ts @@ -33,15 +33,36 @@ export interface CreatePreviewRequest { preset_id?: string volume?: number } + /** 批量预览数量(1~10),默认1。N>1 时返回 N 个独立变体任务 */ + preview_count?: number + /** 各变体独立标题文字:长度1=共用,长度=preview_count=独立,空数组=使用 title_config.text */ + titles?: string[] + /** 各变体独立配音素材库ID:长度1=共用,长度=preview_count=独立,空数组=回退 voice_library_id */ + voice_library_ids?: string[] + /** 各变体独立封面URL:长度1=共用,长度=preview_count=独立(预览阶段通常为空) */ + cover_urls?: string[] } -/** 创建预览任务响应 */ -export interface CreatePreviewResponse { +/** 单个预览变体任务 */ +export interface PreviewVariantItem { task_id: string - status: PreviewStatus + status: string + progress: number is_preview: boolean + variant_index: number resolution: string - created_at: string + video_url: string + duration: number + error_message: string + title_text: string + voice_library_id: string + created_at?: string | null +} + +/** 创建预览任务响应(单变体,preview_count=1 时 items 长度为1) */ +export interface CreatePreviewResponse { + items: PreviewVariantItem[] + total: number /** 后端自动关联的编辑计划 ID(用于 fallback 路径传递 source_edit_plan_id) */ source_edit_plan_id?: string } diff --git a/apps/web/src/api/tasks/types.ts b/apps/web/src/api/tasks/types.ts index 42e41995b..246c8debe 100644 --- a/apps/web/src/api/tasks/types.ts +++ b/apps/web/src/api/tasks/types.ts @@ -92,6 +92,14 @@ export interface CreateGenerationTaskRequest { preset_id?: string volume?: number } + /** 批量生成数量(1~10),默认1。不传=单条旧逻辑 */ + count?: number + /** 各变体独立标题文字:长度1=共用,长度=count=独立,空数组=使用 title_config/custom_title */ + titles?: string[] + /** 各变体独立配音素材库ID:长度1=共用,长度=count=独立,空数组=回退 voice_library_id */ + voice_library_ids?: string[] + /** 各变体独立封面URL:长度1=共用,长度=count=独立,空数组=回退 cover_url */ + cover_urls?: string[] } /** 单个生成任务详情(对齐后端 GenerationTaskResponse) */ diff --git a/apps/web/src/components/auth/WechatQrModal/WechatQrModal.css b/apps/web/src/components/auth/WechatQrModal/WechatQrModal.css new file mode 100644 index 000000000..3f5ea062d --- /dev/null +++ b/apps/web/src/components/auth/WechatQrModal/WechatQrModal.css @@ -0,0 +1,74 @@ +.xx-wechat-qr-modal { + position: relative; + padding: 8px 0 4px; + min-height: 320px; + display: flex; + flex-direction: column; + align-items: center; + justify-content: center; +} + +/* 常驻二维码容器(WxLogin 渲染目标) */ +.xx-wechat-qr-container { + display: flex; + justify-content: center; + min-height: 260px; +} + +/* loading / error 遮罩层,覆盖在二维码容器之上 */ +.xx-wechat-qr-overlay { + position: absolute; + inset: 0; + display: flex; + flex-direction: column; + align-items: center; + justify-content: center; + background: #fff; + text-align: center; + color: #666; +} + +.xx-wechat-qr-overlay p { + margin-top: 16px; + margin-bottom: 0; +} + +.xx-wechat-qr-container iframe { + border: none; +} + +.xx-wechat-qr-tip { + margin: 12px 0 0; + color: #666; + font-size: 14px; +} + +.xx-wechat-qr-error { + text-align: center; + width: 100%; +} + +.xx-wechat-qr-error-msg { + color: #ef4444; + font-size: 14px; + line-height: 1.6; + margin: 0 0 16px; + word-break: break-word; +} + +.xx-wechat-qr-error-actions { + display: flex; + flex-direction: column; + align-items: center; + gap: 12px; +} + +.xx-wechat-qr-fallback { + background: none; + border: none; + color: var(--primary-color, #3b82f6); + cursor: pointer; + font-size: 13px; + padding: 0; + text-decoration: underline; +} diff --git a/apps/web/src/components/auth/WechatQrModal/index.tsx b/apps/web/src/components/auth/WechatQrModal/index.tsx new file mode 100644 index 000000000..59e254c43 --- /dev/null +++ b/apps/web/src/components/auth/WechatQrModal/index.tsx @@ -0,0 +1,265 @@ +/** + * 微信扫码二维码弹窗(登录 / 绑定复用) + * + * 微信官方嵌入式二维码方案:弹窗内用 new WxLogin({ self_redirect: true }) 渲染二维码, + * 扫码后微信重定向到本站回调页(在二维码 iframe 内加载),回调页通过 postMessage + * 把成功/失败结果通知本弹窗(消息协议见 ./messages)。 + * + * 兜底:获取授权链接成功但 WxLogin JS 加载失败/超时时,自动回退整页跳转授权 + * (与旧流程一致);获取授权链接本身失败时在弹窗内展示错误并提供重试。 + */ +import React, { useEffect, useRef, useState } from "react" +import { Spin } from "antd" +import Modal from "@/components/ui/Modal" +import Button from "@/components/ui/Button" +import { + getWechatAuthUrl, + getWechatBindUrl, + getCurrentUser, + normalizeUser, + type User, +} from "@/api/auth" +import { useAuthStore } from "@/store/authStore" +import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh" +import { getErrorMessage } from "@/api/errors" +import { loadWxLoginScript, parseWxAuthUrl } from "@/api/auth/wxLogin" +import { isWechatQrMessage, type WechatQrScene } from "./messages" +import "./WechatQrModal.css" + +export interface WechatQrModalProps { + open: boolean + scene: WechatQrScene + onClose: () => void + /** 登录场景成功回调(needOnboarding=true 时调用方应跳昵称引导页) */ + onLoginSuccess?: (needOnboarding: boolean) => void + /** 绑定场景成功回调(调用方刷新用户信息/提示) */ + onBindSuccess?: () => void +} + +type QrStatus = "loading" | "qrcode" | "error" + +const CONTAINER_ID: Record = { + login: "wechat-qr-login-container", + bind: "wechat-qr-bind-container", +} + +const STATE_STORAGE_KEY: Record = { + login: "wechat_state", + bind: "wechat_bind_state", +} + +/** + * 等待二维码容器挂载到 DOM。antd Modal 内容通过 portal 渲染且带进场动画, + * 父组件 effect 首次执行时容器可能尚未出现在 document 中。 + */ +function waitForContainer(id: string, timeoutMs = 3000): Promise { + return new Promise((resolve) => { + const start = Date.now() + const check = () => { + const el = document.getElementById(id) + if (el) { + resolve(el) + return + } + if (Date.now() - start > timeoutMs) { + resolve(null) + return + } + setTimeout(check, 50) + } + check() + }) +} + +const WechatQrModal: React.FC = ({ + open, + scene, + onClose, + onLoginSuccess, + onBindSuccess, +}) => { + const setAuth = useAuthStore((state) => state.setAuth) + const setUser = useAuthStore((state) => state.setUser) + const [status, setStatus] = useState("loading") + const [errorMsg, setErrorMsg] = useState("") + /** 刷新二维码计数:变化时重新请求授权链接并重渲染 */ + const [renderSeq, setRenderSeq] = useState(0) + /** 最新授权链接,用于"整页打开"兜底 */ + const authUrlRef = useRef(null) + + const isLogin = scene === "login" + + // 初始化:获取授权链接 → 加载 WxLogin JS → 内嵌渲染二维码 + useEffect(() => { + if (!open) return + let cancelled = false + authUrlRef.current = null + setStatus("loading") + setErrorMsg("") + + const init = async () => { + try { + const fetchUrl = isLogin ? getWechatAuthUrl : getWechatBindUrl + const result = await fetchUrl() + if (cancelled) return + // 写 state(整页跳转兜底路径的回调页也会清理它) + localStorage.setItem(STATE_STORAGE_KEY[scene], result.state) + authUrlRef.current = result.auth_url + + const params = parseWxAuthUrl(result.auth_url, result.state) + if (!params) { + // 授权链接格式异常:直接整页跳转,由微信侧/回调页兜底 + window.location.href = result.auth_url + return + } + + const WxLogin = await loadWxLoginScript() + if (cancelled) return + // 等 Modal portal 中的容器挂载完成 + const container = await waitForContainer(CONTAINER_ID[scene]) + if (cancelled) return + if (!container) { + window.location.href = result.auth_url + return + } + container.innerHTML = "" + new WxLogin({ + self_redirect: true, + id: CONTAINER_ID[scene], + appid: params.appid, + scope: "snsapi_login", + redirect_uri: params.redirect_uri, + state: params.state, + style: "black", + }) + if (!cancelled) setStatus("qrcode") + } catch (err) { + if (cancelled) return + if (authUrlRef.current) { + // 授权链接已拿到但二维码脚本加载失败/超时:回退整页跳转 + window.location.href = authUrlRef.current + return + } + // 授权链接接口本身失败:弹窗内展示真实原因,允许重试 + setErrorMsg(getErrorMessage(err, "微信服务暂不可用,请稍后重试")) + setStatus("error") + } + } + + init() + return () => { + cancelled = true + } + }, [open, scene, isLogin, renderSeq]) + + // 监听 iframe 内回调页 postMessage 回来的扫码结果 + useEffect(() => { + if (!open) return + + const handleMessage = async (event: MessageEvent) => { + // 只接受同源消息 + if (event.origin !== window.location.origin) return + if (!isWechatQrMessage(event.data, scene)) return + const msg = event.data + + if (msg.success) { + if (isLogin) { + // iframe 内回调页已把 token 写入 localStorage(同源共享), + // 父窗口同步内存登录态后交给调用方跳转 + try { + const userData = await getCurrentUser() + const user = normalizeUser(userData) as User + setAuth( + user, + localStorage.getItem("access_token") || "", + localStorage.getItem("refresh_token"), + ) + scheduleProactiveRefresh() + } catch { + // token 已持久化,即使这里失败路由守卫/刷新也能恢复登录态 + } + onLoginSuccess?.(msg.payload?.needOnboarding ?? false) + } else { + try { + const userData = await getCurrentUser() + setUser(normalizeUser(userData) as User) + } catch { + // 绑定结果以后端为准,调用方 invalidateQueries 会兜底刷新 + } + onBindSuccess?.() + } + return + } + + // 失败:弹窗内展示回调页透传的真实原因,提供刷新/整页跳转 + setErrorMsg(msg.detail || "微信授权失败,请重试") + setStatus("error") + } + + window.addEventListener("message", handleMessage) + return () => window.removeEventListener("message", handleMessage) + }, [open, scene, isLogin, onLoginSuccess, onBindSuccess, setAuth, setUser]) + + const handleRefresh = () => setRenderSeq((seq) => seq + 1) + + const handleFullPageRedirect = () => { + if (authUrlRef.current) { + window.location.href = authUrlRef.current + } + } + + return ( + +
+ {/* 二维码容器常驻:WxLogin 在 loading 阶段就会把 iframe 渲染进来, + 不能按 status 条件渲染,否则 effect 里永远找不到容器 */} +
+ + {status === "loading" && ( +
+ +

正在生成微信二维码...

+
+ )} + + {status === "qrcode" && ( +

请使用微信扫描二维码{isLogin ? "登录" : "绑定账号"}

+ )} + + {status === "error" && ( +
+

{errorMsg}

+
+ + {authUrlRef.current && ( + + )} +
+
+ )} +
+ + ) +} + +export default WechatQrModal diff --git a/apps/web/src/components/auth/WechatQrModal/messages.ts b/apps/web/src/components/auth/WechatQrModal/messages.ts new file mode 100644 index 000000000..365eba137 --- /dev/null +++ b/apps/web/src/components/auth/WechatQrModal/messages.ts @@ -0,0 +1,71 @@ +/** + * 微信扫码弹窗与 iframe 内回调页之间的 postMessage 消息协议 + * + * 流程:弹窗内 WxLogin(self_redirect:true) 渲染的二维码 iframe 扫码后, + * 微信重定向到本站回调页(同源,在 iframe 内加载);回调页完成换 token/绑定后, + * 通过 window.parent.postMessage 把结果通知弹窗,弹窗负责关闭/展示错误/同步登录态。 + */ + +/** 扫码场景:登录 / 绑定 */ +export type WechatQrScene = "login" | "bind" + +export interface WechatQrSuccessPayload { + /** 登录场景:是否需要昵称引导(新用户或资料未完善) */ + needOnboarding?: boolean +} + +export interface WechatQrMessageData { + /** 固定协议标识,父窗口只认该 source */ + source: "xiaoxia-wechat-qr" + /** 场景,需与弹窗发起时一致(login/bind),父窗口据此过滤 */ + scene: WechatQrScene + /** 成功 / 失败 */ + success: boolean + /** 失败时的真实原因(已在回调页拼好,含后端 detail) */ + detail?: string + payload?: WechatQrSuccessPayload +} + +export const WECHAT_QR_MESSAGE_SOURCE = "xiaoxia-wechat-qr" + +/** 判断收到的 message 是否为本协议消息(且场景匹配) */ +export function isWechatQrMessage( + data: unknown, + scene: WechatQrScene, +): data is WechatQrMessageData { + if (!data || typeof data !== "object") return false + const msg = data as Partial + return msg.source === WECHAT_QR_MESSAGE_SOURCE && msg.scene === scene +} + +/** 当前页面是否运行在 iframe(弹窗内嵌二维码)中 */ +export function isInIframe(): boolean { + try { + return window.parent !== window + } catch { + // 跨域访问 window.parent 可能抛异常,按非 iframe 处理 + return false + } +} + +/** + * iframe 内回调页向父窗口上报扫码结果。同源回调页加载,targetOrigin 限定本站 origin。 + */ +export function postWechatQrResult( + scene: WechatQrScene, + success: boolean, + options?: { detail?: string; needOnboarding?: boolean }, +): void { + if (!isInIframe()) return + const data: WechatQrMessageData = { + source: WECHAT_QR_MESSAGE_SOURCE, + scene, + success, + detail: options?.detail, + payload: + success && options?.needOnboarding !== undefined + ? { needOnboarding: options.needOnboarding } + : undefined, + } + window.parent.postMessage(data, window.location.origin) +} diff --git a/apps/web/src/components/common/ChunkErrorBoundary.tsx b/apps/web/src/components/common/ChunkErrorBoundary.tsx new file mode 100644 index 000000000..5403751c6 --- /dev/null +++ b/apps/web/src/components/common/ChunkErrorBoundary.tsx @@ -0,0 +1,85 @@ +/** + * 全局错误边界:专门兜底"发版后旧标签页懒加载 chunk 失效"导致的白屏, + * 同时兜住页面级渲染崩溃,避免任何未捕获错误导致整页白屏无反馈。 + * + * 捕获到 ChunkLoadError / Failed to fetch dynamically imported module: + * 1. 首次:自动整页刷新一次(sessionStorage 标记,刷新后 index.html 重新拉取, + * 拿到新 chunk 引用,白屏自愈) + * 2. 刷新后仍失败(标记未过期):不再自动刷新,显示"系统已更新,请点击刷新" + * 兜底界面,由用户手动点击 + * + * 其他非 chunk 错误:显示通用错误页 + "返回首页"按钮(跳首页而非刷新当前 URL, + * 避免刷新后再次命中同一路由崩溃形成死循环)。 + */ +import React from "react" +import { Button, Result } from "antd" +import { + getChunkReloadedAt, + goHomeRecover, + isChunkLoadError, + reloadForChunkError, +} from "@/utils/chunkLoadError" + +interface Props { + children: React.ReactNode +} + +interface State { + error: Error | null + isChunkError: boolean + /** 捕获错误时是否已经自动刷新过(决定显示自动刷新中还是手动兜底) */ + alreadyReloaded: boolean +} + +class ChunkErrorBoundary extends React.Component { + state: State = { error: null, isChunkError: false, alreadyReloaded: false } + + static getDerivedStateFromError(error: Error): State { + const chunk = isChunkLoadError(error) + return { + error, + isChunkError: chunk, + alreadyReloaded: chunk ? getChunkReloadedAt() !== null : false, + } + } + + componentDidCatch(error: Error): void { + // 仅 chunk 错误且本次会话没自动刷新过 → 打标记并整页刷新(自愈) + if (isChunkLoadError(error) && getChunkReloadedAt() === null) { + reloadForChunkError() + } + } + + render(): React.ReactNode { + const { error, isChunkError, alreadyReloaded } = this.state + if (!error) return this.props.children + + if (isChunkError && !alreadyReloaded) { + // 已打标记、componentDidCatch 里已触发 reload;极短瞬间展示加载中 + return ( + + ) + } + + // 手动兜底统一跳首页(整页导航):chunk 失效时脱离旧 chunk 引用; + // 业务崩溃时绕开当前报错路由,避免刷新-再崩死循环 + return ( + + {isChunkError ? "刷新并返回首页" : "返回首页"} + + } + /> + ) + } +} + +export default ChunkErrorBoundary diff --git a/apps/web/src/main.tsx b/apps/web/src/main.tsx index a4253a61a..33cdde527 100644 --- a/apps/web/src/main.tsx +++ b/apps/web/src/main.tsx @@ -9,6 +9,7 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query" import { ConfigProvider, App as AntApp } from "antd" import zhCN from "antd/locale/zh_CN" import router from "./router" +import ChunkErrorBoundary from "./components/common/ChunkErrorBoundary" import { scheduleProactiveRefresh } from "./api/auth/tokenRefresh" // 应用启动时,如果用户已登录,立即调度主动 token 刷新 @@ -99,7 +100,9 @@ ReactDOM.createRoot(document.getElementById("root")!).render( - + + + diff --git a/apps/web/src/pages/assets/AssetLibrary.tsx b/apps/web/src/pages/assets/AssetLibrary.tsx index 17e379623..5e6ace7e7 100644 --- a/apps/web/src/pages/assets/AssetLibrary.tsx +++ b/apps/web/src/pages/assets/AssetLibrary.tsx @@ -42,6 +42,7 @@ const AssetLibrary: React.FC = () => { assetsError, assetsErrorObj, refetchAssets, + stalledAssetIds, searchText, setSearchText, filterType, @@ -77,6 +78,7 @@ const AssetLibrary: React.FC = () => { removeUpload, clearFinished, uploading, + transferActive, activeCount, pendingCount, } = useAssetUpload({ effectiveLibId }) @@ -168,6 +170,7 @@ const AssetLibrary: React.FC = () => { {/* 上传区域 */} { selectedIds={selectedIds} diagnosingId={diagnosingId} uploadProgressMap={uploadProgressMap} + stalledAssetIds={stalledAssetIds} onRetry={refetchAssets} onToggleSelect={toggleSelect} onDiagnose={handleDiagnose} diff --git a/apps/web/src/pages/assets/assets.css b/apps/web/src/pages/assets/assets.css index f237ae738..29043aee7 100644 --- a/apps/web/src/pages/assets/assets.css +++ b/apps/web/src/pages/assets/assets.css @@ -831,6 +831,20 @@ color: #ef4444; } +.xx-upload-queue-error-detail { + margin-top: 4px; + font-size: 12px; + line-height: 1.5; + color: #ef4444; + word-break: break-word; + white-space: normal; +} + +.xx-upload-queue-error-hint { + margin-top: 2px; + color: #b45309; +} + .xx-upload-queue-actions { display: flex; gap: 6px; @@ -1041,3 +1055,25 @@ background: #fef2f2; color: #dc2626; } + +/* 上传入口禁用态(直传进行中,防重复提交,Issue #1714) */ +.xx-asset-upload-btn:disabled { + opacity: 0.6; + cursor: not-allowed; +} +.xx-asset-upload-btn:disabled:hover { + opacity: 0.6; +} +.xx-asset-upload-btn:disabled:active { + transform: none; +} + +/* 处理超时遮罩:创建超过 10 分钟仍在处理中(疑似后端卡住),停止转圈并警示 */ +.xx-asset-thumb-stalled { + background: rgba(217, 119, 6, 0.28); + color: #fde68a; + backdrop-filter: blur(2px); +} +.xx-asset-thumb-stalled :first-child { + font-size: var(--font-size-2xl); +} diff --git a/apps/web/src/pages/assets/components/AssetCard.tsx b/apps/web/src/pages/assets/components/AssetCard.tsx index b0b0ea21c..e88a3f6fd 100644 --- a/apps/web/src/pages/assets/components/AssetCard.tsx +++ b/apps/web/src/pages/assets/components/AssetCard.tsx @@ -22,6 +22,8 @@ export interface AssetCardProps { diagnosing?: boolean /** 上传中实时进度(仅 uploading 态有值;ingesting 后由后端状态接管) */ uploadProgress?: { progress: number; uploading: boolean } + /** 处理超过 10 分钟仍未就绪(疑似后端卡住):停止转圈并提示处理超时 */ + stalled?: boolean onToggle: () => void onDiagnose: () => void onPlay: () => void @@ -33,6 +35,7 @@ const AssetCard: React.FC = ({ selected, diagnosing, uploadProgress, + stalled, onToggle, onDiagnose, onPlay, @@ -67,11 +70,15 @@ const AssetCard: React.FC = ({
)} - {/* 转码/处理中遮罩 */} + {/* 转码/处理中遮罩(卡死超过 10 分钟时停止转圈,提示超时) */} {asset.loading && !isUploading && ( -
- - 转码处理中 +
+ {stalled ? : } + {stalled ? "处理超时,可重试上传" : "转码处理中"}
)} @@ -129,7 +136,7 @@ const AssetCard: React.FC = ({

- + {asset.duration && {asset.duration}}
diff --git a/apps/web/src/pages/assets/components/AssetGridSection.tsx b/apps/web/src/pages/assets/components/AssetGridSection.tsx index c88307734..2af6d87ff 100644 --- a/apps/web/src/pages/assets/components/AssetGridSection.tsx +++ b/apps/web/src/pages/assets/components/AssetGridSection.tsx @@ -19,6 +19,8 @@ export interface AssetGridSectionProps { selectedIds: Set diagnosingId: string | null uploadProgressMap?: UploadProgressMap + /** 创建超过 10 分钟仍在处理中的素材 id(疑似后端卡住),卡片提示处理超时 */ + stalledAssetIds?: Set onRetry?: () => void onToggleSelect: (id: string) => void onDiagnose: (asset: AssetItem) => void @@ -34,6 +36,7 @@ export const AssetGridSection: React.FC = ({ selectedIds, diagnosingId, uploadProgressMap, + stalledAssetIds, onRetry, onToggleSelect, onDiagnose, @@ -76,6 +79,7 @@ export const AssetGridSection: React.FC = ({ selected={selectedIds.has(asset.id)} diagnosing={diagnosingId === asset.id} uploadProgress={uploadProgressMap?.get(asset.id)} + stalled={stalledAssetIds?.has(asset.id)} onToggle={() => onToggleSelect(asset.id)} onDiagnose={() => onDiagnose(asset)} onPlay={() => onPlay(asset)} diff --git a/apps/web/src/pages/assets/components/AssetUploadZone.tsx b/apps/web/src/pages/assets/components/AssetUploadZone.tsx index bb9abdac8..c06d3f515 100644 --- a/apps/web/src/pages/assets/components/AssetUploadZone.tsx +++ b/apps/web/src/pages/assets/components/AssetUploadZone.tsx @@ -4,10 +4,13 @@ * - 拖拽文件到内容区任意位置同样触发上传(不再占用大面积虚线框) */ import React, { useRef, useState } from "react" +import { message } from "antd" import { PlusOutlined, CloudUploadOutlined } from "@ant-design/icons" export interface AssetUploadZoneProps { uploading: boolean + /** 有文件正在本地指纹/prepare/直传(非服务端转码),此时禁用入口防重复提交 */ + transferActive: boolean activeCount: number pendingCount: number onUpload: (files: File[]) => void @@ -15,6 +18,7 @@ export interface AssetUploadZoneProps { export const AssetUploadZone: React.FC = ({ uploading, + transferActive, activeCount, pendingCount, onUpload, @@ -27,6 +31,11 @@ export const AssetUploadZone: React.FC = ({ const pickFiles = (list: FileList | null) => { if (!list || list.length === 0) return + // 直传进行中拦截重复触发:相同文件仍由入队去重兜底,这里先给明确反馈 + if (transferActive) { + message.warning("文件正在上传中,请等待当前上传完成后再添加") + return + } onUpload(Array.from(list)) } @@ -58,10 +67,18 @@ export const AssetUploadZone: React.FC = ({ {uploading ? ( diff --git a/apps/web/src/pages/assets/components/UploadQueuePanel.tsx b/apps/web/src/pages/assets/components/UploadQueuePanel.tsx index 1dead1405..e84823792 100644 --- a/apps/web/src/pages/assets/components/UploadQueuePanel.tsx +++ b/apps/web/src/pages/assets/components/UploadQueuePanel.tsx @@ -12,7 +12,8 @@ import { ReloadOutlined, CloseOutlined, } from "@ant-design/icons" -import type { UploadItem } from "../hooks/useAssetUpload" +import type { UploadItem, UploadFailStage } from "../hooks/useAssetUpload" +import { COMPLETE_RETRY_HINT } from "../hooks/useAssetUpload" export interface UploadQueuePanelProps { items: UploadItem[] @@ -29,6 +30,13 @@ const STATUS_TEXT: Record = { error: "上传失败", } +/** 失败阶段中文名:让用户一眼看到失败发生在哪一步 */ +const FAIL_STAGE_TEXT: Record = { + prepare: "准备上传阶段", + transfer: "文件传输阶段", + complete: "确认入库阶段", +} + const UploadQueuePanel: React.FC = ({ items, onRetry, @@ -80,16 +88,34 @@ const UploadQueuePanel: React.FC = ({ ) : null}
{it.duplicated ? "素材已存在,已跳过" : STATUS_TEXT[it.status]} + {it.status === "preparing" && it.hint ? `(${it.hint})` : ""} {it.status === "uploading" ? ` ${it.progress}%` : ""} - {it.status === "error" && it.error ? `:${it.error}` : ""} + {it.status === "error" && it.failedStage + ? `(${FAIL_STAGE_TEXT[it.failedStage]})` + : ""}
+ {it.status === "error" && it.error ? ( +
+ {it.error.split("\n").map((line, idx) => + line === COMPLETE_RETRY_HINT ? ( +
+ {line} +
+ ) : ( +
{line}
+ ), + )} +
+ ) : null}
{it.status === "error" && ( @@ -137,6 +138,13 @@ const Login: React.FC = () => { 还没有账号? 立即注册 + + setWechatQrOpen(false)} + onLoginSuccess={handleWechatQrSuccess} + /> ) } diff --git a/apps/web/src/pages/auth/WechatBindCallback.tsx b/apps/web/src/pages/auth/WechatBindCallback.tsx new file mode 100644 index 000000000..641f0d4bd --- /dev/null +++ b/apps/web/src/pages/auth/WechatBindCallback.tsx @@ -0,0 +1,118 @@ +/** + * 微信绑定回调页(已登录用户在设置页发起"绑定微信"扫码后回到这里) + * 用 code 调绑定接口把微信关联到当前账号,成功后回设置页 + * + * 两种运行环境: + * - 整页跳转授权(旧流程/兜底):本页整页加载,成功/失败后 navigate 回设置页 + * - 弹窗内嵌二维码(WxLogin self_redirect):本页在同源 iframe 内加载, + * 结果通过 postMessage 通知父窗口弹窗,不做页面导航 + */ +import React, { useEffect, useState } from "react" +import { useSearchParams, useNavigate } from "react-router-dom" +import { Spin } from "antd" +import { bindWechat, normalizeUser } from "@/api/auth" +import { getErrorMessage } from "@/api/errors" +import { useAuthStore } from "@/store/authStore" +import { isInIframe, postWechatQrResult } from "@/components/auth/WechatQrModal/messages" + +const WechatBindCallback: React.FC = () => { + const [searchParams] = useSearchParams() + const navigate = useNavigate() + const setUser = useAuthStore((state) => state.setUser) + const [error, setError] = useState(null) + const inIframe = isInIframe() + + useEffect(() => { + const code = searchParams.get("code") + const state = searchParams.get("state") + + const fail = (message: string) => { + if (inIframe) { + // 弹窗模式:把真实原因上报父窗口在 Modal 内展示 + postWechatQrResult("bind", false, { detail: message }) + return + } + setError(message) + } + + if (!code || !state) { + fail("无效的回调参数,请回到设置页重新扫码绑定") + return + } + + const handleBind = async () => { + // state 校验由后端 state store 一次性消费兜底(前端不再比对 localStorage, + // 微信内打开/跨浏览器场景本地无 state 会误杀);清理绑定前写入的 state + localStorage.removeItem("wechat_bind_state") + + try { + const result = await bindWechat(code, state) + setUser(normalizeUser(result.user)) + + if (inIframe) { + // 弹窗模式:通知父窗口关闭弹窗并刷新绑定状态 + postWechatQrResult("bind", true) + return + } + + // 用 replace 回设置页,query 携带成功标记由设置页提示 + navigate("/app/profile?wechat_bind=success", { replace: true }) + } catch (err) { + // 绑定失败直接在本页展示/上报真实原因(如微信已被其他账号绑定),不静默跳走 + fail(`微信绑定失败:${getErrorMessage(err, "请回到设置页重试")}`) + } + } + + handleBind() + }, [searchParams, navigate, setUser, inIframe]) + + if (error) { + return ( +
+
+

{error}

+ +
+
+ ) + } + + return ( +
+
+ +

正在绑定微信...

+
+
+ ) +} + +export default WechatBindCallback diff --git a/apps/web/src/pages/auth/WechatCallback.tsx b/apps/web/src/pages/auth/WechatCallback.tsx index 3dc75a6e9..2b7ecbaef 100644 --- a/apps/web/src/pages/auth/WechatCallback.tsx +++ b/apps/web/src/pages/auth/WechatCallback.tsx @@ -1,40 +1,63 @@ /** * 微信登录回调页 + * 扫码授权后由微信重定向回来:用 code 换登录态, + * 新用户/资料未完善 → 跳昵称引导页;老用户 → 回来源页/首页 + * + * 两种运行环境: + * - 整页跳转授权(旧流程/兜底):本页整页加载,按上述逻辑导航 + * - 弹窗内嵌二维码(WxLogin self_redirect):本页在同源 iframe 内加载, + * 成功/失败均通过 postMessage 通知父窗口弹窗,不做页面导航 */ import React, { useEffect, useState } from "react" import { useSearchParams, useNavigate } from "react-router-dom" -import { Spin, message } from "antd" +import { Spin } from "antd" import { wechatCallback, getCurrentUser, normalizeUser, type User } from "@/api/auth" +import { getErrorMessage } from "@/api/errors" import { useAuthStore } from "@/store/authStore" -import BindContactModal from "@/components/auth/BindContactModal" +import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh" +import { isInIframe, postWechatQrResult } from "@/components/auth/WechatQrModal/messages" const WechatCallback: React.FC = () => { const [searchParams] = useSearchParams() const navigate = useNavigate() const setAuth = useAuthStore((state) => state.setAuth) const [loading, setLoading] = useState(true) - const [showBindModal, setShowBindModal] = useState(false) const [error, setError] = useState(null) + const inIframe = isInIframe() useEffect(() => { const code = searchParams.get("code") const state = searchParams.get("state") - if (!code || !state) { - setError("无效的回调参数") + const fail = (message: string) => { + if (inIframe) { + // 弹窗模式:把真实原因上报父窗口在 Modal 内展示,本页保持"处理中"即可 + postWechatQrResult("login", false, { detail: message }) + return + } + setError(message) setLoading(false) + } + + // 微信重定向出错时(如用户拒绝授权 error=access_denied)直接展示/上报原因 + const wxErrorCode = searchParams.get("error") + const wxErrDesc = searchParams.get("error_description") + if (wxErrorCode || wxErrDesc) { + const reason = [wxErrorCode, wxErrDesc].filter(Boolean).join(":") + fail(`微信授权失败:${reason}`) + return + } + + if (!code || !state) { + fail("无效的回调参数,请重新扫码登录") return } const handleCallback = async () => { try { - // 校验 state,防止 CSRF - const savedState = localStorage.getItem("wechat_state") - if (!savedState || savedState !== state) { - setError("安全校验失败,请重新登录") - setLoading(false) - return - } + // state 的 CSRF 校验由后端 state store 一次性消费兜底(前端不再比对 + // localStorage——微信内打开、跨浏览器等场景本地没有 state,会误杀正常回调); + // 清理登录前写入的 state,避免残留 localStorage.removeItem("wechat_state") const result = await wechatCallback(code, state) @@ -49,41 +72,34 @@ const WechatCallback: React.FC = () => { const userData = await getCurrentUser() const user: User = normalizeUser(userData) setAuth(user, result.access_token, result.refresh_token) + scheduleProactiveRefresh() - if (result.binding_complete) { - // 已绑定,跳转到登录前页面或首页 - message.success("登录成功") - const redirect = localStorage.getItem("login_redirect") || "/" - localStorage.removeItem("login_redirect") - navigate(redirect, { replace: true }) - } else { - // 未绑定,显示绑定弹窗 - setLoading(false) - setShowBindModal(true) + // 新用户 或 资料未完善(如上次中断没填昵称)→ 强制昵称引导 + const needOnboarding = result.is_new_user || user.profile_completed === false + + if (inIframe) { + // 弹窗模式:token 已写入同源 localStorage,通知父窗口同步登录态并跳转 + postWechatQrResult("login", true, { needOnboarding }) + return } + + if (needOnboarding) { + navigate("/welcome/wechat", { replace: true }) + return + } + + // 老用户:回登录前页面或首页 + const redirect = localStorage.getItem("login_redirect") || "/" + localStorage.removeItem("login_redirect") + navigate(redirect, { replace: true }) } catch (err) { - setError("登录失败,请重试") - setLoading(false) + // 透传后端真实错误(如 state 过期、code 已消费、接口异常),禁止吞成通用提示 + fail(`微信登录失败:${getErrorMessage(err, "请重试或更换登录方式")}`) } } handleCallback() - }, [searchParams, navigate, setAuth]) - - const handleBindSuccess = (user: User) => { - const setUser = useAuthStore.getState().setUser - setUser(user) - setShowBindModal(false) - message.success("绑定成功") - const redirect = localStorage.getItem("login_redirect") || "/" - localStorage.removeItem("login_redirect") - navigate(redirect, { replace: true }) - } - - const handleBindCancel = () => { - setShowBindModal(false) - navigate("/login") - } + }, [searchParams, navigate, setAuth, inIframe]) if (loading) { return ( @@ -98,49 +114,39 @@ const WechatCallback: React.FC = () => { >
-

正在登录...

-
- - ) - } - - if (error) { - return ( -
-
-

{error}

- +

微信登录中...

) } return ( - +
+
+

{error}

+ +
+
) } diff --git a/apps/web/src/pages/auth/WechatOnboarding.tsx b/apps/web/src/pages/auth/WechatOnboarding.tsx new file mode 100644 index 000000000..38c62f345 --- /dev/null +++ b/apps/web/src/pages/auth/WechatOnboarding.tsx @@ -0,0 +1,116 @@ +/** + * 微信新用户昵称引导页 + * 新微信用户首次登录后强制填写昵称,完成后才进入主界面 + */ +import React, { useRef } from "react" +import { Form, Input, message } from "antd" +import { Navigate, useNavigate } from "react-router-dom" +import { useMutation } from "@tanstack/react-query" +import { updateProfile } from "@/api/auth" +import { getErrorMessage, isErrorMsgShown } from "@/api/errors" +import { useAuthStore } from "@/store/authStore" +import Button from "@/components/ui/Button" +import "./Login.css" + +interface OnboardingFormValues { + display_name: string +} + +const WechatOnboarding: React.FC = () => { + const navigate = useNavigate() + const setUser = useAuthStore((state) => state.setUser) + const isAuthenticated = useAuthStore((state) => state.isAuthenticated) + const user = useAuthStore((state) => state.user) + const hasAccessToken = Boolean(localStorage.getItem("access_token")) + const [form] = Form.useForm() + // 同步防连点守卫:antd loading 要等 React 重渲染后才禁用按钮, + // 连点两次时第一次的 mutation 刚触发、重渲染未发生,第二次 click 仍会进来 + // (截图里 PATCH /me 405 出现两次就是连点导致的重复提交) + const submittingRef = useRef(false) + + const saveMutation = useMutation({ + mutationFn: (displayName: string) => updateProfile({ display_name: displayName }), + }) + + // 已登录且资料已完善的用户不该停留在引导页 + if (isAuthenticated && hasAccessToken && user?.profile_completed === true) { + return + } + // 未登录(如手动输入 URL)回登录页 + if (!isAuthenticated || !hasAccessToken) { + return + } + + const onFinish = async (values: OnboardingFormValues) => { + if (submittingRef.current) return + submittingRef.current = true + try { + const updated = await saveMutation.mutateAsync(values.display_name.trim()) + // 后端返回的 profile_completed 以最新资料为准,前端同步标记完善 + setUser({ ...updated, profile_completed: true }) + message.success("欢迎加入小虾智剪!") + const redirect = localStorage.getItem("login_redirect") || "/app/dashboard" + localStorage.removeItem("login_redirect") + navigate(redirect, { replace: true }) + } catch (err) { + // 透传后端真实原因(如接口异常/校验失败);拦截器已弹过的不重复弹 + if (!isErrorMsgShown(err)) { + message.error(`昵称保存失败:${getErrorMessage(err, "请稍后重试")}`) + } + submittingRef.current = false + } + // 成功时页面跳走,不复位 + } + + return ( +
+
+
+
+ 🦐 + 小虾智剪 +
+

欢迎使用微信登录,请先设置您的昵称

+
+ +
+ + + + + + + +
+
+
+ ) +} + +export default WechatOnboarding diff --git a/apps/web/src/pages/generate/GeneratePage.tsx b/apps/web/src/pages/generate/GeneratePage.tsx index db681da43..a31a87f58 100644 --- a/apps/web/src/pages/generate/GeneratePage.tsx +++ b/apps/web/src/pages/generate/GeneratePage.tsx @@ -1,12 +1,12 @@ /** - * 智能剪辑页面 — 前端实时预览架构 - * 6 步向导:选择模板 → 素材 → 配音 → 标题(含预览) → 确认生成 → 选择封面 + * 智能剪辑页面(Issue #1677 多视频批量生成,修正版) + * 固定 6 步向导:模板(弹数量) → 素材 → 配音 → 标题 → 确认生成 → 封面,单视频与批量完全一致 * * 架构: - * - 步骤 4 右侧显示 FrontendPreviewPlayer 实时预览 - * - 步骤 5 右侧内联播放生成中的/最终视频 - * - 步骤 6 封面从最终成片中智能选帧(MediaKit) - * - 点"确认生成"时调用 createGenerationTask 创建一次服务器渲染任务 + * - 预览全部为纯前端 Canvas 实时播放(FrontendPreviewPlayer),不调任何后端渲染接口: + * N=1 单播放器;N>1 CanvasPreviewGrid(variantSeed 让素材排布/起始点不同,画面有差异) + * - 步骤5确认生成:正式生成接口(count + titles[]/voice_library_ids[]/cover_urls[]), + * 批量时逐任务独立进度/失败重试(BatchGenerationGrid) */ import React, { useMemo, useState, useEffect, useRef, useCallback } from "react" import { message } from "antd" @@ -17,12 +17,15 @@ import { useCloneProgress } from "@/hooks/useCloneProgress" import CloneModal from "@/components/voice/CloneModal" import GenerateHeader from "./components/GenerateHeader" import FrontendPreviewPlayer from "./components/FrontendPreviewPlayer" +import CanvasPreviewGrid from "./components/CanvasPreviewGrid" +import PreviewCountModal from "./components/PreviewCountModal" import GenerateStepsBar from "./components/GenerateStepsBar" import GenerateStepContent from "./components/GenerateStepContent" import GenerateStepActions from "./components/GenerateStepActions" import { useGenerateFormState } from "./hooks/useGenerateFormState" import { useStepNavigation } from "./hooks/useStepNavigation" import { useGenerateVideo } from "./hooks/useGenerateVideo" + import { usePreviewAssets } from "./hooks/usePreviewAssets" import { useTitleStyleUpdaters } from "./hooks/useStep4Title/useTitleStyleUpdaters" import { getAssetsByKind } from "@/api/assets" @@ -53,10 +56,8 @@ const GeneratePage: React.FC = () => { selectedVoice, setSelectedVoice, voiceMode, - setVoiceMode, selectedClonedVoice, - setSelectedClonedVoice, - presetVoices, + cloneModalOpen, setCloneModalOpen, videoRatio, @@ -72,8 +73,51 @@ const GeneratePage: React.FC = () => { setStoredSourceEditPlanId, serverClips, setServerClips, + previewCount, + setPreviewCount, + previewTitles, + setPreviewTitles, + voiceModePerVideo, + setVoiceModePerVideo, + voiceLibraryIds, + setVoiceLibraryIds, + previewCovers, + setPreviewCovers, + selectedVariantIds, + setSelectedVariantIds, } = formState + const isBatch = previewCount > 1 + + /* ── 配音选择同步:共用配音 ↔ 变体数组 ── */ + // 触发场景:①共用配音变化 ②批量模式进入/退出 ③独立→共用切换(需把所有变体刷成共用配音) + // 独立模式下:仅同步变体[0](其选择器绑定共用配音),用户单独选择的其他变体不覆盖 + const prevVoiceSyncRef = useRef({ + voice: selectedVoice, + batch: isBatch, + perVideo: voiceModePerVideo, + }) + useEffect(() => { + const prev = prevVoiceSyncRef.current + const voiceChanged = prev.voice !== selectedVoice + const modeChanged = prev.batch !== isBatch || prev.perVideo !== voiceModePerVideo + prevVoiceSyncRef.current = { voice: selectedVoice, batch: isBatch, perVideo: voiceModePerVideo } + if (!voiceChanged && !modeChanged) return + if (!isBatch) return + if (!voiceModePerVideo) { + // 共用模式(含刚从独立切回):所有变体跟随共用配音,未选择的补默认值 + setVoiceLibraryIds((prevIds) => (prevIds || []).map((id) => id || selectedVoice)) + } else if (voiceChanged) { + // 独立模式下共用配音变化:仅同步变体[0](与共用选择器绑定),其余不覆盖 + setVoiceLibraryIds((prevIds) => + (prevIds || []).map((id, i) => (i === 0 ? selectedVoice : id)), + ) + } + }, [selectedVoice, isBatch, voiceModePerVideo, setVoiceLibraryIds]) + + /* ── 数量选择弹窗 ── */ + const [countModalOpen, setCountModalOpen] = useState(false) + /* ── 标题样式回调 ── */ const styleUpdaters = useTitleStyleUpdaters({ titleSettings, @@ -88,6 +132,8 @@ const GeneratePage: React.FC = () => { const [previewVoiceAudioUrl, setPreviewVoiceAudioUrl] = useState(null) const ttsAbortRef = useRef(null) + // TTS 试听文案:批量跟随变体0标题(仅取首项,避免编辑其他变体标题触发多余 TTS 请求) + const variant0Title = isBatch ? previewTitles?.[0] || "" : "" useEffect(() => { const voiceAsset = voiceMaterials.find((m) => m.id === selectedVoice) @@ -96,8 +142,10 @@ const GeneratePage: React.FC = () => { return } + // 批量模式下 TTS 文案跟随变体0标题;单视频跟随主标题 + const ttsTitle = isBatch ? variant0Title || "" : titleSettings.title const voiceId = selectedClonedVoice || selectedVoice - if (!voiceId || !titleSettings.title) { + if (!voiceId || !ttsTitle) { setPreviewVoiceAudioUrl(null) return } @@ -107,7 +155,7 @@ const GeneratePage: React.FC = () => { ttsAbortRef.current = controller let cancelled = false - previewTts({ text: titleSettings.title, voice_id: voiceId }) + previewTts({ text: ttsTitle, voice_id: voiceId }) .then((res) => { if (!cancelled && res.audio_url) { setPreviewVoiceAudioUrl(res.audio_url) @@ -124,10 +172,17 @@ const GeneratePage: React.FC = () => { cancelled = true controller.abort() } - }, [selectedVoice, selectedClonedVoice, titleSettings.title, voiceMaterials]) + }, [ + selectedVoice, + selectedClonedVoice, + titleSettings.title, + variant0Title, + isBatch, + voiceMaterials, + ]) /* ── 克隆声音 ── */ - const { clones: clonedVoices, addClone, hasProcessing } = useCloneProgress() + const { addClone } = useCloneProgress() const handleCloneSuccess = (voice: VoiceClone) => { addClone(voice) @@ -156,17 +211,25 @@ const GeneratePage: React.FC = () => { [bgm, currentTemplate], ) - /* ── 加载素材详情(供前端预览播放器使用 + 配音时长校验) ── */ + /* ── 加载素材详情(供前端预览播放器使用) ── */ const previewAssetsEnabled = previewAssetIds.length > 0 const { assets: previewAssets, ready: previewAssetsReady } = usePreviewAssets( previewAssetIds, previewAssetsEnabled, ) - /* ── 预览就绪:素材已加载,且有模板 ── */ - const previewReady = useMemo( - () => previewAssetsReady && !!currentTemplate, - [previewAssetsReady, currentTemplate], + /* ── 预览就绪:纯前端 Canvas 预览,素材详情加载完即可秒开(单视频/批量一致) ── */ + const previewReady = previewAssetsReady && !!currentTemplate + + /* ── 勾选变体 ── */ + const toggleVariantSelect = useCallback( + (index: number) => { + setSelectedVariantIds((prev) => { + const list = prev || [] + return list.includes(index) ? list.filter((i) => i !== index) : [...list, index].sort() + }) + }, + [setSelectedVariantIds], ) /* ── 视频生成核心逻辑 ── */ @@ -176,8 +239,10 @@ const GeneratePage: React.FC = () => { generated, generateError, generatedVideos, + batchTasks, generate: handleGenerate, retry: handleRetryGenerate, + retryBatchTask: handleRetryBatchTask, dismissError: handleDismissError, download: handleDownload, share: handleShare, @@ -199,27 +264,92 @@ const GeneratePage: React.FC = () => { sourceEditPlanId: storedSourceEditPlanId || sourceEditPlanId, previewTaskId, bgmConfig, + previewCount, + variantTitles: previewTitles, + variantVoiceLibraryIds: voiceLibraryIds, + voiceModePerVideo, + variantCoverUrls: previewCovers, + selectedVariantIndexes: isBatch ? selectedVariantIds : undefined, onGenerationSuccess: () => { setPreviewTaskId(null) setStoredSourceEditPlanId(null) }, }) - /* ── 步骤4「确认生成视频」:校验标题/预览 → 创建最终渲染任务 → 成功后进入步骤5 ── */ + /* ── 数量弹窗确认:设置数量 + 同步批量数组长度 + 进入步骤2 ── */ + const handleCountConfirm = useCallback( + (count: number) => { + setPreviewCount(count) + setCountModalOpen(false) + // 同步批量数组长度 + setPreviewTitles((prev) => { + const list = prev || [] + const base = list[0] || titleSettings.title || "" + return Array.from({ length: count }, (_, i) => list[i] ?? (i === 0 ? base : "")) + }) + setVoiceLibraryIds((prev) => { + const list = prev || [] + return Array.from({ length: count }, (_, i) => list[i] ?? selectedVoice ?? "") + }) + setPreviewCovers((prev) => { + const list = prev || [] + return Array.from({ length: count }, (_, i) => list[i] ?? "") + }) + setSelectedVariantIds(Array.from({ length: count }, (_, i) => i)) + setCurrentStep(2) + }, + [ + setPreviewCount, + setPreviewTitles, + setVoiceLibraryIds, + setPreviewCovers, + setSelectedVariantIds, + setCurrentStep, + titleSettings.title, + selectedVoice, + ], + ) + + /* ── 步骤4「确认生成视频」:校验通过 → 创建正式生成任务 → 跳步骤5看实时进展 ── */ const handleConfirmGenerate = useCallback(async () => { - if (!titleSettings.title.trim()) { - message.warning("请选择或输入标题") - return + // 标题校验:批量只校验已勾选的变体;单视频校验主标题 + if (isBatch) { + if (selectedVariantIds.length === 0) { + message.warning("请至少勾选一个视频") + return + } + const missing = selectedVariantIds.some((i) => !previewTitles[i]?.trim()) + if (missing) { + message.warning("请为每个勾选的视频输入标题") + return + } + } else if (!titleSettings.title?.trim()) { + // 与 buildPayload.validateGenerateInputs 一致:AI 自动选标题模式(aiAutoSelect) + // 允许空标题由后端生成;手动模式必须填写,避免提交空标题 + if (!titleSettings.aiAutoSelect) { + message.warning("请先选择或输入标题") + return + } } if (!previewReady) { - message.warning("预览视频正在加载,请稍候") + message.warning("预览素材正在加载,请稍候") return } const ok = await handleGenerate() if (ok) { + // 单视频与批量一致:任务创建成功后进入步骤5「确认生成」看实时渲染进展 setCurrentStep(5) } - }, [titleSettings.title, previewReady, handleGenerate, setCurrentStep]) + }, [ + isBatch, + selectedVariantIds, + previewTitles, + titleSettings.aiAutoSelect, + titleSettings.title, + previewReady, + handleGenerate, + setCurrentStep, + ]) /* ── 步骤导航 ── */ const { goNext, goPrev } = useStepNavigation({ @@ -230,13 +360,21 @@ const GeneratePage: React.FC = () => { selectedMaterials, smartSelectedIds, titleSettings, - previewReady, generated, + onOpenCountModal: () => setCountModalOpen(true), }) - /* ── 最终成片(步骤5/6 右侧播放) ── */ + /* ── 最终成片(单视频右侧播放) ── */ const finalVideo = generatedVideos[0] + /* ── 布局 class:步骤4标题页=预览+标题侧栏;步骤5/6批量=整行宽;步骤1~3=整行宽 ── */ + const layoutClassName = useMemo(() => { + if (currentStep < 4) return "xx-generate-layout full-width" + if (currentStep === 4) return "xx-generate-layout step4-layout" + // 步骤5/6:批量网格需要整行宽度;单视频保持 表单+右侧成片 两栏 + return isBatch ? "xx-generate-layout full-width" : "xx-generate-layout" + }, [currentStep, isBatch]) + /* ================================================================ 渲染 ================================================================ */ @@ -247,8 +385,61 @@ const GeneratePage: React.FC = () => { -
- {/* ════ 左侧:表单区 ════ */} +
+ {/* ════ 步骤4:左侧预览大区域(纯前端 Canvas 实时预览) ════ */} + {currentStep === 4 && !!currentTemplate && ( +
+ {!isBatch ? ( + /* 单视频:前端 Canvas 实时预览(与旧版一致,零回归) */ + 0} + serverClips={serverClips} + voiceAudioUrl={previewVoiceAudioUrl || undefined} + titleSettings={{ + title: titleSettings.title, + size: titleSettings.size, + font: titleSettings.font, + color: titleSettings.color, + position: titleSettings.position as "top" | "center" | "bottom" | "custom", + bold: titleSettings.bold, + italic: titleSettings.italic, + stroke: titleSettings.stroke, + shadow: titleSettings.shadow, + posX: titleSettings.posX, + posY: titleSettings.posY, + }} + onTitlePositionChange={styleUpdaters.updateTitlePosition} + /> + ) : ( + /* 批量:N 个前端 Canvas 预览网格(不调任何后端渲染接口,秒开) */ +
+
+

🎬 {previewCount} 个视频预览

+ + 实时预览,勾选要生成的视频 + +
+ +
+ )} +
+ )} + + {/* ════ 右侧:步骤1~3 表单 / 步骤4 标题边栏 / 步骤5 确认生成进度 / 步骤6 封面 ════ */}
{ selectedVoice={selectedVoice} onSelectedVoiceChange={setSelectedVoice} onServerClipsChange={setServerClips} - voiceMode={voiceMode} - onVoiceModeChange={setVoiceMode} - selectedClonedVoice={selectedClonedVoice} - onSelectedClonedVoiceChange={setSelectedClonedVoice} - clonedVoices={clonedVoices} - addClone={addClone} - hasProcessing={hasProcessing} - cloneModalOpen={cloneModalOpen} - onCloneModalOpenChange={setCloneModalOpen} generating={generating} generated={generated} generateError={generateError} progress={progress} generatedVideos={generatedVideos} onRetry={handleRetryGenerate} + onRetryBatchTask={handleRetryBatchTask} onDismissError={handleDismissError} - presetVoices={presetVoices} + batchTasks={batchTasks} + previewCount={previewCount} + previewTitles={previewTitles} + onPreviewTitlesChange={setPreviewTitles} + voiceModePerVideo={voiceModePerVideo} + onVoiceModePerVideoChange={setVoiceModePerVideo} + voiceLibraryIds={voiceLibraryIds} + onVoiceLibraryIdsChange={setVoiceLibraryIds} + previewCovers={previewCovers} + onPreviewCoversChange={setPreviewCovers} + selectedVariantIds={selectedVariantIds} /> { generating={generating} generated={generated} generateError={generateError} + selectedCount={isBatch ? selectedVariantIds.length : 1} />
- {/* ════ 右侧:步骤4实时预览,步骤5/6最终视频 ════ */} -
- {currentStep === 4 && !!currentTemplate && ( - 0} - serverClips={serverClips} - voiceAudioUrl={previewVoiceAudioUrl || undefined} - titleSettings={{ - title: titleSettings.title, - size: titleSettings.size, - font: titleSettings.font, - color: titleSettings.color, - position: titleSettings.position as "top" | "center" | "bottom" | "custom", - bold: titleSettings.bold, - italic: titleSettings.italic, - stroke: titleSettings.stroke, - shadow: titleSettings.shadow, - posX: titleSettings.posX, - posY: titleSettings.posY, - }} - onTitlePositionChange={styleUpdaters.updateTitlePosition} - /> - )} - {currentStep >= 5 && generated && finalVideo && ( + {/* ════ 步骤5/6(单视频):右侧成片播放器 ════ */} + {currentStep >= 5 && !isBatch && generated && finalVideo && ( +
- )} -
+
+ )}
+ {/* 数量选择弹窗 */} + setCountModalOpen(false)} + /> + {/* 音色克隆弹窗 */} void +} + +const BatchGenerationGrid: React.FC = ({ + tasks, + titles, + onRetryTask, +}) => { + const sorted = [...tasks].sort( + (a, b) => Number(a.variantIndex || 0) - Number(b.variantIndex || 0), + ) + + return ( +
+
+

🎬 正在生成 {tasks.length} 个视频

+ + 完成 {tasks.filter((t) => t.status === "completed").length} / {tasks.length} + +
+
+ {sorted.map((task) => { + const title = titles[task.variantIndex] || `视频 ${task.variantIndex + 1}` + const video = (task.videos?.[0] || null) as GeneratedVideo | null + return ( +
+
+ + {task.status === "completed" ? ( + + ) : task.status === "failed" ? ( + + ) : ( + + )} + 视频 {task.variantIndex + 1}:{title} + +
+ +
+ {task.status === "running" && ( + <> +
+
+
+
{Math.round(task.progress)}%
+ + )} + {task.status === "completed" && video && ( +
+
+ ) + })} +
+
+ ) +} + +export default BatchGenerationGrid diff --git a/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx b/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx new file mode 100644 index 000000000..4dd74d0f3 --- /dev/null +++ b/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx @@ -0,0 +1,97 @@ +/** + * 批量前端 Canvas 实时预览网格(Issue #1677 修正方案) + * + * N 个 FrontendPreviewPlayer 网格排列: + * - 纯前端 Canvas + video 元素实时播放素材片段,不调任何后端渲染接口 + * - variantSeed 让每个变体素材排布/起始点不同,画面有可见差异 + * - 各自叠加独立标题浮层(variantTitle),标题样式全局共用 + * - 勾选框决定提交时生成哪些变体 + */ +import React from "react" +import type { AssetItem } from "@/api/assets" +import type { EditingTemplate } from "@/api/editing-planner" +import type { TitleSettings } from "../types" +import FrontendPreviewPlayer from "./FrontendPreviewPlayer" + +interface CanvasPreviewGridProps { + count: number + assets: AssetItem[] + template: EditingTemplate | null + videoRatio: string + titles: string[] + titleSettings: TitleSettings + /** 共用配音预览音频(仅第 1 个变体播放,避免多路音频重叠) */ + voiceAudioUrl?: string + /** 勾选的变体序号 */ + selectedIds: number[] + onToggleSelect: (index: number) => void + /** 生成中禁止勾选 */ + selectable?: boolean +} + +const CanvasPreviewGrid: React.FC = ({ + count, + assets, + template, + videoRatio, + titles, + titleSettings, + voiceAudioUrl, + selectedIds, + onToggleSelect, + selectable = true, +}) => { + // count 上限已在源头 PreviewCountModal 的数量选择(1~MAX_PREVIEW_COUNT=10)clamp, + // 这里完整渲染所有变体,保证每个变体都有勾选/预览入口,UI 与数据不脱节 + return ( +
+ {Array.from({ length: count }, (_, i) => { + const checked = selectedIds.includes(i) + return ( +
+
+ +
+ 0} + variantSeed={i + 1} + variantTitle={titles[i] || ""} + voiceAudioUrl={i === 0 ? voiceAudioUrl : undefined} + compact + titleSettings={{ + title: titles[i] || "", + size: titleSettings.size, + font: titleSettings.font, + color: titleSettings.color, + position: titleSettings.position as "top" | "center" | "bottom" | "custom", + bold: titleSettings.bold, + italic: titleSettings.italic, + stroke: titleSettings.stroke, + shadow: titleSettings.shadow, + posX: titleSettings.posX, + posY: titleSettings.posY, + }} + /> +
+ ) + })} +
+ ) +} + +export default CanvasPreviewGrid diff --git a/apps/web/src/pages/generate/components/FrontendPreviewPlayer.tsx b/apps/web/src/pages/generate/components/FrontendPreviewPlayer.tsx index 35387d1fe..b9b721b70 100644 --- a/apps/web/src/pages/generate/components/FrontendPreviewPlayer.tsx +++ b/apps/web/src/pages/generate/components/FrontendPreviewPlayer.tsx @@ -41,6 +41,16 @@ interface FrontendPreviewPlayerProps { posY?: number | null } onTitlePositionChange?: (posX: number, posY: number) => void + /** + * 变体种子(批量生成 #1677):同一批素材在不同变体中采用不同的素材顺序与 + * 片段起始点,让 N 个 Canvas 预览画面有差异(纯前端随机剪辑模拟,不调后端)。 + * 0 / 不传 = 单视频,排布与旧版完全一致(零回归)。 + */ + variantSeed?: number + /** 变体标题文字(批量时每个预览独立标题,叠加在画面上);不传用 titleSettings.title */ + variantTitle?: string + /** 紧凑模式(批量网格中使用,缩小内边距/标题尺寸) */ + compact?: boolean } function formatTime(seconds: number): string { @@ -52,10 +62,23 @@ function formatTime(seconds: number): string { /** * 将素材映射为播放片段(复用原逻辑) */ +/** 简单可复现随机数(mulberry32),同一种子产出稳定排布,避免每次渲染抖动 */ +function seededRandom(seed: number): () => number { + let a = seed >>> 0 + return () => { + a |= 0 + a = (a + 0x6d2b79f5) | 0 + let t = Math.imul(a ^ (a >>> 15), 1 | a) + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t + return ((t ^ (t >>> 14)) >>> 0) / 4294967296 + } +} + function buildPlaybackSegments( assets: AssetItem[], template: EditingTemplate | null, serverClips?: EditPlanClip[], + variantSeed = 0, ): PlaybackSegment[] { if (!assets.length) return [] @@ -79,18 +102,40 @@ function buildPlaybackSegments( } } - // Fallback: 本地构建片段(与旧行为一致) + // Fallback: 本地构建片段 + // variantSeed=0(单视频):与旧行为完全一致(素材原序、起始点 0),零回归 + // variantSeed>0(批量变体):素材顺序按种子轮换 + 片段起始点在素材内偏移, + // 模拟后端"AI 随机剪辑出不同版本",让 N 个预览画面有可见差异 const templateSegments = template?.segments || [] const segments: PlaybackSegment[] = [] + const orderedAssets = variantSeed > 0 ? [...assets] : assets + if (variantSeed > 0 && orderedAssets.length > 1) { + const rand = seededRandom(variantSeed * 7919 + 13) + // 素材轮换:把数组旋转 (seed % n) 位,再对后半段做一次稳定交换 + const n = orderedAssets.length + const rotate = variantSeed % n + orderedAssets.push(...orderedAssets.splice(0, rotate)) + const swapA = Math.floor(rand() * n) + const swapB = Math.floor(rand() * n) + if (swapA !== swapB) { + ;[orderedAssets[swapA], orderedAssets[swapB]] = [orderedAssets[swapB], orderedAssets[swapA]] + } + } - assets.forEach((asset, i) => { + orderedAssets.forEach((asset, i) => { const assetDuration = asset.duration || asset.metadata?.duration || 30 const tplSeg = templateSegments[i] || templateSegments[templateSegments.length - 1] const segDuration = tplSeg ? Math.min(tplSeg.duration_max, Math.max(tplSeg.duration_min, assetDuration)) : Math.min(assetDuration, 10) - const startTime = 0 + let startTime = 0 + if (variantSeed > 0 && assetDuration - segDuration > 1) { + const rand = seededRandom(variantSeed * 104729 + i * 31 + 7) + // 起始点在素材可用区间内随机偏移(至少留 0.5s 余量) + const maxStart = Math.max(0, assetDuration - segDuration - 0.5) + startTime = Math.round(rand() * maxStart * 10) / 10 + } const endTime = Math.min(startTime + segDuration, assetDuration) const videoUrl = asset.file_url || asset.storage_key @@ -109,11 +154,16 @@ const FrontendPreviewPlayer: React.FC = ({ voiceAudioUrl, titleSettings, onTitlePositionChange, + variantSeed = 0, + variantTitle, + compact = false, }) => { const segments = useMemo( - () => buildPlaybackSegments(assets, template, serverClips), - [assets, template, serverClips], + () => buildPlaybackSegments(assets, template, serverClips, variantSeed), + [assets, template, serverClips, variantSeed], ) + // 批量变体:标题文字取 variantTitle,样式仍由全局 titleSettings 控制 + const effectiveTitle = variantTitle ?? titleSettings?.title // ── ASS 坐标系参数(与后端 ass_subtitle_builder.py 一致) ── const TITLE_MARGIN_TOP = 120 @@ -234,7 +284,7 @@ const FrontendPreviewPlayer: React.FC = ({ // ── Canvas 播放器(WebCodecs 路径) ── const canvasTitle = titleSettings ? { - text: titleSettings.title || "标题预览", + text: effectiveTitle || "标题预览", fontSize: titleSettings.size, fontFamily: titleSettings.font || "思源黑体", color: titleSettings.color || "#ffffff", @@ -520,13 +570,15 @@ const FrontendPreviewPlayer: React.FC = ({ style={{ position: "relative", width: "100%", - maxWidth: 280, + maxWidth: compact ? "100%" : 280, + margin: compact ? 0 : "0 auto", aspectRatio: "9 / 16", background: "#0a0a0a", - borderRadius: 24, + borderRadius: compact ? 10 : 24, overflow: "hidden", - boxShadow: - "0 4px 6px -1px rgba(0,0,0,0.3), 0 20px 50px -12px rgba(0,0,0,0.5), inset 0 0 0 1px rgba(255,255,255,0.06)", + boxShadow: compact + ? "inset 0 0 0 1px rgba(255,255,255,0.06)" + : "0 4px 6px -1px rgba(0,0,0,0.3), 0 20px 50px -12px rgba(0,0,0,0.5), inset 0 0 0 1px rgba(255,255,255,0.06)", }} > {/* ── Canvas 渲染层(WebCodecs 路径) ── */} @@ -610,8 +662,8 @@ const FrontendPreviewPlayer: React.FC = ({ ? { top: "50%", transform: "translate(-50%, -50%)" } : { bottom: `${titleBottomPct}%` }), }), - pointerEvents: "auto", - cursor: onTitlePositionChange ? "grab" : "default", + pointerEvents: onTitlePositionChange && variantSeed === 0 ? "auto" : "none", + cursor: onTitlePositionChange && variantSeed === 0 ? "grab" : "default", touchAction: "none", userSelect: "none", WebkitUserSelect: "none", @@ -641,7 +693,7 @@ const FrontendPreviewPlayer: React.FC = ({ : undefined, }} > - {titleSettings.title.split(/[//]/).map((part, i) => ( + {(effectiveTitle || "").split(/[//]/).map((part, i) => ( {i > 0 &&
} {part} diff --git a/apps/web/src/pages/generate/components/GenerateStepActions.tsx b/apps/web/src/pages/generate/components/GenerateStepActions.tsx index 87f5d49fc..c25e75abe 100644 --- a/apps/web/src/pages/generate/components/GenerateStepActions.tsx +++ b/apps/web/src/pages/generate/components/GenerateStepActions.tsx @@ -1,9 +1,9 @@ /** - * GeneratePage 步骤底部操作按钮 + * GeneratePage 步骤底部操作按钮(Issue #1677 修正:固定 6 步) * * 步骤 1~3:上一步 / 下一步 - * 步骤 4(标题+预览):上一步 / 确认生成视频(点击后直接创建最终渲染任务,成功后跳转步骤5) - * 步骤 5(确认生成):上一步 / 下一步(渲染中禁用,渲染完成后可进入封面) + * 步骤 4(选择标题):「✨ 确认生成视频 / 确认生成 N 个视频」→ 创建正式生成任务,成功后跳步骤5 + * 步骤 5(确认生成):渲染进度页,全部完成后「下一步:选择封面」;仅上一步 * 步骤 6(选择封面):仅上一步 */ import React from "react" @@ -12,14 +12,16 @@ export interface GenerateStepActionsProps { currentStep: number onPrev: () => void onNext: () => void - /** 步骤4:确认生成视频(校验 + 创建渲染任务 + 成功后进入步骤5) */ + /** 步骤4:确认生成视频(校验 + 创建渲染任务) */ onConfirmGenerate: () => void | Promise generating: boolean generated: boolean generateError: string | null + /** 批量模式下勾选的视频数量(N=1 时为1) */ + selectedCount?: number } -export const GenerateStepActions: React.FC = ({ +const GenerateStepActions: React.FC = ({ currentStep, onPrev, onNext, @@ -27,9 +29,10 @@ export const GenerateStepActions: React.FC = ({ generating, generated, generateError, + selectedCount = 1, }) => { const renderPrimaryButton = () => { - /* 步骤 1~3:上一步 / 下一步(必填校验由 useStepNavigation.goNext 统一处理) */ + /* 步骤 1~3:上一步 / 下一步 */ if (currentStep < 4) { return ( ) } if (generateError) { return ( - ) - } - if (generated) { - return ( - ) } return ( ) } - /* 步骤 5:渲染中禁用,完成后下一步进入封面 */ + /* 步骤 5:确认生成进度页 — 全部完成后下一步进封面 */ if (currentStep === 5) { + if (generated) { + return ( + + ) + } return ( - ) } - /* 步骤 6(最后一步):无主按钮 */ + /* 步骤 6(封面,最后一步):无主按钮 */ return null } diff --git a/apps/web/src/pages/generate/components/GenerateStepContent.tsx b/apps/web/src/pages/generate/components/GenerateStepContent.tsx index 1b160b5ec..639c99aa5 100644 --- a/apps/web/src/pages/generate/components/GenerateStepContent.tsx +++ b/apps/web/src/pages/generate/components/GenerateStepContent.tsx @@ -1,20 +1,20 @@ /** * GeneratePage 步骤内容渲染 - * 步骤顺序(6步):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6) + * 步骤顺序(6步,Issue #1677 修正):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6) + * 步骤4预览(Canvas 网格)与步骤5进度(批量渲染网格)由 GeneratePage 直接渲染在左侧大区域。 */ import React from "react" import type { EditingTemplate } from "@/api/editing-planner" import type { EditPlanClip } from "@/api/template-editor" -import type { PresetVoiceItem } from "@/api/voices" -import type { VoiceClone } from "@/api/voice-clone" import type { CoverConfig } from "../types/cover" import type { TitleSettings } from "../types" import Step1TemplateSelect from "../components/Step1TemplateSelect" import Step2MaterialSelect from "../components/Step2MaterialSelect" -import Step3VoiceSelect from "../components/Step5VoiceSelect" +import Step3VoiceWithMode from "./Step3VoiceWithMode" import Step4TitleSettings from "../components/Step4TitleSettings" -import Step5ConfirmGenerate from "../components/Step7ConfirmGenerate" import Step6CoverSettings from "../components/Step6CoverSettings" +import BatchGenerationGrid from "./BatchGenerationGrid" +import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling" import type { GeneratedVideo } from "@/api/template-editor" export interface GenerateStepContentProps { @@ -50,15 +50,6 @@ export interface GenerateStepContentProps { selectedVoice: string onSelectedVoiceChange: (id: string) => void onServerClipsChange: (clips: EditPlanClip[]) => void - voiceMode: "preset" | "custom" | "clone" - onVoiceModeChange: (mode: "preset" | "custom" | "clone") => void - selectedClonedVoice: string - onSelectedClonedVoiceChange: (id: string) => void - clonedVoices: VoiceClone[] - addClone: (voice: VoiceClone) => void - hasProcessing: boolean - cloneModalOpen: boolean - onCloneModalOpenChange: (open: boolean) => void /* 生成 */ generating: boolean generated: boolean @@ -66,13 +57,26 @@ export interface GenerateStepContentProps { progress: number generatedVideos: GeneratedVideo[] onRetry: () => void + onRetryBatchTask: (taskId: string) => void onDismissError: () => void - /* 其他 */ - presetVoices: PresetVoiceItem[] + /** 批量:每个正式生成任务的独立状态(步骤5进度网格) */ + batchTasks: BatchTaskState[] /** BGM 开关 */ bgm: boolean /** BGM 配置(来自模板) */ bgmConfig?: { enabled: boolean; music_id?: string } + /* ── 批量生成(#1677)── */ + previewCount: number + previewTitles: string[] + onPreviewTitlesChange: (titles: string[]) => void + voiceModePerVideo: boolean + onVoiceModePerVideoChange: (v: boolean) => void + voiceLibraryIds: string[] + onVoiceLibraryIdsChange: (ids: string[]) => void + previewCovers: string[] + onPreviewCoversChange: (urls: string[]) => void + /** 批量模式勾选的变体索引(封面卡片按勾选顺序展示) */ + selectedVariantIds?: number[] } export const GenerateStepContent: React.FC = (props) => { @@ -104,17 +108,24 @@ export const GenerateStepContent: React.FC = (props) = selectedVoice, onSelectedVoiceChange, onServerClipsChange, - voiceMode, - selectedClonedVoice, - clonedVoices, generating, generated, generateError, progress, - generatedVideos, onRetry, - onDismissError, - presetVoices, + generatedVideos, + batchTasks, + onRetryBatchTask, + previewCount, + previewTitles, + onPreviewTitlesChange, + voiceModePerVideo, + onVoiceModePerVideoChange, + voiceLibraryIds, + onVoiceLibraryIdsChange, + previewCovers, + onPreviewCoversChange, + selectedVariantIds, } = props /* 当前模板的 segments,传给 Step2 构建 clips */ @@ -146,9 +157,14 @@ export const GenerateStepContent: React.FC = (props) = ) case 3: return ( - ) case 4: @@ -167,31 +183,66 @@ export const GenerateStepContent: React.FC = (props) = onApplyPreset={onApplyPreset} activePreset={activePreset} titlePresets={titlePresets} + previewCount={previewCount} + previewTitles={previewTitles} + onPreviewTitlesChange={onPreviewTitlesChange} /> ) case 5: + /* 确认生成页:批量=逐任务进度网格;单视频=进度状态卡(成片播放器在左侧大区域) */ + if (previewCount > 1) { + return ( + + ) + } + /* 单视频:渲染进度 / 失败重试 / 完成提示(成片播放器在右侧栏) */ return ( - +
+

🎬 确认生成

+ {generating && ( +
+
+
+
+ ⏳ 视频渲染中… {Math.round(progress)}% +
+
+ 生成过程中可以切换到其他页面,完成后可在任务历史查看 +
+
+
+
+
+
+
+ )} + {generateError && !generating && ( +
+
+
生成失败
+
{generateError}
+
+ +
+ )} + {generated && !generating && ( +
+
+
✅ 视频生成完成!
+
右侧可预览成片,点击「下一步」选择封面
+
+
+ )} +
) case 6: return ( @@ -201,6 +252,11 @@ export const GenerateStepContent: React.FC = (props) = selectedTemplate={selectedTemplate} titleSettings={titleSettings} generatedVideos={generatedVideos} + previewCount={previewCount} + previewTitles={previewTitles} + previewCovers={previewCovers} + onPreviewCoversChange={onPreviewCoversChange} + selectedVariantIndexes={selectedVariantIds} /> ) default: diff --git a/apps/web/src/pages/generate/components/PreviewCountModal.tsx b/apps/web/src/pages/generate/components/PreviewCountModal.tsx new file mode 100644 index 000000000..f8a334bf7 --- /dev/null +++ b/apps/web/src/pages/generate/components/PreviewCountModal.tsx @@ -0,0 +1,127 @@ +/** + * 生成数量选择弹窗(Issue #1677) + * Step1 选完模板点「下一步」时弹出:要生成几个视频?(1~10) + * 默认 1,回车 = 1(零额外操作) + */ +import React, { useState, useEffect, useRef } from "react" +import { MAX_PREVIEW_COUNT } from "../constants" + +interface PreviewCountModalProps { + open: boolean + /** 默认值(上次选择,默认1) */ + defaultCount?: number + onConfirm: (count: number) => void + onCancel: () => void +} + +const PreviewCountModal: React.FC = ({ + open, + defaultCount = 1, + onConfirm, + onCancel, +}) => { + const [count, setCount] = useState(defaultCount) + const inputRef = useRef(null) + + useEffect(() => { + if (open) { + setCount(defaultCount) + // 弹窗打开后聚焦并选中,方便直接回车=默认1 + setTimeout(() => inputRef.current?.focus(), 50) + } + }, [open, defaultCount]) + + const clamp = (n: number) => Math.max(1, Math.min(MAX_PREVIEW_COUNT, n || 1)) + + const handleConfirm = () => { + onConfirm(clamp(count)) + } + + const handleKeyDown = (e: React.KeyboardEvent) => { + if (e.key === "Enter") { + e.preventDefault() + handleConfirm() + } + if (e.key === "Escape") { + onCancel() + } + } + + if (!open) return null + + return ( +
+
e.stopPropagation()}> +

要生成几个视频?

+

+ 素材共用,AI 随机剪辑出不同版本,每个视频可独立设置标题、配音和封面 +

+ +
+ + setCount(clamp(parseInt(e.target.value, 10) || 1))} + onKeyDown={handleKeyDown} + className="xx-count-input" + /> + +
+ +
+ {[1, 3, 5, 10].map((n) => ( + + ))} +
+ +
+ + +
+

+ 直接按回车 = 生成 1 个 +

+
+
+ ) +} + +export default PreviewCountModal diff --git a/apps/web/src/pages/generate/components/Step3VoiceWithMode.tsx b/apps/web/src/pages/generate/components/Step3VoiceWithMode.tsx new file mode 100644 index 000000000..ce832ae1d --- /dev/null +++ b/apps/web/src/pages/generate/components/Step3VoiceWithMode.tsx @@ -0,0 +1,102 @@ +/** + * Step3 配音选择(Issue #1677 批量生成) + * - 单视频 / 共用模式:与原配音选择完全一致 + * - 独立模式(开关开启):N 个配音选择器,每个视频独立选择 + */ +import React from "react" +import Step3VoiceSelect from "./Step5VoiceSelect" + +interface Step3VoiceWithModeProps { + previewCount: number + /** 共用配音ID */ + selectedVoice: string + onSelectedVoiceChange: (id: string) => void + /** 是否独立配音 */ + voiceModePerVideo: boolean + onVoiceModePerVideoChange: (v: boolean) => void + /** 各变体独立配音ID */ + voiceLibraryIds: string[] + onVoiceLibraryIdsChange: (ids: string[]) => void +} + +const Step3VoiceWithMode: React.FC = ({ + previewCount, + selectedVoice, + onSelectedVoiceChange, + voiceModePerVideo, + onVoiceModePerVideoChange, + voiceLibraryIds, + onVoiceLibraryIdsChange, +}) => { + const isBatch = previewCount > 1 + + if (!isBatch) { + return ( + + ) + } + + return ( +
+ {/* 共用/独立切换 */} +
+
+
+ 🎙️ 配音方式:{voiceModePerVideo ? "每个视频独立配音" : "所有视频共用配音"} +
+
+ {voiceModePerVideo + ? `为 ${previewCount} 个视频分别选择不同配音` + : "所有视频使用同一个配音(默认)"} +
+
+
onVoiceModePerVideoChange(!voiceModePerVideo)} + role="switch" + aria-checked={voiceModePerVideo} + tabIndex={0} + onKeyDown={(e) => { + if (e.key === "Enter" || e.key === " ") { + e.preventDefault() + onVoiceModePerVideoChange(!voiceModePerVideo) + } + }} + > +
+
+
+ + {!voiceModePerVideo ? ( + + ) : ( +
+ {Array.from({ length: previewCount }, (_, i) => ( + { + const next = [...voiceLibraryIds] + next[i] = id + onVoiceLibraryIdsChange(next) + }} + /> + ))} +
+ )} +
+ ) +} + +export default Step3VoiceWithMode diff --git a/apps/web/src/pages/generate/components/Step4TitleSettings.tsx b/apps/web/src/pages/generate/components/Step4TitleSettings.tsx index 8a457b52b..6f8500519 100644 --- a/apps/web/src/pages/generate/components/Step4TitleSettings.tsx +++ b/apps/web/src/pages/generate/components/Step4TitleSettings.tsx @@ -1,17 +1,23 @@ /** - * Step 4 选择标题(合并原 Step4 标题输入 + Step5 标题样式面板) + * Step 4 选择标题(Issue #1677 批量生成) * - * 左侧:标题文字输入 + AI生成标题 + 样式设置(位置/字号/字体/颜色/样式/预设) - * 右侧:FrontendPreviewPlayer 实时预览(由 GeneratePage 统一渲染) + * 布局(由 GeneratePage 编排):左侧大区域实时预览(单=大播放器,批量=Canvas 网格), + * 右侧边栏标题设置。本组件渲染在右侧边栏: + * - 单视频:AI 标题生成器 + AutoComplete 标题库(与旧版完全一致,零回归) + * - 批量:N 个独立标题输入框(AutoComplete 支持标题库选择)+ 批量 AI 生成 + * (一次生成 N 个标题,分别填入各变体,可单独换一个) + * - 标题样式(字体/颜色/位置/大小/粗斜描边/预设):全局统一 */ -import React from "react" -import { AutoComplete } from "antd" -import { PlayCircleOutlined } from "@ant-design/icons" +import React, { useMemo, useState } from "react" +import { Input, message } from "antd" +import { LoadingOutlined } from "@ant-design/icons" import type { TitleSettings } from "../types" import { POSITION_OPTIONS, FONT_OPTIONS } from "../constants" import { useStep4Title } from "../hooks/useStep4Title" import AiTitleGenerator from "./title/AiTitleGenerator" +import TitleLibraryAutoComplete from "./title/TitleLibraryAutoComplete" import TitleStylePanel from "./title/TitleStylePanel" +import { AI_TITLE_TEMPLATES } from "../constants" interface Step4TitleSettingsProps { titleSettings: TitleSettings @@ -29,6 +35,42 @@ interface Step4TitleSettingsProps { onApplyPreset: (presetKey: string) => void activePreset: string | null titlePresets: { key: string; label: string; previewStyle: React.CSSProperties }[] + /* ── 批量生成(#1677)── */ + /** 生成数量 */ + previewCount?: number + /** 每个变体的标题文字(长度=previewCount) */ + previewTitles?: string[] + onPreviewTitlesChange?: (titles: string[]) => void +} + +/** 从本地 AI 标题模板池按主题词生成 N 个不同标题(与单视频 AI 生成同源) */ +function buildBatchAiTitles(topic: string, count: number): string[] { + const styles: Array<"catchy" | "emotional" | "informative"> = [ + "catchy", + "emotional", + "informative", + ] + const pool: string[] = [] + styles.forEach((style) => { + const templates = AI_TITLE_TEMPLATES[style] || [] + templates.forEach((tpl) => pool.push(tpl.replace(/\{topic\}/g, topic))) + }) + // 洗牌后取前 count 个;不足则轮转补齐 + const shuffled = [...pool].sort(() => Math.random() - 0.5) + const out: string[] = [] + for (let i = 0; i < count; i++) { + out.push(shuffled[i % shuffled.length] || "") + } + return out +} + +function extractTopic(text: string): string { + const keywords = text + .replace(/[,。!?、,.!?]/g, " ") + .split(/\s+/) + .filter(Boolean) + if (keywords.length === 0) return "这个话题" + return keywords.slice(0, 3).join("") } const Step4TitleSettings: React.FC = (props) => { @@ -44,122 +86,195 @@ const Step4TitleSettings: React.FC = (props) => { onApplyPreset, activePreset, titlePresets, + previewCount = 1, + previewTitles, + onPreviewTitlesChange, } = props + const isBatch = previewCount > 1 + const [batchAiLoading, setBatchAiLoading] = useState(false) + const [batchAiTopic, setBatchAiTopic] = useState("") + + /** 更新单个变体标题;变体0同步写回 titleSettings.title(全局样式面板/草稿/TTS 链路依赖) */ + const updateVariantTitle = (index: number, val: string) => { + if (!previewTitles || !onPreviewTitlesChange) return + const next = [...previewTitles] + next[index] = val + onPreviewTitlesChange(next) + if (index === 0) { + t.updateTitle(val) + } + } + + /** 批量 AI 生成:按主题词生成标题,分别填入 N 个变体 */ + const handleBatchAiGenerate = async (onlyEmpty = false) => { + if (!onPreviewTitlesChange || !previewTitles) return + const topic = (batchAiTopic || t.aiTitleInput || "").trim() + if (!topic) { + message.warning("请先输入主题词,例如:萌宠日常、旅行vlog") + return + } + setBatchAiLoading(true) + try { + // 与单视频一致:本地模板模拟 AI 生成(1200ms 体验延迟) + await new Promise((resolve) => setTimeout(resolve, 800)) + const picked = buildBatchAiTitles(extractTopic(topic), previewCount) + const next = [...previewTitles] + for (let i = 0; i < previewCount; i++) { + if (onlyEmpty && next[i]?.trim()) continue + if (picked[i]) next[i] = picked[i] + } + onPreviewTitlesChange(next) + if (next[0]) t.updateTitle(next[0]) + message.success(`已为 ${previewCount} 个视频生成标题,可单独修改`) + } finally { + setBatchAiLoading(false) + } + } + + const titleOptions = useMemo( + () => t.userTitles.map((ut) => ({ label: ut.content, value: ut.content })), + [t.userTitles], + ) + return ( -
+

📝 选择标题

- {/* AI 自动选择模式 */} - {t.titleSettings.aiAutoSelect && ( + {!isBatch ? ( + /* ── 单视频:原有 AI 标题 + 输入框(保持不变,零回归) ── */ <> -
- AI 自动选择标题 -
-
-
-
- - {/* 显示当前 AI 选中的标题(只读)+ 换一个按钮 */} -
- -
- {t.titleSettings.title || "AI 将自动为你选择标题"} - -
-
- - )} - - {/* 手动选择模式 */} - {!t.titleSettings.aiAutoSelect && ( - <> - - -
- AI 自动选择标题 -
-
-
-
- -
- - t.updateTitle(val || "")} - options={t.userTitles.map((ut) => ({ - label: ut.content, - value: ut.content, - }))} - filterOption={(inputValue, option) => { - const title = (option?.label || option?.value || "") as string - return title.toLowerCase().includes((inputValue || "").toLowerCase()) - }} - notFoundContent={ - t.userTitles.length === 0 ? ( - - 标题库为空,请前往「标题管理」添加 + {t.titleSettings.aiAutoSelect ? ( + <> +
+ AI 自动选择标题 +
+
+
+
+
+ +
+ + {(previewTitles?.[0] ?? t.titleSettings.title) || "AI 将自动为你选择标题"} - ) : null - } - /> -
+ +
+
+ + ) : ( + <> + + +
+ AI 自动选择标题 +
+
+
+
+
+ + { + t.updateTitle(val || "") + onPreviewTitlesChange?.([val || ""]) + }} + options={titleOptions} + /> +
+ + )} + ) : ( + /* ── 批量:AI 批量生成 + N 个独立标题输入框(AutoComplete 支持标题库) ── */ +
+
+ 为每个视频输入独立标题,修改会实时叠加到左侧对应视频上。标题样式(字体/颜色/位置)全局统一。 +
+ + {/* 批量 AI 标题 */} +
+ { + setBatchAiTopic(e.target.value) + t.setAiTitleInput(e.target.value) + }} + maxLength={30} + size="small" + style={{ flex: 1 }} + /> + + +
+ + {Array.from({ length: previewCount }, (_, i) => ( +
+ + updateVariantTitle(i, val)} + options={titleOptions} + /> +
+ ))} +
)} - {/* 标题样式面板(原 Step5) */} -
- - - 右侧为实时预览,调整样式即时生效 - -
- + {/* 标题样式面板(全局共用) */} void + /** 卡片标题(独立配音模式下显示"视频 N 的配音"),默认"选择配音" */ + heading?: string + /** 描述文案 */ + description?: string + /** 是否使用紧凑卡片样式(独立配音模式下 N 个并排) */ + compact?: boolean } /** 获取素材实际时长(优先顶层 duration,fallback 到 metadata.duration) */ @@ -43,6 +49,9 @@ const formatFileSize = (bytes?: number): string => { const Step5VoiceSelect: React.FC = ({ selectedVoice, onSelectedVoiceChange, + heading = "🎙️ 选择配音", + description = "从配音库中选择已上传的素材,点击卡片可预览播放", + compact = false, }) => { const navigate = useNavigate() const [playingId, setPlayingId] = useState(null) @@ -148,14 +157,14 @@ const Step5VoiceSelect: React.FC = ({ return (
-

🎙️ 选择配音

-

- 从配音库中选择已上传的素材,点击卡片可预览播放 -

+

{heading}

+

{description}

diff --git a/apps/web/src/pages/generate/components/Step6CoverSettings.tsx b/apps/web/src/pages/generate/components/Step6CoverSettings.tsx index d394a5f9a..41eaba9c0 100755 --- a/apps/web/src/pages/generate/components/Step6CoverSettings.tsx +++ b/apps/web/src/pages/generate/components/Step6CoverSettings.tsx @@ -1,9 +1,16 @@ -import React from "react" +/** + * Step 5 选择封面(Issue #1677 批量生成改造) + * - 单视频:保留原封面流程(自动生成/封面设置模板/封面预览) + * - N 个视频:N 张封面卡片,每张带对应视频标题,可逐个自动生成或上传 + */ +import React, { useRef } from "react" import { Modal, Spin } from "antd" +import { LoadingOutlined } from "@ant-design/icons" import type { CoverConfig } from "../types/cover" import type { GeneratedVideo } from "@/api/template-editor" import type { TitleSettings } from "../types" import { useStep6Cover } from "../hooks/useStep6Cover" +import { useBatchCovers } from "../hooks/useBatchCovers" import Button from "@/components/ui/Button" import CoverSettingsModal from "./cover-settings/CoverSettingsModal" import CoverEditorModal from "./cover-settings/CoverEditorModal" @@ -17,6 +24,15 @@ interface Step6CoverSettingsProps { titleSettings?: TitleSettings /** 确认生成步骤产出的最终视频列表 */ generatedVideos: GeneratedVideo[] + /* ── 批量生成(#1677)── */ + previewCount?: number + /** 每个变体的标题文字 */ + previewTitles?: string[] + /** 每个变体的封面URL(按变体索引) */ + previewCovers?: string[] + onPreviewCoversChange?: (urls: string[]) => void + /** 勾选的变体索引(批量封面按此顺序展示,与最终成片顺序一致) */ + selectedVariantIndexes?: number[] } const Step6CoverSettings: React.FC = (props) => { @@ -46,13 +62,173 @@ const Step6CoverSettings: React.FC = (props) => { generatedVideos: props.generatedVideos, }) - const handleAutoGenerate = () => { - generateAutoCover() - } + const previewCount = props.previewCount || 1 + const isBatch = previewCount > 1 + const previewTitles = props.previewTitles || [] + const previewCovers = props.previewCovers || [] + /** 卡片展示的变体索引顺序:批量=勾选顺序(与成片顺序一致),单视频=[0] */ + const cardIndexes = + isBatch && props.selectedVariantIndexes?.length + ? props.selectedVariantIndexes + : Array.from({ length: previewCount }, (_, i) => i) + const uploadInputRef = useRef(null) + const uploadTargetRef = useRef(0) + + const completedVideos = props.generatedVideos.filter((v) => v.status === "completed") + const batchTitles = cardIndexes.map((vi) => previewTitles[vi] || "") + const batchCoversList = cardIndexes.map((vi) => previewCovers[vi] || "") + const batchCovers = useBatchCovers({ + selectedTemplate: props.selectedTemplate || "", + generatedVideos: props.generatedVideos, + titles: batchTitles, + titleStyle: { + font: props.titleSettings?.font || "思源黑体", + size: props.titleSettings?.size || 28, + color: props.titleSettings?.color || "#ffffff", + position: props.titleSettings?.position || "top", + bold: props.titleSettings?.bold ?? true, + stroke: props.titleSettings?.stroke ?? true, + shadow: props.titleSettings?.shadow ?? false, + }, + covers: batchCoversList, + onCoversChange: (urls) => { + // 按卡片顺序写回对应变体索引 + const next = [...(props.previewCovers || [])] + cardIndexes.forEach((vi, cardPos) => { + next[vi] = urls[cardPos] || "" + }) + props.onPreviewCoversChange?.(next) + }, + }) - // 预览图:优先 thumbnail_url,其次 upload_url const previewUrl = coverSettings.thumbnail_url || coverSettings.upload_url + const handleUploadClick = (variantIndex: number) => { + uploadTargetRef.current = variantIndex + uploadInputRef.current?.click() + } + + const handleFileChange = (e: React.ChangeEvent) => { + const file = e.target.files?.[0] + e.target.value = "" + if (file) { + const variantIndex = uploadTargetRef.current + const cardPos = cardIndexes.indexOf(variantIndex) + if (cardPos >= 0) void batchCovers.uploadOne(cardPos, file) + } + } + + /* ── 批量封面 ── */ + if (isBatch) { + return ( +
+

🖼️ 选择封面

+ +
+ 🎬 共 {completedVideos.length} 个成片,封面将从对应成片中智能选帧并叠加该视频的标题 +
+ +
+ +
+ +
+ {cardIndexes.map((variantIndex, cardPos) => { + const url = batchCoversList[cardPos] + const isLoading = batchCovers.loadingIndex === cardPos + const isUploading = batchCovers.uploadingIndex === cardPos + const title = batchTitles[cardPos] + return ( +
+
视频 {variantIndex + 1}
+
+ {isLoading || isUploading ? ( +
+ } /> + {isLoading ? "AI 选帧中…" : "上传中…"} +
+ ) : url ? ( + {`视频${variantIndex + ) : ( +
+ 🖼️ + 未设置封面 +
+ )} +
9:16
+
+ {title && ( +
+ 标题:{title} +
+ )} +
+ + +
+
+ ) + })} +
+ + +
+ ) + } + + /* ── 单视频:原有流程保持不变 ── */ return (

🖼️ 选择封面

@@ -75,7 +251,7 @@ const Step6CoverSettings: React.FC = (props) => { )}
-
+ +
+
+ +
+

微信账号

+
+
+ 💬 +
+ {wechatBound ? ( + <> + + 已绑定微信{user?.wechat_nickname ? `(${user.wechat_nickname})` : ""} + +

可使用微信扫码登录本账号

+ + ) : ( + <> + 未绑定微信 +

绑定后可使用微信扫码快速登录

+ + )} +
+
+
+ {wechatBound ? ( + + ) : ( + + )} +
+
+
+ + setWechatBindOpen(false)} + onBindSuccess={handleBindSuccess} + />
) } diff --git a/apps/web/src/router/ProtectedRoute.tsx b/apps/web/src/router/ProtectedRoute.tsx index 2c49851ff..943511cc3 100644 --- a/apps/web/src/router/ProtectedRoute.tsx +++ b/apps/web/src/router/ProtectedRoute.tsx @@ -6,10 +6,16 @@ import { useAuthStore } from "@/store/authStore" export const ProtectedRoute = ({ children }: { children: React.ReactNode }) => { const isAuthenticated = useAuthStore((state) => state.isAuthenticated) const hasAccessToken = Boolean(localStorage.getItem("access_token")) + const profileCompleted = useAuthStore((state) => state.user?.profile_completed !== false) if (!isAuthenticated || !hasAccessToken) { return } + // 微信新用户未完成昵称引导时,禁止进入主界面 + if (!profileCompleted) { + return + } + return <>{children} } diff --git a/apps/web/src/router/appRoutes.tsx b/apps/web/src/router/appRoutes.tsx index 2858217c3..579b1b9a6 100644 --- a/apps/web/src/router/appRoutes.tsx +++ b/apps/web/src/router/appRoutes.tsx @@ -1,6 +1,7 @@ import { Navigate, type RouteObject } from "react-router-dom" import MainLayout from "@/components/layout/MainLayout" import { ProtectedRoute } from "./ProtectedRoute" +import { lazyRoute } from "./lazyRoute" /** * 受保护的 /app 子路由 @@ -13,202 +14,118 @@ const appChildren: RouteObject[] = [ }, { path: "dashboard", - lazy: () => - import("@/pages/dashboard/Dashboard").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/dashboard/Dashboard")), }, { path: "assets", - lazy: () => - import("@/pages/assets/AssetLibrary").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/assets/AssetLibrary")), }, { path: "titles", - lazy: () => - import("@/pages/titles/TitleLibrary").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/titles/TitleLibrary")), }, { path: "voices", - lazy: () => - import("@/pages/voices/VoiceLibrary").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/voices/VoiceLibrary")), }, { path: "templates", - lazy: () => - import("@/pages/templates/TemplateLibrary").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/templates/TemplateLibrary")), }, { path: "generate", - lazy: () => - import("@/pages/generate/GeneratePage").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/generate/GeneratePage")), }, { path: "history", - lazy: () => - import("@/pages/history/TaskHistory").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/history/TaskHistory")), }, { path: "products", - lazy: () => - import("@/pages/products/ProductLibrary").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/products/ProductLibrary")), }, { path: "products/:id", - lazy: () => - import("@/pages/products/ProductDetail").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/products/ProductDetail")), }, { path: "tasks", - lazy: () => - import("@/pages/tasks/TaskCenter").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/tasks/TaskCenter")), }, { path: "editing-planner", - lazy: () => - import("@/pages/editing-planner/EditingPlanner").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/editing-planner/EditingPlanner")), }, { path: "my-templates", - lazy: () => - import("@/pages/my-templates/MyTemplates").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/my-templates/MyTemplates")), }, { path: "voice-clone", - lazy: () => - import("@/pages/voice-clone/VoiceClone").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/voice-clone/VoiceClone")), }, { path: "voice-materials", - lazy: () => - import("@/pages/voice-materials/VoiceMaterialLibrary").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/voice-materials/VoiceMaterialLibrary")), }, { path: "my-voices", - lazy: () => - import("@/pages/my-voices/MyVoices").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/my-voices/MyVoices")), }, { path: "accounts", - lazy: () => - import("@/pages/accounts/Accounts").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/accounts/Accounts")), }, { path: "duplication", - lazy: () => - import("@/pages/duplication/DuplicationUpload").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/duplication/DuplicationUpload")), }, { path: "duplication/results", - lazy: () => - import("@/pages/duplication/DuplicationResults").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/duplication/DuplicationResults")), }, { path: "duplication/:id", - lazy: () => - import("@/pages/duplication/DuplicationDetail").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/duplication/DuplicationDetail")), }, { path: "subscription", - lazy: () => - import("@/pages/subscription/Plans").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/subscription/Plans")), }, { path: "subscription/upgrade", - lazy: () => - import("@/pages/subscription/UpgradeSubscription").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/subscription/UpgradeSubscription")), }, { path: "subscription/billing", - lazy: () => - import("@/pages/subscription/Billing").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/subscription/Billing")), }, { path: "profile", - lazy: () => - import("@/pages/profile/Settings").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/profile/Settings")), }, { path: "admin", children: [ { index: true, - lazy: () => - import("@/pages/admin/AdminComingSoon").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")), }, { path: "users", - lazy: () => - import("@/pages/admin/AdminComingSoon").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")), }, { path: "analytics", - lazy: () => - import("@/pages/admin/AdminComingSoon").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")), }, { path: "monitor", - lazy: () => - import("@/pages/admin/AdminComingSoon").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")), }, { path: "logs", - lazy: () => - import("@/pages/admin/AdminComingSoon").then((m) => ({ - Component: m.default, - })), + lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")), }, ], }, diff --git a/apps/web/src/router/lazyRoute.ts b/apps/web/src/router/lazyRoute.ts new file mode 100644 index 000000000..2a1970935 --- /dev/null +++ b/apps/web/src/router/lazyRoute.ts @@ -0,0 +1,40 @@ +import type { LazyRouteFunction, RouteObject } from "react-router-dom" +import { isChunkLoadError } from "@/utils/chunkLoadError" + +/** + * 给 React Router data router 的路由懒加载包一层自动重试: + * + * - 网络抖动 / 瞬态失败:自动重试最多 2 次(间隔 300ms / 800ms),用户无感恢复 + * - 发版后旧 chunk 404(chunk 文件名已不存在):重试也拿不到旧文件名, + * 重试耗尽后抛出,由全局 ChunkErrorBoundary 捕获并引导整页刷新 + * (刷新后 index.html 是 no-cache 的,会拿到新 chunk 引用) + */ +const RETRY_DELAYS_MS = [300, 800] +const RETRY_COUNT = RETRY_DELAYS_MS.length + +const sleep = (ms: number) => new Promise((r) => setTimeout(r, ms)) + +export const lazyRoute = ( + factory: () => Promise<{ default: React.ComponentType }>, +): LazyRouteFunction => { + return async () => { + let lastError: unknown + for (let attempt = 0; attempt <= RETRY_COUNT; attempt++) { + try { + const mod = await factory() + if (!mod.default) { + throw new Error("lazyRoute: 目标模块缺少 default 导出") + } + return { Component: mod.default } + } catch (err) { + lastError = err + // 非 chunk 加载错误(代码 bug 等)立即抛出,不浪费重试 + if (!isChunkLoadError(err)) throw err + if (attempt < RETRY_COUNT) { + await sleep(RETRY_DELAYS_MS[attempt]) + } + } + } + throw lastError + } +} diff --git a/apps/web/src/router/publicRoutes.tsx b/apps/web/src/router/publicRoutes.tsx index e610b941d..48c7245fb 100644 --- a/apps/web/src/router/publicRoutes.tsx +++ b/apps/web/src/router/publicRoutes.tsx @@ -5,6 +5,8 @@ import Register from "@/pages/auth/Register" import ForgotPassword from "@/pages/auth/ForgotPassword" import ResetPassword from "@/pages/auth/ResetPassword" import WechatCallback from "@/pages/auth/WechatCallback" +import WechatOnboarding from "@/pages/auth/WechatOnboarding" +import WechatBindCallback from "@/pages/auth/WechatBindCallback" import { useAuthStore } from "@/store/authStore" /** 首页路由组件:已登录跳 dashboard,未登录显示落地页 */ @@ -45,4 +47,12 @@ export const publicRoutes: RouteObject[] = [ path: "/auth/wechat/callback", element: , }, + { + path: "/auth/wechat/bind/callback", + element: , + }, + { + path: "/welcome/wechat", + element: , + }, ] diff --git a/apps/web/src/store/authStore.ts b/apps/web/src/store/authStore.ts index cef495665..3f6a2991c 100644 --- a/apps/web/src/store/authStore.ts +++ b/apps/web/src/store/authStore.ts @@ -13,6 +13,12 @@ interface User { display_name: string is_email_verified: boolean email_verified: boolean + wechat_bound?: boolean + wechat_nickname?: string + avatar_url?: string + phone?: string + phone_verified?: boolean + profile_completed?: boolean } interface AuthState { diff --git a/apps/web/src/test/api/assets.test.ts b/apps/web/src/test/api/assets.test.ts index 71f5788d7..cce914b2e 100644 --- a/apps/web/src/test/api/assets.test.ts +++ b/apps/web/src/test/api/assets.test.ts @@ -242,6 +242,20 @@ describe("assets API", () => { await expect(completeDirectUpload({ name: "test-item" })).resolves.not.toThrow() }) + it("请求体携带 file_size(后端同名兜底去重的大小校验依赖它)", async () => { + await completeDirectUpload({ + project_id: "p-1", + library_id: "l-1", + storage_key: "uploads/k.mp4", + file_size: 12345, + } as never) + const completeCalls = mockPost.mock.calls.filter( + ([u]: [string]) => u === "/upload/direct/complete", + ) + expect(completeCalls).toHaveLength(1) + expect(completeCalls[0][1]).toMatchObject({ file_size: 12345 }) + }) + it("should reject on API error", async () => { mockGet.mockRejectedValue(new Error("Network error")) mockPost.mockRejectedValue(new Error("Network error")) @@ -259,6 +273,114 @@ describe("assets API", () => { }) }) + describe("uploadAssetDirect skip_transfer 短路", () => { + it("prepare 返回 skip_transfer=true → 直接返回 duplicated,不调 transfer/complete", async () => { + mockPost.mockImplementation((url: string) => { + if (url === "/upload/direct/prepare") { + return Promise.resolve({ + data: { + upload_url: "https://oss/x", + method: "POST", + storage_key: "uploads/skip/y.mp4", + expires_at: "2099", + fields: {}, + max_size_bytes: 1e9, + asset_id: "existing-asset", + skip_transfer: true, + duplicated: true, + }, + }) + } + if (url === "/upload/direct/complete") { + throw new Error("complete 不应被调用") + } + throw new Error("unexpected url " + url) + }) + const putSpy = vi.spyOn(globalThis, "XMLHttpRequest") + const file = new File(["x"], "x.mp4", { type: "video/mp4" }) + const result = await uploadAssetDirect({ file, library_id: "lib-1" }) + expect(result.duplicated).toBe(true) + expect(result.asset_id).toBe("existing-asset") + // complete 未被调用(mockPost 只记录 prepare,complete 若调用会抛 "不应被调用") + const completeCalls = mockPost.mock.calls.filter( + ([u]: [string]) => u === "/upload/direct/complete", + ) + expect(completeCalls).toHaveLength(0) + putSpy.mockRestore() + }) + + it("prepare 返回 skip_transfer=false → 走老流程(complete 被调用)", async () => { + mockPost.mockImplementation((url: string) => { + if (url === "/upload/direct/prepare") { + return Promise.resolve({ + data: { + upload_url: "https://oss/x", + method: "POST", + storage_key: "uploads/normal/y.mp4", + expires_at: "2099", + fields: {}, + max_size_bytes: 1e9, + asset_id: "new-asset", + }, + }) + } + if (url === "/upload/direct/complete") { + return Promise.resolve({ + data: { + storage_key: "uploads/normal/y.mp4", + ingest_job_id: "job-1", + url: "https://oss/y.mp4", + duplicated: false, + asset_id: "new-asset", + }, + }) + } + throw new Error("unexpected url " + url) + }) + // mock XMLHttpRequest:send 之后下一 tick 触发 onload 让 transfer 立即成功 + const origOpen = XMLHttpRequest.prototype.open + const origSend = XMLHttpRequest.prototype.send + const origSetReadyState = Object.getOwnPropertyDescriptor( + XMLHttpRequest.prototype, + "readyState", + ) as PropertyDescriptor | undefined + const origStatus = Object.getOwnPropertyDescriptor(XMLHttpRequest.prototype, "status") + Object.defineProperty(XMLHttpRequest.prototype, "readyState", { + configurable: true, + writable: true, + value: 4, + }) + Object.defineProperty(XMLHttpRequest.prototype, "status", { + configurable: true, + writable: true, + value: 200, + }) + XMLHttpRequest.prototype.open = vi.fn() as unknown as typeof origOpen + XMLHttpRequest.prototype.send = vi.fn(function (this: XMLHttpRequest) { + // 下一 tick 触发 onload(模拟 XHR 异步完成) + setTimeout(() => this.onload?.(new ProgressEvent("load")), 0) + }) as unknown as typeof origSend + const file = new File(["x"], "x.mp4", { type: "video/mp4" }) + const result = await uploadAssetDirect({ file, library_id: "lib-1" }) + expect(result.duplicated).toBeFalsy() + expect(result.asset_id).toBe("new-asset") + const completeCalls = mockPost.mock.calls.filter( + ([u]: [string]) => u === "/upload/direct/complete", + ) + expect(completeCalls).toHaveLength(1) + // complete 请求必须带上 file_size,否则后端同名兜底会误杀同名新视频 + expect(completeCalls[0][1]).toMatchObject({ file_size: file.size }) + XMLHttpRequest.prototype.open = origOpen + XMLHttpRequest.prototype.send = origSend + if (origSetReadyState) { + Object.defineProperty(XMLHttpRequest.prototype, "readyState", origSetReadyState) + } + if (origStatus) { + Object.defineProperty(XMLHttpRequest.prototype, "status", origStatus) + } + }) + }) + describe("getIngestJob", () => { it("should resolve successfully", async () => { await expect(getIngestJob("test-jobId")).resolves.not.toThrow() diff --git a/apps/web/src/test/api/uploadDedup.test.ts b/apps/web/src/test/api/uploadDedup.test.ts new file mode 100644 index 000000000..526df30fb --- /dev/null +++ b/apps/web/src/test/api/uploadDedup.test.ts @@ -0,0 +1,130 @@ +/** + * 上传去重/幂等工具单测(Issue #1714) + */ +import { describe, it, expect, vi } from "vitest" +import { + computeFileHash, + findDuplicateInQueue, + HASH_FULL_READ_LIMIT, + HASH_SAMPLE_CHUNK, + makeClientUploadId, + makeFileFingerprint, +} from "@/api/assets/uploadDedup" + +const makeFile = (name: string, size = 100, lastModified = 1_700_000_000_000) => + new File([new Uint8Array(size)], name, { type: "video/mp4", lastModified }) + +describe("makeFileFingerprint", () => { + it("同一文件(name+size+lastModified 相同)指纹一致", () => { + const a = makeFile("a.mp4", 1000, 12345) + const b = makeFile("a.mp4", 1000, 12345) + expect(makeFileFingerprint(a)).toBe(makeFileFingerprint(b)) + }) + + it("文件名/大小/修改时间任一不同指纹即不同", () => { + const base = makeFile("a.mp4", 1000, 100) + expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("b.mp4", 1000, 100))) + expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("a.mp4", 1001, 100))) + expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("a.mp4", 1000, 101))) + }) +}) + +describe("findDuplicateInQueue", () => { + const queue = [ + { fileKey: "k1", status: "preparing" }, + { fileKey: "k2", status: "uploading" }, + { fileKey: "k3", status: "ingesting" }, + { fileKey: "k4", status: "done" }, + { fileKey: "k5", status: "error" }, + ] + + it("在途状态(preparing/uploading/ingesting/done)命中重复", () => { + expect(findDuplicateInQueue(queue, "k1")?.status).toBe("preparing") + expect(findDuplicateInQueue(queue, "k2")?.status).toBe("uploading") + expect(findDuplicateInQueue(queue, "k3")?.status).toBe("ingesting") + expect(findDuplicateInQueue(queue, "k4")?.status).toBe("done") + }) + + it("未命中返回 null", () => { + expect(findDuplicateInQueue(queue, "missing")).toBeNull() + }) + + it("排除 error 状态后,失败项不算重复(允许重新激活)", () => { + expect(findDuplicateInQueue(queue, "k5", ["error"])).toBeNull() + }) + + it("同时排除 done 后,已完成项也不算重复", () => { + expect(findDuplicateInQueue(queue, "k4", ["error", "done"])).toBeNull() + // 但在途的仍然命中 + expect(findDuplicateInQueue(queue, "k1", ["error", "done"])).not.toBeNull() + }) +}) + +describe("makeClientUploadId", () => { + it("生成带前缀且互不相同的幂等 token", () => { + const ids = new Set(Array.from({ length: 20 }, () => makeClientUploadId())) + expect(ids.size).toBe(20) + for (const id of ids) expect(id.startsWith("up_")).toBe(true) + }) +}) + +describe("computeFileHash", () => { + it("相同内容 hash 一致、不同内容 hash 不同", async () => { + const f1 = makeFile("a.mp4", 4096) + const f2 = makeFile("b.mp4", 4096) + // 两个文件都是 0 填充,内容相同 → hash 一致 + expect(await computeFileHash(f1)).toBe(await computeFileHash(f2)) + + const f3 = new File([new Uint8Array(4096).fill(7)], "c.mp4", { type: "video/mp4" }) + expect(await computeFileHash(f1)).not.toBe(await computeFileHash(f3)) + }) + + it("返回 64 位十六进制(SHA-256,与后端 file_hash 长度一致)", async () => { + const hash = await computeFileHash(makeFile("a.mp4", 1024)) + expect(hash).toMatch(/^[0-9a-f]{64}$/) + }) +}) + +describe("computeFileHash 大文件抽样(>64MB)", () => { + it("抽样路径正常返回 64 位 hex,且大小不同则 hash 不同", async () => { + // mock 一个「声称」300MB 的 File:slice 返回小 buffer 即可,不真分配 300MB + const makeBig = (declaredSize: number, head: number) => { + const f = new File([new Uint8Array([head, 2, 3])], "big.mov", { type: "video/quicktime" }) + Object.defineProperty(f, "size", { value: declaredSize, configurable: true }) + // slice 仍按真实内容返回小片段(头尾片段内容由底层小 buffer 决定) + return f + } + const h1 = await computeFileHash(makeBig(300 * 1024 * 1024, 1)) + const h2 = await computeFileHash(makeBig(301 * 1024 * 1024, 1)) + expect(h1).toMatch(/^[0-9a-f]{64}$/) + // 声明大小不同 → 写入的 64 位 size 字段不同 → hash 必须不同(锁定 setBigUint64 路径) + expect(h1).not.toBe(h2) + }) + + it("≤64MB 走全量读取(slice 一次覆盖整个文件)", async () => { + const f = new File([new Uint8Array(1024).fill(9)], "full.mp4", { type: "video/mp4" }) + Object.defineProperty(f, "size", { value: HASH_FULL_READ_LIMIT, configurable: true }) + const sliceSpy = vi.spyOn(f, "slice") + await computeFileHash(f) + // 全量路径:唯一一次 slice 为 (0, size) + expect(sliceSpy).toHaveBeenCalledTimes(1) + expect(sliceSpy).toHaveBeenCalledWith(0, HASH_FULL_READ_LIMIT) + sliceSpy.mockRestore() + }) + + it(">64MB 只读取头尾各 16MB 抽样,绝不整文件读入内存", async () => { + const f = new File([new Uint8Array(1024).fill(9)], "big.mp4", { type: "video/mp4" }) + Object.defineProperty(f, "size", { value: HASH_FULL_READ_LIMIT + 1, configurable: true }) + const sliceSpy = vi.spyOn(f, "slice") + await computeFileHash(f) + // 抽样路径:两次 slice —— 头部 (0, 16MB) 与尾部 (size-16MB, size) + expect(sliceSpy).toHaveBeenCalledTimes(2) + expect(sliceSpy).toHaveBeenNthCalledWith(1, 0, HASH_SAMPLE_CHUNK) + expect(sliceSpy).toHaveBeenNthCalledWith( + 2, + HASH_FULL_READ_LIMIT + 1 - HASH_SAMPLE_CHUNK, + HASH_FULL_READ_LIMIT + 1, + ) + sliceSpy.mockRestore() + }) +}) diff --git a/apps/web/src/test/api/wxLogin.test.ts b/apps/web/src/test/api/wxLogin.test.ts new file mode 100644 index 000000000..5cc38730d --- /dev/null +++ b/apps/web/src/test/api/wxLogin.test.ts @@ -0,0 +1,64 @@ +import { describe, expect, it, vi, beforeEach, afterEach } from "vitest" + +describe("wxLogin 工具", () => { + describe("parseWxAuthUrl", () => { + it("从微信授权链接解析出 appid/redirect_uri/state(redirect_uri 解码)", async () => { + const { parseWxAuthUrl } = await import("@/api/auth/wxLogin") + const authUrl = + "https://open.weixin.qq.com/connect/qrconnect?appid=wxb7ae80b48e53980d" + + "&redirect_uri=https%3A%2F%2Fstaging.xiaoxiajianji.com%2Fauth%2Fwechat%2Fcallback" + + "&response_type=code&scope=snsapi_login&state=abc123#wechat_redirect" + const params = parseWxAuthUrl(authUrl) + expect(params).not.toBeNull() + expect(params?.appid).toBe("wxb7ae80b48e53980d") + expect(params?.redirect_uri).toBe("https://staging.xiaoxiajianji.com/auth/wechat/callback") + expect(params?.state).toBe("abc123") + }) + + it("链接里缺 state 时回退使用 stateFallback", async () => { + const { parseWxAuthUrl } = await import("@/api/auth/wxLogin") + const authUrl = + "https://open.weixin.qq.com/connect/qrconnect?appid=wx123" + + "&redirect_uri=https%3A%2F%2Fexample.com%2Fcb" + const params = parseWxAuthUrl(authUrl, "fallback-state") + expect(params?.state).toBe("fallback-state") + }) + + it("缺 appid 或 redirect_uri 时返回 null(调用方应回退整页跳转)", async () => { + const { parseWxAuthUrl } = await import("@/api/auth/wxLogin") + expect(parseWxAuthUrl("https://open.weixin.qq.com/connect/qrconnect?appid=wx123")).toBeNull() + expect(parseWxAuthUrl("not a url")).toBeNull() + }) + }) + + describe("loadWxLoginScript", () => { + beforeEach(() => { + vi.resetModules() + document.head.querySelectorAll("script[src*='wxLogin']").forEach((el) => el.remove()) + delete (window as unknown as { WxLogin?: unknown }).WxLogin + }) + afterEach(() => { + vi.restoreAllMocks() + }) + + it("window.WxLogin 已存在时直接复用,不重复插入 script", async () => { + const fakeCtor = vi.fn() + ;(window as unknown as { WxLogin: unknown }).WxLogin = fakeCtor + const { loadWxLoginScript } = await import("@/api/auth/wxLogin") + const ctor = await loadWxLoginScript() + expect(ctor).toBe(fakeCtor) + expect(document.head.querySelector("script[src*='wxLogin']")).toBeNull() + }) + + it("脚本 onerror 时 reject(调用方据此回退整页跳转)", async () => { + const { loadWxLoginScript } = await import("@/api/auth/wxLogin") + const promise = loadWxLoginScript() + const script = document.head.querySelector( + "script[src*='wxLogin']", + ) as HTMLScriptElement | null + expect(script).not.toBeNull() + script?.dispatchEvent(new Event("error")) + await expect(promise).rejects.toThrow(/加载失败/) + }) + }) +}) diff --git a/apps/web/src/test/components/ChunkErrorBoundary.test.tsx b/apps/web/src/test/components/ChunkErrorBoundary.test.tsx new file mode 100644 index 000000000..c049dcb95 --- /dev/null +++ b/apps/web/src/test/components/ChunkErrorBoundary.test.tsx @@ -0,0 +1,79 @@ +import { describe, it, expect, beforeEach, afterEach, vi } from "vitest" +import { render, screen, fireEvent } from "@testing-library/react" +import { Button } from "antd" +import { useState } from "react" +import ChunkErrorBoundary from "@/components/common/ChunkErrorBoundary" +import * as chunkUtils from "@/utils/chunkLoadError" + +// reload 函数 mock 掉(jsdom 不支持真实 window.location.reload) +vi.mock("@/utils/chunkLoadError", async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + reloadForChunkError: vi.fn(), + goHomeRecover: vi.fn(), + } +}) +const { reloadForChunkError, goHomeRecover } = vi.mocked(chunkUtils) + +/** 渲染时直接抛错的子组件 */ +const Boom: React.FC<{ error: Error }> = ({ error }) => { + throw error +} + +/** 点击按钮后才抛 chunk 错误的子组件 */ +const ChunkBoomButton: React.FC = () => { + const [boom, setBoom] = useState(false) + if (boom) { + throw new TypeError("Failed to fetch dynamically imported module: /assets/x.js") + } + return +} + +const renderBoundary = (ui: React.ReactNode) => + render({ui}) + +beforeEach(() => { + sessionStorage.clear() + vi.clearAllMocks() + // error boundary 捕获后 React 会打 error log,静默掉 + vi.spyOn(console, "error").mockImplementation(() => {}) +}) + +afterEach(() => { + vi.restoreAllMocks() + sessionStorage.clear() +}) + +describe("ChunkErrorBoundary", () => { + it("正常渲染 children", () => { + renderBoundary(
hello-child
) + expect(screen.getByText("hello-child")).toBeInTheDocument() + }) + + it("首次捕获 chunk 错误 → 自动刷新(reloadForChunkError)并显示自动刷新提示", () => { + renderBoundary() + fireEvent.click(screen.getByText("boom")) + expect(reloadForChunkError).toHaveBeenCalledTimes(1) + expect(screen.getByText(/正在自动刷新/)).toBeInTheDocument() + }) + + it("已刷新过仍失败 → 不再自动刷新,显示手动兜底按钮", () => { + // 模拟"本会话已经自动刷新过一次" + sessionStorage.setItem("chunk_error_reloaded_at", String(Date.now())) + renderBoundary( + , + ) + expect(reloadForChunkError).not.toHaveBeenCalled() + expect(screen.getByText("系统已更新")).toBeInTheDocument() + // 点击兜底按钮 → goHomeRecover(跳首页,不刷新当前 URL) + fireEvent.click(screen.getByText("刷新并返回首页")) + expect(goHomeRecover).toHaveBeenCalledTimes(1) + }) + + it("非 chunk 错误 → 显示通用错误页,不触发 chunk 自动刷新", () => { + renderBoundary() + expect(reloadForChunkError).not.toHaveBeenCalled() + expect(screen.getByText("页面出现异常")).toBeInTheDocument() + }) +}) diff --git a/apps/web/src/test/components/WechatQrModal.test.tsx b/apps/web/src/test/components/WechatQrModal.test.tsx new file mode 100644 index 000000000..c3910e0d5 --- /dev/null +++ b/apps/web/src/test/components/WechatQrModal.test.tsx @@ -0,0 +1,165 @@ +import { describe, expect, it, vi, beforeEach, afterEach } from "vitest" +import { render, screen, waitFor, cleanup, fireEvent } from "@testing-library/react" +import WechatQrModal from "@/components/auth/WechatQrModal" + +const { mockWxLoginCtor, mockGetAuthUrl, mockGetBindUrl, mockGetCurrentUser } = vi.hoisted(() => ({ + mockWxLoginCtor: vi.fn(), + mockGetAuthUrl: vi.fn(), + mockGetBindUrl: vi.fn(), + mockGetCurrentUser: vi.fn(), +})) + +vi.mock("@/api/auth", () => ({ + getWechatAuthUrl: (...args: unknown[]) => mockGetAuthUrl(...args), + getWechatBindUrl: (...args: unknown[]) => mockGetBindUrl(...args), + getCurrentUser: (...args: unknown[]) => mockGetCurrentUser(...args), + normalizeUser: (u: unknown) => u, +})) + +vi.mock("@/api/auth/wxLogin", () => ({ + loadWxLoginScript: vi.fn(async () => mockWxLoginCtor), + parseWxAuthUrl: vi.fn(() => ({ + appid: "wxb7ae80b48e53980d", + redirect_uri: "https://staging.xiaoxiajianji.com/auth/wechat/callback", + state: "state-from-url", + })), +})) + +vi.mock("@/api/auth/tokenRefresh", () => ({ + scheduleProactiveRefresh: vi.fn(), + cancelProactiveRefresh: vi.fn(), +})) + +const { mockSetAuth, mockSetUser } = vi.hoisted(() => ({ + mockSetAuth: vi.fn(), + mockSetUser: vi.fn(), +})) +vi.mock("@/store/authStore", () => ({ + useAuthStore: (selector: (s: unknown) => unknown) => + selector({ setAuth: mockSetAuth, setUser: mockSetUser }), +})) + +const AUTH_URL = + "https://open.weixin.qq.com/connect/qrconnect?appid=wxb7ae80b48e53980d" + + "&redirect_uri=https%3A%2F%2Fstaging.xiaoxiajianji.com%2Fauth%2Fwechat%2Fcallback&state=st123" + +const postMessage = (data: Record) => + window.dispatchEvent(new MessageEvent("message", { data, origin: window.location.origin })) + +beforeEach(() => { + vi.clearAllMocks() + mockGetAuthUrl.mockResolvedValue({ auth_url: AUTH_URL, state: "st123" }) + mockGetBindUrl.mockResolvedValue({ auth_url: AUTH_URL, state: "st123" }) + mockGetCurrentUser.mockResolvedValue({ id: 1, display_name: "测试用户" }) + localStorage.clear() +}) + +afterEach(() => cleanup()) + +describe("WechatQrModal", () => { + it("open=false 时不渲染弹窗内容", () => { + render() + expect(screen.queryByText("微信扫码登录")).toBeNull() + }) + + it("登录场景:open 后请求授权链接、写入 state、用 WxLogin 渲染二维码", async () => { + render() + await waitFor(() => expect(mockGetAuthUrl).toHaveBeenCalledTimes(1)) + expect(localStorage.getItem("wechat_state")).toBe("st123") + await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1)) + expect(mockWxLoginCtor).toHaveBeenCalledWith( + expect.objectContaining({ + self_redirect: true, + appid: "wxb7ae80b48e53980d", + scope: "snsapi_login", + state: "state-from-url", + redirect_uri: "https://staging.xiaoxiajianji.com/auth/wechat/callback", + }), + ) + expect(screen.getByText(/请使用微信扫描二维码登录/)).toBeTruthy() + }) + + it("绑定场景:请求 bind/url 且写入 wechat_bind_state", async () => { + render() + await waitFor(() => expect(mockGetBindUrl).toHaveBeenCalledTimes(1)) + expect(mockGetAuthUrl).not.toHaveBeenCalled() + expect(localStorage.getItem("wechat_bind_state")).toBe("st123") + await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1)) + }) + + it("获取授权链接失败时弹窗内展示错误并提供刷新", async () => { + mockGetAuthUrl.mockRejectedValueOnce({ + response: { status: 500, data: { detail: "微信服务内部错误" } }, + }) + render() + expect(await screen.findByText(/微信服务内部错误/)).toBeTruthy() + expect(screen.getByText("刷新二维码")).toBeTruthy() + // 点刷新后重新请求 + fireEvent.click(screen.getByText("刷新二维码")) + await waitFor(() => expect(mockGetAuthUrl).toHaveBeenCalledTimes(2)) + }) + + it("登录成功消息:同步登录态并回调 onLoginSuccess(needOnboarding)", async () => { + const onSuccess = vi.fn() + localStorage.setItem("access_token", "tok-123") + render() + await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1)) + + postMessage({ + source: "xiaoxia-wechat-qr", + scene: "login", + success: true, + payload: { needOnboarding: true }, + }) + + await waitFor(() => expect(onSuccess).toHaveBeenCalledWith(true)) + expect(mockGetCurrentUser).toHaveBeenCalled() + expect(mockSetAuth).toHaveBeenCalledWith(expect.objectContaining({ id: 1 }), "tok-123", null) + }) + + it("登录失败消息:弹窗内展示回调页透传的真实原因", async () => { + render() + await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1)) + + postMessage({ + source: "xiaoxia-wechat-qr", + scene: "login", + success: false, + detail: "微信登录失败:state 已过期或已被使用", + }) + + expect(await screen.findByText(/state 已过期或已被使用/)).toBeTruthy() + }) + + it("绑定成功消息:刷新用户并回调 onBindSuccess", async () => { + const onBindSuccess = vi.fn() + render() + await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1)) + + postMessage({ source: "xiaoxia-wechat-qr", scene: "bind", success: true }) + + await waitFor(() => expect(onBindSuccess).toHaveBeenCalledTimes(1)) + expect(mockSetUser).toHaveBeenCalled() + }) + + it("忽略跨源消息和其他场景的消息", async () => { + const onSuccess = vi.fn() + render() + await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1)) + + // 跨源 + window.dispatchEvent( + new MessageEvent("message", { + data: { source: "xiaoxia-wechat-qr", scene: "login", success: true }, + origin: "https://evil.example.com", + }), + ) + // 场景不符(bind 消息发给 login 弹窗) + postMessage({ source: "xiaoxia-wechat-qr", scene: "bind", success: true }) + // 无协议标识 + postMessage({ foo: "bar" }) + + await new Promise((r) => setTimeout(r, 50)) + expect(onSuccess).not.toHaveBeenCalled() + }) +}) diff --git a/apps/web/src/test/pages/Settings.test.tsx b/apps/web/src/test/pages/Settings.test.tsx index e0614b0bf..3d0711b64 100644 --- a/apps/web/src/test/pages/Settings.test.tsx +++ b/apps/web/src/test/pages/Settings.test.tsx @@ -1,8 +1,8 @@ -import { describe, expect, it, vi } from "vitest" -import { render, screen } from "@testing-library/react" +import { describe, expect, it, vi, beforeEach } from "vitest" +import { render, screen, fireEvent, waitFor } from "@testing-library/react" import { MemoryRouter } from "react-router-dom" +import { QueryClient, QueryClientProvider } from "@tanstack/react-query" -// mock PageHead 简单mock vi.mock("@/components/layout/PageHead", () => ({ default: ({ title, description }: { title: string; description?: string }) => (
@@ -12,51 +12,142 @@ vi.mock("@/components/layout/PageHead", () => ({ ), })) +const mockSetUser = vi.fn() +const mockInvalidate = vi.fn() +let authState: Record = { + user: { + id: "1", + user_id: "1", + username: "testuser", + email: "test@example.com", + display_name: "Test User", + wechat_bound: false, + }, + isAuthenticated: true, + setUser: mockSetUser, +} + vi.mock("@/store/authStore", () => ({ - useAuthStore: (selector: (state: any) => any) => - selector({ + useAuthStore: (selector: (state: unknown) => unknown) => selector(authState), +})) + +const getCurrentUserMock = vi.fn(async () => authState.user as Record) +const updateProfileMock = vi.fn() +const getWechatBindUrlMock = vi.fn(async () => ({ + auth_url: "https://wx.example/auth", + state: "s1", +})) +const unbindWechatMock = vi.fn(async () => ({ success: true })) + +vi.mock("@/api/auth", () => ({ + getCurrentUser: () => getCurrentUserMock(), + updateProfile: (d: unknown) => updateProfileMock(d), + getWechatBindUrl: () => getWechatBindUrlMock(), + unbindWechat: () => unbindWechatMock(), +})) + +vi.mock("antd", async () => { + const actual = await vi.importActual("antd") + return { ...actual, message: { success: vi.fn(), error: vi.fn() } } +}) + +import Settings from "@/pages/profile/Settings" + +const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false }, mutations: { retry: false } }, +}) + +const renderPage = () => + render( + + + + + , + ) + +describe("Settings Page", () => { + beforeEach(() => { + vi.clearAllMocks() + authState = { user: { id: "1", user_id: "1", username: "testuser", email: "test@example.com", display_name: "Test User", - is_email_verified: true, - email_verified: true, + wechat_bound: false, }, isAuthenticated: true, - }), -})) - -import Settings from "@/pages/profile/Settings" - -describe("Settings Page", () => { - it("should render without crashing", () => { - render( - - - , - ) - expect(screen.getByText("个人设置")).toBeTruthy() + setUser: mockSetUser, + } }) - it("should display user info", () => { - render( - - - , - ) + it("渲染个人设置与用户信息", () => { + renderPage() + expect(screen.getByText("个人设置")).toBeTruthy() expect(screen.getByDisplayValue("testuser")).toBeTruthy() expect(screen.getByDisplayValue("test@example.com")).toBeTruthy() }) - it("should show save button is disabled", () => { - render( - - - , + it("未绑定时显示绑定微信按钮,点击跳转微信授权", async () => { + renderPage() + expect(screen.getByText("未绑定微信")).toBeTruthy() + const btn = screen.getByText("绑定微信") + fireEvent.click(btn) + await waitFor(() => { + expect(getWechatBindUrlMock).toHaveBeenCalled() + expect(localStorage.getItem("wechat_bind_state")).toBe("s1") + }) + }) + + it("已绑定时显示状态与解绑按钮,确认后调解绑接口", async () => { + authState.user = { + ...(authState.user as object), + wechat_bound: true, + wechat_nickname: "微信昵称", + } as never + renderPage() + expect(screen.getByText(/已绑定微信/)).toBeTruthy() + fireEvent.click( + screen.getByText( + (_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").replace(/\s/g, "") === "解绑", + ), ) - const button = screen.getByText("保存暂未开放") - expect(button).toBeTruthy() + // antd Modal.confirm 弹确认框(标题+内容均含"解绑微信",用 role=dialog 内的确认按钮) + await waitFor(() => { + expect(document.querySelector(".ant-modal-confirm")).toBeTruthy() + }) + fireEvent.click( + screen.getByText( + (_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").includes("确定解绑"), + ), + ) + await waitFor(() => { + expect(unbindWechatMock).toHaveBeenCalled() + }) + }) + + it("修改昵称后保存按钮可用,点击调用更新接口", async () => { + renderPage() + const saveBtn = screen.getByText( + (_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").replace(/\s/g, "") === "保存", + ) + expect(saveBtn.closest("button")?.disabled).toBe(true) + fireEvent.change(screen.getByDisplayValue("Test User"), { + target: { value: "新昵称" }, + }) + await waitFor(() => { + expect(saveBtn.closest("button")?.disabled).toBe(false) + }) + updateProfileMock.mockResolvedValueOnce({ + id: "1", + display_name: "新昵称", + wechat_bound: false, + }) + fireEvent.click(saveBtn) + await waitFor(() => { + expect(updateProfileMock).toHaveBeenCalledWith({ display_name: "新昵称" }) + }) }) }) diff --git a/apps/web/src/test/pages/assets/useAssetUpload.test.tsx b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx index 3bc612669..2ad8d634c 100644 --- a/apps/web/src/test/pages/assets/useAssetUpload.test.tsx +++ b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx @@ -26,22 +26,32 @@ interface FakeHandle { fields: Record max_size_bytes: number asset_id: string + duplicated?: boolean + skip_transfer?: boolean } transfer: ReturnType complete: ReturnType /** 手动结束传输(transfer 被调用后挂载);finish(true) 以失败结束 */ finish: (fail?: boolean) => void + /** complete 已被调用的次数 */ + completeCalls: { resolve: () => void; reject: (err: unknown) => void }[] } -let activeTransfers = 0 -let maxConcurrent = 0 - /** - * 创建一个假 handle:transfer 返回挂起的 promise, - * finish 槽位在 transfer executor 同步执行时挂载,测试中调用 finish() 控制成败 + * 创建一个假 handle: + * - transfer 返回挂起的 promise,finish()/finish(true) 控制成败 + * - complete 每次调用返回独立的挂起 promise,由 completeCalls 记录控制, + * 成功调 resolve(idx) / 失败调 reject(idx)(模拟超时) */ -const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?: boolean }) => { - const h = { +const makeFakeHandle = (opts: { + id: string + duplicated?: boolean + failTransfer?: boolean + completeAuto?: boolean + /** prepare 阶段就命中去重:prepare 响应 skip_transfer/duplicated=true */ + prepareDedup?: boolean +}) => { + const h: FakeHandle = { prepared: { upload_url: "https://oss.example.com/u", method: "POST", @@ -50,17 +60,43 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?: fields: {}, max_size_bytes: 2_000_000_000, asset_id: opts.id, + duplicated: opts.prepareDedup ? true : undefined, + skip_transfer: opts.prepareDedup ? true : undefined, }, transfer: vi.fn(), - complete: vi.fn().mockResolvedValue({ - storage_key: "uploads/x/y.mp4", - ingest_job_id: opts.duplicated ? "" : "job-1", - url: "https://oss.example.com/u", - duplicated: opts.duplicated, - asset_id: opts.id, - }), - finish: (() => {}) as (fail?: boolean) => void, + complete: vi.fn(), + finish: () => {}, + completeCalls: [], } + + h.complete.mockImplementation( + () => + new Promise<{ + storage_key: string + ingest_job_id: string + url: string + duplicated: boolean + asset_id: string + }>((resolve, reject) => { + h.completeCalls.push({ + resolve: () => + resolve({ + storage_key: "uploads/x/y.mp4", + ingest_job_id: opts.duplicated ? "" : `job-${opts.id}`, + url: "https://oss.example.com/u", + duplicated: !!opts.duplicated, + asset_id: opts.id, + }), + reject, + }) + // 默认立即成功,保持旧用例简单 + if (opts.completeAuto !== false) { + const idx = h.completeCalls.length - 1 + Promise.resolve().then(() => h.completeCalls[idx]?.resolve()) + } + }), + ) + h.transfer.mockImplementation( () => new Promise((_resolve, reject) => { @@ -78,17 +114,24 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?: type FakeHandleLike = ReturnType +let activeTransfers = 0 +let maxConcurrent = 0 + /** prepare mock:调用序号生成稳定 id,立即把 handle(含 finish 槽位)推入数组 */ const installPrepareMock = ( handles: FakeHandleLike[], - optOverrides?: (id: string) => { duplicated?: boolean; failTransfer?: boolean }, + optOverrides?: (id: string) => { + duplicated?: boolean + failTransfer?: boolean + completeAuto?: boolean + }, ) => { let callNo = 0 ;(prepareDirectUploadHandle as unknown as ReturnType).mockImplementation( async () => { const id = `asset-${callNo++}` const overrides = optOverrides?.(id) ?? {} - const h = makeFakeHandle({ id, ...overrides }) + const h = makeFakeHandle({ id, completeAuto: true, ...overrides }) handles.push(h) await new Promise((r) => setTimeout(r, 10)) return h @@ -187,6 +230,11 @@ describe("useAssetUpload", () => { }) await waitFor(() => expect(result.current.uploadItems[0].status).toBe("error")) + // 失败卡片记录失败阶段与完整错误原因(不再只显示"上传失败") + const failed = result.current.uploadItems[0] + expect(failed.failedStage).toBe("transfer") + expect(failed.error).toContain("OSS boom") + // 重试:重新 prepare(handles[1] 成功) const tempId = result.current.uploadItems[0].tempId await act(async () => { @@ -223,4 +271,142 @@ describe("useAssetUpload", () => { expect(result.current.uploadItems[0].status).toBe("done") }) }) + it("同一文件多次选择不重复入队(指纹去重)", async () => { + const handles: FakeHandleLike[] = [] + installPrepareMock(handles) + + const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), { + wrapper: createWrapper(), + }) + + // 同一文件(name+size+lastModified 完全一致)第一次入队 + const sameFile = mp4("same.mp4") + await act(async () => { + result.current.enqueueUploads([sameFile]) + }) + await waitFor(() => expect(handles.length).toBe(1)) + expect(result.current.uploadItems).toHaveLength(1) + + // transfer 挂起期间,再次选择同一文件(模拟用户反复点选/拖拽) + await act(async () => { + result.current.enqueueUploads([sameFile]) + }) + await act(async () => { + result.current.enqueueUploads([sameFile]) + }) + // 队列表只有 1 项、prepare 只有 1 次 + expect(result.current.uploadItems).toHaveLength(1) + expect(handles.length).toBe(1) + + // 完成后再次重复选择(已 done):仍然不新增 + await act(async () => { + handles[0].finish() + }) + await waitFor(() => expect(result.current.uploadItems[0].status).toBe("done")) + await act(async () => { + result.current.enqueueUploads([sameFile]) + }) + expect(result.current.uploadItems).toHaveLength(1) + expect(handles.length).toBe(1) + + // 不同文件正常入队 + await act(async () => { + result.current.enqueueUploads([mp4("other.mp4")]) + }) + await waitFor(() => expect(handles.length).toBe(2)) + expect(result.current.uploadItems).toHaveLength(2) + }) + + it("complete 失败(超时)后重试:只重发 complete,不重新 prepare/直传", async () => { + const handles: FakeHandleLike[] = [] + installPrepareMock(handles, () => ({ completeAuto: false })) + + const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), { + wrapper: createWrapper(), + }) + + await act(async () => { + result.current.enqueueUploads([mp4("slow.mp4")]) + }) + await waitFor(() => expect(handles.length).toBe(1)) + await waitFor(() => expect(handles[0].transfer).toHaveBeenCalled()) + await act(async () => { + handles[0].finish() + }) + // complete 被调用但挂起 + await waitFor(() => expect(handles[0].complete).toHaveBeenCalledTimes(1)) + + // 模拟 complete 超时(后端记录可能已建成) + await act(async () => { + handles[0].completeCalls[0]?.reject(new Error("complete timeout (ECONNABORTED)")) + }) + const tempId = result.current.uploadItems[0].tempId + await waitFor(() => { + const it = result.current.uploadItems.find((x) => x.tempId === tempId) + expect(it?.status).toBe("error") + expect(it?.failedStage).toBe("complete") + // 卡片同时展示真实失败原因与"重试不会重新上传"提示 + expect(it?.error).toContain("complete timeout") + expect(it?.error).toContain("不会重新上传文件") + }) + + // 点重试:pump 复用 handle,只再调一次 complete(transfer/prepare 不重复) + await act(async () => { + result.current.retryUpload(tempId) + }) + await waitFor(() => expect(handles[0].complete).toHaveBeenCalledTimes(2)) + expect(handles.length).toBe(1) // 没有重新 prepare + expect(handles[0].transfer).toHaveBeenCalledTimes(1) // 没有重新直传 + + // 第二次 complete 成功 + await act(async () => { + handles[0].completeCalls[1]?.resolve() + }) + await waitFor(() => { + expect(result.current.uploadItems.find((x) => x.tempId === tempId)?.status).toBe("done") + }) + }) + + it("prepare 阶段失败:标记 prepare 阶段并保留后端错误明细", async () => { + ;(prepareDirectUploadHandle as unknown as ReturnType).mockRejectedValueOnce({ + isAxiosError: true, + response: { status: 500, data: { detail: "签名服务内部错误" } }, + message: "Request failed with status code 500", + }) + + const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), { + wrapper: createWrapper(), + }) + + await act(async () => { + result.current.enqueueUploads([mp4("prep-fail.mp4")]) + }) + await waitFor(() => expect(result.current.uploadItems[0]?.status).toBe("error")) + const it = result.current.uploadItems[0] + expect(it.failedStage).toBe("prepare") + expect(it.error).toContain("签名服务内部错误") + }) + it("prepare 返回 skip_transfer=true 时立即跳过 transfer+complete,标记 done+duplicated", async () => { + const h = makeFakeHandle({ id: "a-skip", prepareDedup: true }) + ;(prepareDirectUploadHandle as unknown as ReturnType).mockImplementation( + async () => h, + ) + + const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), { + wrapper: createWrapper(), + }) + + await act(async () => { + result.current.enqueueUploads([mp4("skip-transfer.mp4")]) + }) + + await waitFor(() => { + expect(h.transfer).not.toHaveBeenCalled() + expect(h.complete).not.toHaveBeenCalled() + const it = result.current.uploadItems[0] + expect(it?.status).toBe("done") + expect(it?.duplicated).toBe(true) + expect(it?.assetId).toBe("a-skip") + }) + }) }) diff --git a/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx b/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx new file mode 100644 index 000000000..c9cadc752 --- /dev/null +++ b/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx @@ -0,0 +1,144 @@ +import { describe, expect, it, vi, beforeEach, afterEach } from "vitest" +import { render, screen, waitFor, cleanup } from "@testing-library/react" +import { MemoryRouter } from "react-router-dom" +import WechatBindCallback from "@/pages/auth/WechatBindCallback" + +const mockNavigate = vi.fn() +const mockSetUser = vi.fn() +const mockParams = new URLSearchParams({ code: "bind_code", state: "bind_state" }) +const mockSearchParams = [mockParams] as const + +const localStorageStore: Record = {} +vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => localStorageStore[key] || null) +vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => { + localStorageStore[key] = val +}) +vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => { + delete localStorageStore[key] +}) + +let bindError: unknown = null +const mockBindResult = { user: { id: "u1", wechat_bound: true } } + +vi.mock("react-router-dom", async () => { + const actual = await vi.importActual("react-router-dom") + return { + ...actual, + useNavigate: () => mockNavigate, + useSearchParams: () => mockSearchParams, + } +}) + +vi.mock("@/api/auth", () => ({ + bindWechat: vi.fn(async () => { + if (bindError) throw bindError + return mockBindResult + }), + normalizeUser: (u: unknown) => u, +})) + +vi.mock("@/store/authStore", () => ({ + useAuthStore: (selector: (state: unknown) => unknown) => selector({ setUser: mockSetUser }), +})) + +// iframe 场景:默认非 iframe;用例可 mockReturnValue(true) +const { mockIsInIframe, mockPostResult } = vi.hoisted(() => ({ + mockIsInIframe: vi.fn(() => false), + mockPostResult: vi.fn(), +})) +vi.mock("@/components/auth/WechatQrModal/messages", () => ({ + isInIframe: () => mockIsInIframe(), + postWechatQrResult: (...args: unknown[]) => mockPostResult(...args), +})) + +const renderPage = () => + render( + + + , + ) + +describe("WechatBindCallback Page", () => { + afterEach(() => { + cleanup() + }) + + beforeEach(() => { + vi.clearAllMocks() + mockIsInIframe.mockReturnValue(false) + bindError = null + Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k)) + mockParams.set("code", "bind_code") + mockParams.set("state", "bind_state") + localStorageStore.wechat_bind_state = "bind_state" + }) + + it("绑定成功跳转设置页并携带 success 标记", async () => { + renderPage() + await waitFor(() => { + expect(mockNavigate).toHaveBeenCalledWith("/app/profile?wechat_bind=success", { + replace: true, + }) + }) + expect(mockSetUser).toHaveBeenCalled() + }) + + it("本地无 wechat_bind_state(微信内/跨浏览器)不再误杀,绑定正常完成", async () => { + delete localStorageStore.wechat_bind_state + renderPage() + await waitFor(() => { + expect(mockNavigate).toHaveBeenCalledWith("/app/profile?wechat_bind=success", { + replace: true, + }) + }) + }) + + it("后端报错(微信已被其他账号绑定)时页面透传真实原因,不静默跳走", async () => { + bindError = { + isAxiosError: true, + response: { status: 409, data: { detail: "该微信已绑定其他账号" } }, + message: "Request failed with status code 409", + } + renderPage() + await waitFor(() => { + expect(screen.getByText(/该微信已绑定其他账号/)).toBeTruthy() + }) + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("缺少 code/state 时提示无效回调", async () => { + mockParams.delete("code") + renderPage() + await waitFor(() => { + expect(screen.getByText(/无效的回调参数/)).toBeTruthy() + }) + }) + + describe("iframe(弹窗内嵌二维码)场景", () => { + it("绑定成功时 postMessage 通知父窗口,不做 navigate", async () => { + mockIsInIframe.mockReturnValue(true) + renderPage() + await waitFor(() => { + expect(mockPostResult).toHaveBeenCalledWith("bind", true) + }) + expect(mockSetUser).toHaveBeenCalled() + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("绑定失败时把真实原因 postMessage 给父窗口", async () => { + mockIsInIframe.mockReturnValue(true) + bindError = { + isAxiosError: true, + response: { status: 409, data: { detail: "该微信已绑定其他账号" } }, + } + renderPage() + await waitFor(() => { + expect(mockPostResult).toHaveBeenCalledWith("bind", false, { + detail: expect.stringContaining("该微信已绑定其他账号"), + }) + }) + expect(screen.queryByText(/返回设置/)).toBeNull() + expect(mockNavigate).not.toHaveBeenCalled() + }) + }) +}) diff --git a/apps/web/src/test/pages/auth/WechatCallback.test.tsx b/apps/web/src/test/pages/auth/WechatCallback.test.tsx index fe558f930..7aa16f3c4 100644 --- a/apps/web/src/test/pages/auth/WechatCallback.test.tsx +++ b/apps/web/src/test/pages/auth/WechatCallback.test.tsx @@ -1,79 +1,223 @@ -import { describe, expect, it, vi, beforeEach } from "vitest" -import { render, screen } from "@testing-library/react" +import { describe, expect, it, vi, beforeEach, afterEach } from "vitest" +import { render, screen, waitFor, cleanup } from "@testing-library/react" import { MemoryRouter } from "react-router-dom" import WechatCallback from "@/pages/auth/WechatCallback" +const mockNavigate = vi.fn() +const mockSetAuth = vi.fn() + +// useSearchParams 返回模块级稳定引用(数组元素同一 URLSearchParams 实例), +// 避免每次 render 返回新数组/新实例导致 useEffect 依赖变化重跑 +const mockParams = new URLSearchParams({ code: "test_code", state: "test_state" }) +const mockSearchParams = [mockParams] as const +const mockAuthState = { setAuth: mockSetAuth } + +// 文件级 localStorage mock(避免每个用例重复 spy 导致链式污染) +const localStorageStore: Record = {} +vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => localStorageStore[key] || null) +vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => { + localStorageStore[key] = val +}) +vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => { + delete localStorageStore[key] +}) + +let mockCallbackResult: Record = {} +let mockCurrentUser: Record = {} +let callbackError: unknown = null + vi.mock("react-router-dom", async () => { const actual = await vi.importActual("react-router-dom") return { ...actual, - useNavigate: () => vi.fn(), - useSearchParams: () => [new URLSearchParams({ code: "test_code", state: "test_state" })], + useNavigate: () => mockNavigate, + useSearchParams: () => mockSearchParams, } }) vi.mock("@/api/auth", () => ({ - wechatCallback: vi.fn(() => new Promise(() => {})), // pending promise,保持loading - getCurrentUser: vi.fn(), + wechatCallback: vi.fn(async () => { + if (callbackError) throw callbackError + return mockCallbackResult + }), + getCurrentUser: vi.fn(async () => mockCurrentUser), normalizeUser: (u: unknown) => u, })) +vi.mock("@/api/auth/tokenRefresh", () => ({ + scheduleProactiveRefresh: vi.fn(), + cancelProactiveRefresh: vi.fn(), +})) + vi.mock("@/store/authStore", () => ({ - useAuthStore: () => ({ - setAuth: vi.fn(), - }), + useAuthStore: (selector: (state: unknown) => unknown) => selector({ setAuth: mockSetAuth }), })) -vi.mock("@/components/auth/BindContactModal", () => ({ - default: ({ open }: { open: boolean }) => ( -
- BindContactModal -
- ), +// iframe 场景:默认非 iframe;用例可 mockReturnValue(true) +const { mockIsInIframe, mockPostResult } = vi.hoisted(() => ({ + mockIsInIframe: vi.fn(() => false), + mockPostResult: vi.fn(), +})) +vi.mock("@/components/auth/WechatQrModal/messages", () => ({ + isInIframe: () => mockIsInIframe(), + postWechatQrResult: (...args: unknown[]) => mockPostResult(...args), })) -vi.mock("antd", async () => { - const actual = await vi.importActual("antd") - return { - ...actual, - message: { - success: vi.fn(), - error: vi.fn(), - }, - } -}) +const renderPage = () => + render( + + + , + ) describe("WechatCallback Page", () => { + afterEach(() => { + cleanup() + }) + beforeEach(() => { - // mock localStorage,设置wechat_state匹配,让校验通过 - const store: Record = { - wechat_state: "test_state", + vi.clearAllMocks() + mockIsInIframe.mockReturnValue(false) + callbackError = null + // 默认正常回调参数;用例可改写 mockParams 模拟 error 重定向 + Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k)) + mockParams.set("code", "test_code") + mockParams.set("state", "test_state") + localStorageStore.wechat_state = "test_state" + mockCallbackResult = { + access_token: "at", + refresh_token: "rt", + is_new_user: false, } - vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => store[key] || null) - vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => { - store[key] = val + mockCurrentUser = { + id: "u1", + display_name: "老用户", + profile_completed: true, + } + }) + + it("老用户登录成功跳转首页/来源页", async () => { + renderPage() + await waitFor(() => { + expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true }) }) - vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => { - delete store[key] + expect(mockSetAuth).toHaveBeenCalled() + }) + + it("新用户(is_new_user)跳转昵称引导页", async () => { + mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: true } + mockCurrentUser = { id: "u2", display_name: "微信用户", profile_completed: false } + renderPage() + await waitFor(() => { + expect(mockNavigate).toHaveBeenCalledWith("/welcome/wechat", { replace: true }) }) }) - it("should render without crashing", () => { - const { container } = render( - - - , - ) - expect(container).toBeTruthy() + it("is_new_user=false 但 profile_completed=false(上次中断)也跳引导页", async () => { + mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: false } + mockCurrentUser = { id: "u3", display_name: "微信用户", profile_completed: false } + renderPage() + await waitFor(() => { + expect(mockNavigate).toHaveBeenCalledWith("/welcome/wechat", { replace: true }) + }) }) - it("should show loading state while processing", () => { - render( - - - , - ) - // wechatCallback 返回 pending promise,所以应该显示 loading - expect(screen.getByText("正在登录...")).toBeTruthy() + it("本地无 wechat_state(微信内打开/跨浏览器场景)不再误杀,正常完成登录", async () => { + delete localStorageStore.wechat_state + renderPage() + await waitFor(() => { + expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true }) + }) + // state 已被清理 + expect(localStorageStore.wechat_state).toBeUndefined() + }) + + it("后端返回 detail 错误时,页面透传真实原因(不再吞成通用提示)", async () => { + callbackError = { + isAxiosError: true, + response: { status: 400, data: { detail: "微信授权码已过期,请重新扫码" } }, + message: "Request failed with status code 400", + } + renderPage() + await waitFor(() => { + expect(screen.getByText(/微信授权码已过期,请重新扫码/)).toBeTruthy() + }) + expect(screen.queryByText(/^微信登录失败,请重试$/)).toBeNull() + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("微信重定向带 error(用户拒绝授权)时展示授权失败原因", async () => { + for (const k of Array.from(mockParams.keys())) mockParams.delete(k) + mockParams.set("error", "access_denied") + mockParams.set("error_description", "The+user+denied+the+request") + renderPage() + await waitFor(() => { + expect(screen.getByText(/微信授权失败/)).toBeTruthy() + expect(screen.getByText(/access_denied/)).toBeTruthy() + }) + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("缺少 code/state 参数时提示无效回调", async () => { + mockParams.delete("code") + renderPage() + await waitFor(() => { + expect(screen.getByText(/无效的回调参数/)).toBeTruthy() + }) + }) + + it("处理中显示 loading", () => { + renderPage() + expect(screen.getByText("微信登录中...")).toBeTruthy() + }) + + describe("iframe(弹窗内嵌二维码)场景", () => { + it("登录成功时 postMessage 通知父窗口(needOnboarding=false),不做 navigate", async () => { + mockIsInIframe.mockReturnValue(true) + renderPage() + await waitFor(() => { + expect(mockPostResult).toHaveBeenCalledWith("login", true, { needOnboarding: false }) + }) + expect(mockSetAuth).toHaveBeenCalled() + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("新用户成功时上报 needOnboarding=true", async () => { + mockIsInIframe.mockReturnValue(true) + mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: true } + renderPage() + await waitFor(() => { + expect(mockPostResult).toHaveBeenCalledWith("login", true, { needOnboarding: true }) + }) + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("后端报错时把真实原因 postMessage 给父窗口,页面不渲染错误/按钮", async () => { + mockIsInIframe.mockReturnValue(true) + callbackError = { + isAxiosError: true, + response: { status: 400, data: { detail: "state 已过期或已被使用" } }, + } + renderPage() + await waitFor(() => { + expect(mockPostResult).toHaveBeenCalledWith("login", false, { + detail: expect.stringContaining("state 已过期或已被使用"), + }) + }) + expect(screen.queryByText(/返回登录/)).toBeNull() + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("微信重定向 error(拒绝授权)在 iframe 内也上报父窗口", async () => { + mockIsInIframe.mockReturnValue(true) + for (const k of Array.from(mockParams.keys())) mockParams.delete(k) + mockParams.set("error", "access_denied") + renderPage() + await waitFor(() => { + expect(mockPostResult).toHaveBeenCalledWith("login", false, { + detail: expect.stringContaining("access_denied"), + }) + }) + }) }) }) diff --git a/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx new file mode 100644 index 000000000..548d7f52e --- /dev/null +++ b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx @@ -0,0 +1,186 @@ +import { describe, expect, it, vi, beforeEach, afterEach } from "vitest" +import { render, screen, fireEvent, waitFor, cleanup } from "@testing-library/react" +import { MemoryRouter } from "react-router-dom" +import { QueryClient, QueryClientProvider } from "@tanstack/react-query" +import WechatOnboarding from "@/pages/auth/WechatOnboarding" + +const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false }, mutations: { retry: false } }, +}) + +const mockNavigate = vi.fn() +const mockSetUser = vi.fn() +let updateProfileMock = vi.fn() + +vi.mock("react-router-dom", async () => { + const actual = await vi.importActual("react-router-dom") + return { ...actual, useNavigate: () => mockNavigate } +}) + +let authState: Record = {} +vi.mock("@/store/authStore", () => ({ + useAuthStore: (selector: (state: unknown) => unknown) => selector(authState), +})) + +vi.mock("@/api/auth", () => ({ + updateProfile: (data: { display_name: string }) => updateProfileMock(data), +})) + +vi.mock("antd", async () => { + const actual = await vi.importActual("antd") + return { ...actual, message: { success: vi.fn(), error: vi.fn() } } +}) + +const renderPage = () => + render( + + + + + , + ) + +describe("WechatOnboarding 昵称引导页", () => { + afterEach(() => { + cleanup() + }) + + beforeEach(() => { + vi.clearAllMocks() + authState = { + isAuthenticated: true, + user: { id: "u1", display_name: "", profile_completed: false }, + setUser: mockSetUser, + } + localStorage.setItem("access_token", "at") + updateProfileMock = vi.fn(async (data: { display_name: string }) => ({ + id: "u1", + display_name: data.display_name, + profile_completed: true, + })) + }) + + it("未登录时跳转登录页", () => { + authState = { + isAuthenticated: false, + user: null, + setUser: mockSetUser, + } + localStorage.removeItem("access_token") + renderPage() + expect(mockNavigate).not.toHaveBeenCalled() + // Navigate 组件渲染即生效;这里断言页面不含昵称表单 + expect(screen.queryByText("进入小虾智剪")).toBeNull() + }) + + it("资料已完善的用户跳 dashboard", () => { + authState = { + isAuthenticated: true, + user: { id: "u1", display_name: "已起名", profile_completed: true }, + setUser: mockSetUser, + } + renderPage() + expect(screen.queryByText("进入小虾智剪")).toBeNull() + }) + + it("昵称输入框不预填,必须用户自己输入", () => { + renderPage() + expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy() + expect((screen.getByPlaceholderText("请输入您的昵称") as HTMLInputElement).value).toBe("") + }) + + it("新用户可见昵称表单并能提交", async () => { + renderPage() + expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy() + + fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), { + target: { value: "小虾用户" }, + }) + fireEvent.click(screen.getByText("进入小虾智剪")) + + await waitFor(() => { + expect(updateProfileMock).toHaveBeenCalledWith({ display_name: "小虾用户" }) + }) + await waitFor(() => { + expect(mockSetUser).toHaveBeenCalled() + expect(mockNavigate).toHaveBeenCalledWith("/app/dashboard", { replace: true }) + }) + }) + + it("昵称为空时不允许提交(表单校验)", async () => { + renderPage() + fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), { + target: { value: " " }, + }) + fireEvent.click(screen.getByText("进入小虾智剪")) + // 等待表单校验 + await waitFor( + () => { + expect(updateProfileMock).not.toHaveBeenCalled() + }, + { timeout: 1000 }, + ) + }) + + it("连点提交按钮只触发一次请求(防重复提交)", async () => { + // mutation 挂起不立即完成,模拟慢网络下连续双击 + let resolveSubmit: (v: unknown) => void = () => {} + updateProfileMock = vi.fn( + () => + new Promise((resolve) => { + resolveSubmit = resolve + }), + ) + renderPage() + fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), { + target: { value: "小虾用户" }, + }) + const btn = screen.getByText("进入小虾智剪") + fireEvent.click(btn) + // 第一次点击后立即再点(此时重渲染/loading 可能还没生效) + fireEvent.click(btn) + fireEvent.click(btn) + await waitFor(() => { + expect(updateProfileMock).toHaveBeenCalledTimes(1) + }) + // 释放挂起的 Promise,避免泄漏 + resolveSubmit({ id: "u1", display_name: "小虾用户", profile_completed: true }) + }) + + it("提交失败后守卫复位,允许再次提交", async () => { + updateProfileMock = vi + .fn() + .mockRejectedValueOnce({ + isAxiosError: true, + response: { status: 500, data: { detail: "服务内部错误" } }, + }) + .mockResolvedValueOnce({ id: "u1", display_name: "小虾用户", profile_completed: true }) + renderPage() + fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), { + target: { value: "小虾用户" }, + }) + fireEvent.click(screen.getByText("进入小虾智剪")) + await waitFor(() => { + expect(updateProfileMock).toHaveBeenCalledTimes(1) + }) + // 失败后再点一次,应能重新提交 + fireEvent.click(screen.getByText("进入小虾智剪")) + await waitFor(() => { + expect(updateProfileMock).toHaveBeenCalledTimes(2) + }) + }) + + it("提交失败显示错误且不跳转", async () => { + updateProfileMock = vi.fn(async () => { + throw new Error("500") + }) + renderPage() + fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), { + target: { value: "小虾用户" }, + }) + fireEvent.click(screen.getByText("进入小虾智剪")) + await waitFor(() => { + expect(mockNavigate).not.toHaveBeenCalled() + }) + }) +}) diff --git a/apps/web/src/test/pages/generate/smoke.test.tsx b/apps/web/src/test/pages/generate/smoke.test.tsx index 475514f67..9e58e0c6b 100755 --- a/apps/web/src/test/pages/generate/smoke.test.tsx +++ b/apps/web/src/test/pages/generate/smoke.test.tsx @@ -19,6 +19,10 @@ import "@/pages/generate/GeneratePage" import "@/pages/generate/components/Step2MaterialSelect" import "@/pages/generate/components/Step4TitleSettings" import "@/pages/generate/components/Step5VoiceSelect" +import "@/pages/generate/components/Step3VoiceWithMode" +import "@/pages/generate/components/CanvasPreviewGrid" +import "@/pages/generate/components/BatchGenerationGrid" +import "@/pages/generate/components/PreviewCountModal" import "@/pages/generate/components/PreviewVideoPanel" import "@/pages/generate/components/GenerateResultPanel" import "@/pages/generate/components/GenerateStepContent" @@ -45,6 +49,7 @@ describe("GeneratePage module smoke test", () => { }) }) import "@/pages/generate/hooks/useGenerateVideo" +import "@/pages/generate/hooks/useBatchCovers" import "@/pages/generate/hooks/usePreviewAssets" import "@/pages/generate/hooks/useSegmentScheduler" import "@/pages/generate/hooks/generate-video/useGenerationPolling" diff --git a/apps/web/src/test/pages/generate/title-library-autocomplete.test.tsx b/apps/web/src/test/pages/generate/title-library-autocomplete.test.tsx new file mode 100644 index 000000000..b8a05188d --- /dev/null +++ b/apps/web/src/test/pages/generate/title-library-autocomplete.test.tsx @@ -0,0 +1,126 @@ +/** + * TitleLibraryAutoComplete 单测(Issue #1737) + * + * 覆盖: + * - 聚焦空输入框 → 下拉立即展开,展示标题库全部标题(原生 AutoComplete 聚焦不展开,此为本工单核心修复) + * - 输入关键词 → 下拉只显示匹配项 + * - 点击下拉项 → onChange 回填所选标题 + * - 自由输入自定义标题 → onChange 正常透传,不被下拉干扰 + * - 标题库为空 → 聚焦不展开(不出"暂无数据"空壳) + * - 选中后下拉关闭 + */ +import { describe, it, expect, vi } from "vitest" +import { render, screen, waitFor, fireEvent } from "@testing-library/react" +import userEvent from "@testing-library/user-event" +import TitleLibraryAutoComplete from "@/pages/generate/components/title/TitleLibraryAutoComplete" + +const OPTIONS = [ + { label: "永康这家面馆绝了", value: "永康这家面馆绝了" }, + { label: "永康美食探店vlog", value: "永康美食探店vlog" }, + { label: "萌宠日常第一天", value: "萌宠日常第一天" }, +] + +function renderBox(initialValue = "", opts = OPTIONS) { + const onChange = vi.fn() + const result = render( + , + ) + return { onChange, ...result } +} + +/** 聚焦输入框(combobox role) */ +function focusInput() { + const input = screen.getByRole("combobox") as HTMLInputElement + fireEvent.focus(input) + return input +} + +/** 取下拉中实际可见的选项(rc-virtual-list 渲染为 .ant-select-item-option;role=option 的 listbox 是 a11y 哨兵) */ +function getVisibleOptions(): HTMLElement[] { + const dropdown = document.querySelector(".ant-select-dropdown:not(.ant-select-dropdown-hidden)") + if (!dropdown) return [] + return Array.from(dropdown.querySelectorAll(".ant-select-item-option")) as HTMLElement[] +} + +describe("TitleLibraryAutoComplete (#1737)", () => { + it("聚焦空输入框时下拉展开并展示标题库全部标题", async () => { + renderBox() + expect(screen.queryByRole("listbox")).not.toBeInTheDocument() + + focusInput() + + await screen.findByRole("listbox") + await waitFor(() => expect(getVisibleOptions()).toHaveLength(3)) + const options = getVisibleOptions() + expect(options[0]).toHaveTextContent("永康这家面馆绝了") + expect(options[2]).toHaveTextContent("萌宠日常第一天") + }) + + it("输入关键词时下拉只显示匹配项", async () => { + const user = userEvent.setup() + renderBox() + const input = screen.getByRole("combobox") + await user.click(input) + await screen.findByRole("listbox") + + await user.type(input, "永康") + await waitFor(() => expect(getVisibleOptions()).toHaveLength(2)) + const options = getVisibleOptions() + expect(options.every((o) => o.textContent?.includes("永康"))).toBe(true) + }) + + it("点击下拉项后 onChange 回填标题且下拉关闭", async () => { + const user = userEvent.setup() + const { onChange } = renderBox() + const input = screen.getByRole("combobox") as HTMLInputElement + await user.click(input) + await screen.findByRole("listbox") + + await user.click(screen.getByText("萌宠日常第一天")) + + await waitFor(() => { + expect(onChange).toHaveBeenCalledWith("萌宠日常第一天") + }) + await waitFor(() => { + expect(screen.queryByRole("listbox")).not.toBeInTheDocument() + }) + }) + + it("自由输入自定义标题时 onChange 正常透传(不被下拉干扰)", async () => { + const user = userEvent.setup() + const { onChange } = renderBox() + const input = screen.getByRole("combobox") + await user.click(input) + + await user.type(input, "我自己编的标题XYZ") + await waitFor(() => { + expect(onChange).toHaveBeenCalledWith("我自己编的标题XYZ") + }) + // 输入无匹配关键词,下拉无 option 时不阻塞输入 + expect(input).toHaveValue("我自己编的标题XYZ") + }) + + it("标题库为空时聚焦不展开下拉", async () => { + renderBox("", []) + focusInput() + // 等一帧确认没有 listbox + await new Promise((r) => setTimeout(r, 50)) + expect(screen.queryByRole("listbox")).not.toBeInTheDocument() + }) + + it("渲染下拉箭头图标作为可选择提示", () => { + const { container } = renderBox() + // antd 后缀图标在 .ant-select-arrow 内 + expect(container.querySelector(".ant-select-arrow")).toBeInTheDocument() + }) + + it("有初始值时输入框正常展示", () => { + renderBox("已有标题") + expect(screen.getByRole("combobox")).toHaveValue("已有标题") + }) +}) diff --git a/apps/web/src/test/router/lazyRoute.test.ts b/apps/web/src/test/router/lazyRoute.test.ts new file mode 100644 index 000000000..8e64fac7a --- /dev/null +++ b/apps/web/src/test/router/lazyRoute.test.ts @@ -0,0 +1,44 @@ +import { describe, it, expect, vi, afterEach } from "vitest" +import { lazyRoute } from "@/router/lazyRoute" + +const chunkErr = () => new TypeError("Failed to fetch dynamically imported module: /assets/x.js") + +/** fake 模块 */ +const Comp = function Comp() {} +const factoryOk = vi.fn(async () => ({ default: Comp })) + +afterEach(() => { + vi.clearAllMocks() +}) + +describe("lazyRoute", () => { + it("首次成功直接返回 Component", async () => { + const result = await lazyRoute(factoryOk)() + expect(result).toEqual({ Component: Comp }) + expect(factoryOk).toHaveBeenCalledTimes(1) + }) + + it("chunk 失败重试:前两次失败、第三次成功 → 不抛出", async () => { + const f = vi + .fn() + .mockRejectedValueOnce(chunkErr()) + .mockRejectedValueOnce(chunkErr()) + .mockResolvedValueOnce({ default: Comp }) + + const result = await lazyRoute(f as never)() + expect(result).toEqual({ Component: Comp }) + expect(f).toHaveBeenCalledTimes(3) + }) + + it("chunk 失败重试 2 次仍失败 → 抛出", async () => { + const f = vi.fn().mockRejectedValue(chunkErr()) + await expect(lazyRoute(f as never)()).rejects.toThrow(/dynamically imported/) + expect(f).toHaveBeenCalledTimes(3) + }) + + it("非 chunk 错误立即抛出,不重试", async () => { + const f = vi.fn().mockRejectedValue(new Error("业务模块内部报错")) + await expect(lazyRoute(f as never)()).rejects.toThrow("业务模块内部报错") + expect(f).toHaveBeenCalledTimes(1) + }) +}) diff --git a/apps/web/src/test/utils/chunkLoadError.test.ts b/apps/web/src/test/utils/chunkLoadError.test.ts new file mode 100644 index 000000000..b7a5b9997 --- /dev/null +++ b/apps/web/src/test/utils/chunkLoadError.test.ts @@ -0,0 +1,80 @@ +import { describe, it, expect, beforeEach, afterEach, vi } from "vitest" +import { + getChunkReloadedAt, + goHomeRecover, + isChunkLoadError, + reloadForChunkError, +} from "@/utils/chunkLoadError" + +describe("isChunkLoadError", () => { + it("识别 Vite 动态 import 失败", () => { + const err = new TypeError( + "Failed to fetch dynamically imported module: https://x/assets/AssetLibrary-abc.js", + ) + expect(isChunkLoadError(err)).toBe(true) + }) + + it("识别 Webpack 风格 ChunkLoadError", () => { + const err = new Error("Loading chunk 12 failed.") + err.name = "ChunkLoadError" + expect(isChunkLoadError(err)).toBe(true) + }) + + it("识别字符串形式错误", () => { + expect(isChunkLoadError("Error loading dynamically imported module")).toBe(true) + }) + + it("普通错误不命中", () => { + expect(isChunkLoadError(new Error("Cannot read properties of undefined"))).toBe(false) + expect(isChunkLoadError(null)).toBe(false) + expect(isChunkLoadError(undefined)).toBe(false) + expect(isChunkLoadError({ status: 500 })).toBe(false) + }) +}) + +describe("reload 标记", () => { + beforeEach(() => { + sessionStorage.clear() + // jsdom 未实现真实导航,reload 仅打 "not implemented" 警告,静默掉 + vi.spyOn(console, "error").mockImplementation(() => {}) + }) + afterEach(() => { + vi.restoreAllMocks() + sessionStorage.clear() + }) + + it("无标记返回 null", () => { + expect(getChunkReloadedAt()).toBeNull() + }) + + it("reloadForChunkError 写入刷新标记", () => { + expect(() => reloadForChunkError()).not.toThrow() + expect(getChunkReloadedAt()).not.toBeNull() + }) + + it("标记过期(>10min)返回 null", () => { + sessionStorage.setItem("chunk_error_reloaded_at", String(Date.now() - 11 * 60 * 1000)) + expect(getChunkReloadedAt()).toBeNull() + }) + + it("goHomeRecover 清掉标记", () => { + reloadForChunkError() + expect(getChunkReloadedAt()).not.toBeNull() + expect(() => goHomeRecover()).not.toThrow() + expect(sessionStorage.getItem("chunk_error_reloaded_at")).toBeNull() + }) + + it("sessionStorage 抛异常(无痕模式)时降级不崩溃", () => { + const spy = vi.spyOn(Storage.prototype, "getItem").mockImplementation(() => { + throw new Error("Storage disabled") + }) + const setSpy = vi.spyOn(Storage.prototype, "setItem").mockImplementation(() => { + throw new Error("Storage disabled") + }) + expect(getChunkReloadedAt()).toBeNull() + expect(() => reloadForChunkError()).not.toThrow() + expect(() => goHomeRecover()).not.toThrow() + spy.mockRestore() + setSpy.mockRestore() + }) +}) diff --git a/apps/web/src/utils/chunkLoadError.ts b/apps/web/src/utils/chunkLoadError.ts new file mode 100644 index 000000000..7dce61a66 --- /dev/null +++ b/apps/web/src/utils/chunkLoadError.ts @@ -0,0 +1,84 @@ +/** + * 发版后旧标签页懒加载 chunk 失效(白屏)的识别与恢复工具。 + * + * 背景:页面 React Router 的 lazy 动态 import,发版后旧 chunk 文件名被删除, + * 停留在旧标签页的用户点菜单时 import 404,抛出 + * "Failed to fetch dynamically imported module"(Vite)/ ChunkLoadError, + * 不捕获就是整页白屏。 + */ + +/** sessionStorage 标记:最近已经为 chunk 失效自动刷新过一次(带时间戳,10min 有效) */ +const RELOAD_FLAG_KEY = "chunk_error_reloaded_at" +/** 标记有效期:超过后允许再次自动刷新,避免用户手动正常刷新后标记永久残留 */ +const RELOAD_FLAG_TTL_MS = 10 * 60 * 1000 + +/** + * Storage 在 Safari 无痕模式 / 禁用 Cookie 的浏览器 / 严格 iframe 策略下 + * 访问可能抛异常;此处统一容错,拿不到存储就降级为"无标记",绝不能让 + * 错误边界本身因读存储而崩溃。 + */ +const safeStorage = { + getItem: (key: string): string | null => { + try { + return sessionStorage.getItem(key) + } catch { + return null + } + }, + setItem: (key: string, value: string): void => { + try { + sessionStorage.setItem(key, value) + } catch { + /* 存储不可用时静默降级:仅丢失"已刷新"标记,不影响恢复动作 */ + } + }, + removeItem: (key: string): void => { + try { + sessionStorage.removeItem(key) + } catch { + /* ignore */ + } + }, +} + +/** 判断错误是否为懒加载 chunk 加载失败(发版 404 / 网络中断 / 动态 import 失败) */ +export const isChunkLoadError = (error: unknown): boolean => { + if (!error) return false + // Vite: Failed to fetch dynamically imported module: /assets/xxx-yyy.js + // Webpack: ChunkLoadError: Loading chunk xxx failed. + const needle = + error instanceof Error + ? `${error.name} ${error.message}` + : typeof error === "string" + ? error + : "" + return /failed to fetch dynamically imported module|chunkloaderror|loading chunk \d+ failed|error loading dynamically imported module|importing a module script failed/i.test( + needle, + ) +} + +/** 读取上次自动刷新时间戳;过期或不存在返回 null */ +export const getChunkReloadedAt = (): number | null => { + const raw = safeStorage.getItem(RELOAD_FLAG_KEY) + if (!raw) return null + const ts = Number(raw) + if (!Number.isFinite(ts)) return null + if (Date.now() - ts > RELOAD_FLAG_TTL_MS) return null + return ts +} + +/** 标记"已为 chunk 失效自动刷新过",然后刷新页面 */ +export const reloadForChunkError = (): void => { + safeStorage.setItem(RELOAD_FLAG_KEY, String(Date.now())) + window.location.reload() +} + +/** + * 硬恢复:清掉标记后回到首页(整页导航,不是当前 URL 刷新)。 + * - chunk 失效兜底:回到首页会拉取最新 index.html,彻底脱离旧 chunk 引用 + * - 非 chunk 的页面级崩溃:跳首页能绕开当前报错路由,避免"刷新-再崩"死循环 + */ +export const goHomeRecover = (): void => { + safeStorage.removeItem(RELOAD_FLAG_KEY) + window.location.href = "/" +} diff --git a/apps/worker/video_processing/dedup.py b/apps/worker/video_processing/dedup.py index b4c6aaba3..7e6ebdf06 100755 --- a/apps/worker/video_processing/dedup.py +++ b/apps/worker/video_processing/dedup.py @@ -31,25 +31,65 @@ SCENE_CHANGE_THRESHOLD = 30 # 灰度差异阈值 MIN_KEYFRAME_INTERVAL_SEC = 1.0 # 最小关键帧间隔(秒) MAX_KEYFRAMES = 30 # 最大关键帧数 MIN_KEYFRAMES = 5 # 最小关键帧数 +FINGERPRINT_SAMPLE_INTERVAL_SEC = 1.0 # 指纹采样间隔(秒):密集均匀采样,保证两视频时序可对齐 +FINGERPRINT_MAX_SAMPLES = 30 # 长视频采样数上限(超过后采样间隔自动放宽) LONG_VIDEO_SEGMENT_SEC = 30 # 长视频每段秒数 LONG_VIDEO_DURATION_THRESHOLD_SEC = 180 # 3 分钟阈值 MIN_FRAMES_PER_SEGMENT = 2 # 长视频每段最少帧数 -# ── 滑动窗口匹配常量 ──────────────────────────────────────────── -SEGMENT_MATCH_THRESHOLD = 8 # 帧匹配汉明距离阈值 -MIN_CONSECUTIVE_MATCHES = 5 # 最少连续匹配帧数 +# ── 滑动窗口匹配常量(Issue #1702 二次校准) ───────────────────── +# 阈值经 staging 真实数据两轮回归校准(worker 容器内离线实验): +# 第一轮(2026-09-05):同源对 <=12 命中 4/11,异源最小距离 24 → 定 12; +# 第二轮(2026-09-05,证据视频 B->A 仍漏检):扩大样本到该用户全部 +# 15 个真实成片(13 个异源候选)实测: +# - 同源成片对(A 20s / B、C 各 11.75s,1s 密集采样): +# B->A 中位数距离 14,<=16 命中 8/11=0.73;C->A 8/11=0.73 +# - 异源成片对(13 个真实视频):每帧全局最近邻最小距离 18, +# <=16 命中帧数全部为 0(最近邻 18 仅个别帧,中位数 22~28) +# 12 漏掉同源降重对(降重滤镜/字幕/画面扰动把距离从 ~8 推到 14~16); +# 16 对同源命中 0.73+ 且与异源分布(最近邻 >=18)仍有 >=2bit 安全裕度, +# 异源 <=16 命中 0 帧,无误报空间。 +PHASH_THRESHOLD = 16 +SEGMENT_MATCH_THRESHOLD = PHASH_THRESHOLD # 片段匹配阈值与帧匹配统一(#1702:阈值常量统一来源) +MIN_CONSECUTIVE_MATCHES = 5 # 连续匹配默认门槛;短视频自适应 min(5, max(2, 分片数//2)) MAX_GAP = 2 # 允许的最大间隙帧数 +NEIGHBOR_WINDOW = 1 # 分片时序对齐:允许 ±1 邻接偏移(1s 密集采样下即 ±1s,缓解切点不一致) # ── 融合判定常量 ──────────────────────────────────────────────── PHASH_WEIGHT = 0.7 # pHash 权重 HISTOGRAM_WEIGHT = 0.3 # 直方图权重 -MATCH_RATIO_THRESHOLD = 0.7 # 至少 70% 帧匹配 +MATCH_RATIO_THRESHOLD = 0.7 # 全片重复(is_duplicate)至少 70% 帧匹配 +PARTIAL_COVERAGE_THRESHOLD = 0.5 # 局部复用覆盖率 >=50% 也判全片重复 DUPLICATE_THRESHOLD = 0.70 # 融合后相似度阈值 +# ── 降重裁剪规避常量(Issue #1702) ───────────────────────────── +# 成片强制 2-5% random_edge_crop 降重只服务外部平台;自查重指纹取中心 90% +# 区域,使两次不同裁剪的同源画面 pHash 距离回到同分布。 +FINGERPRINT_CENTER_CROP_RATIO = 0.90 + # ── 感知哈希 & 颜色直方图工具函数 ──────────────────────────────── +def center_crop_frame(image: np.ndarray, ratio: float = FINGERPRINT_CENTER_CROP_RATIO) -> np.ndarray: + """取画面中心 ratio 比例区域(裁除四边边缘)。 + + 查重指纹用:random_edge_crop 降重(2-5% 四边随机裁剪)会让同源画面 pHash + 位翻转 12-16,污染自查重(Issue #1702)。算 pHash/颜色直方图前先居中裁除 + 边缘 10%,两次不同裁剪的同源画面中心区域基本重合,指纹不再被降重污染。 + 降重只服务外部平台,不影响内部查重。 + """ + if image is None or image.size == 0: + return image + h, w = image.shape[:2] + ch, cw = int(h * ratio), int(w * ratio) + if ch <= 0 or cw <= 0 or (ch >= h and cw >= w): + return image + y0 = (h - ch) // 2 + x0 = (w - cw) // 2 + return image[y0 : y0 + ch, x0 : x0 + cw] + + def compute_phash(image: np.ndarray, hash_size: int = 8) -> str: """计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。 @@ -101,11 +141,17 @@ def hamming_distance(hash1: str, hash2: str) -> int: def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]: - """Compute color histogram for an image.""" + """Compute BGR color histogram for an image. + + Issue #1702: 每个通道独立做 NORM_L1 归一化(通道内 Σ=1,是概率分布), + 三通道拼接存储。Bhattacharyya 系数对拼接向量直接 Σ√(a*b) 会得到 + 3 通道之和(范围 [0,3],实测 ~14.9 是旧 L2 归一化的错误结果), + 消费方 _bhattacharyya_coefficient 按通道数平均归一到 [0,1]。 + """ hist = [] for i in range(3): h = cv2.calcHist([image], [i], None, [bins], [0, 256]) - h = cv2.normalize(h, h).flatten() + h = cv2.normalize(h, h, norm_type=cv2.NORM_L1).flatten() hist.extend(h) return hist @@ -210,6 +256,30 @@ def detect_keyframe_timestamps( return keyframe_times +def sample_fingerprint_timestamps( + duration: float, + *, + interval_sec: float = FINGERPRINT_SAMPLE_INTERVAL_SEC, + max_samples: int = FINGERPRINT_MAX_SAMPLES, +) -> list[float]: + """指纹采样时间戳:固定间隔密集均匀采样(Issue #1702)。 + + 动态场景检测抽帧(#1659)在两个同源视频上会各自取到不同时刻,切点/取帧 + 错位让对齐帧的 pHash 距离都很大(实测同源对最小距离 12 且配对时序错乱)。 + 改为固定 1s 间隔均匀采样后,复用片段的帧时刻天然对齐,配合 ±1 邻接窗口 + 即可检出同源/局部复用。长视频(>max_samples*interval)自动放宽间隔到 + duration/max_samples,保证分片数有上限。 + """ + if duration <= 0: + return [] + step = interval_sec + n_uniform = int(duration / step) + if n_uniform > max_samples: + step = duration / max_samples + count = max(1, int(duration / step)) + return [step * (i + 0.5) for i in range(count)] + + # ── 数据类 ────────────────────────────────────────────────────── @@ -297,23 +367,33 @@ def find_duplicate_segments( target_chunks: list, *, match_threshold: int = SEGMENT_MATCH_THRESHOLD, - min_consecutive: int = MIN_CONSECUTIVE_MATCHES, + min_consecutive: Optional[int] = None, max_gap: int = MAX_GAP, + neighbor_window: int = NEIGHBOR_WINDOW, ) -> list[DuplicateSegment]: - """滑动窗口时序匹配:找出两组分片之间的重复片段。 + """滑动窗口时序匹配:找出两组分片之间的重复片段(Issue #1702 重构)。 算法: - 1. 对每个 query chunk,找到 target 中汉明距离最小的 chunk - 2. 距离 <= match_threshold 视为匹配 - 3. 找连续匹配的 run(允许 max_gap 帧间隙) - 4. 连续匹配数 >= min_consecutive 的 run 报告为重复片段 + 1. 构建 query×target 全量汉明距离矩阵;每个 query chunk 保留所有 + 距离 <= match_threshold 的候选 target 分片(与帧匹配判定同一阈值)。 + 2. 时序一致贪心对齐:沿 query 时序推进,run 内优先选择与上一匹配帧 + 目标序号连贯(|delta| <= neighbor_window+1,允许 ±1 邻接/时序偏移 + 对齐——1s 密集采样下相邻帧 pHash 接近,最近邻在目标相邻帧间 + 正/反向跳变均属正常,缓解场景切割切点、取帧错位、局部倒退)的 + 候选;同距时偏好小索引(最早对齐位置)。 + 3. 连贯匹配中允许 <= max_gap 帧间隙桥接;断裂后另起新 run——天然 + 支持局部片段复用(复用片段可出现在任意时序位置,各成独立片段)。 + 4. 连续匹配帧数 >= min_consecutive 的 run 报为重复片段。短视频自适应: + min_consecutive = min(5, max(2, len(query_chunks)//2));n=1 时 + 不形成片段,由调用方匹配帧回退兜底。 Args: query_chunks: 查询视频的分片列表(FingerprintChunk 或 dict) target_chunks: 目标视频的分片列表 - match_threshold: 汉明距离匹配阈值 - min_consecutive: 最少连续匹配帧数 + match_threshold: 汉明距离匹配阈值(统一常量 PHASH_THRESHOLD) + min_consecutive: 最少连续匹配帧数;None 时按短视频自适应 max_gap: 允许的最大间隙帧数 + neighbor_window: 时序对齐允许的目标分片序号邻接窗口(正/反向均允许) Returns: DuplicateSegment 列表 @@ -321,95 +401,94 @@ def find_duplicate_segments( if not query_chunks or not target_chunks: return [] - def _get_phash(chunk) -> str: + def _get(chunk, key): if isinstance(chunk, dict): - return chunk["phash_binary"] - return chunk.phash_binary + return chunk[key] + return getattr(chunk, key) - def _get_start(chunk) -> int: - if isinstance(chunk, dict): - return chunk["start_time_ms"] - return chunk.start_time_ms + n, m = len(query_chunks), len(target_chunks) + q_ph = [_get(c, "phash_binary") for c in query_chunks] + t_ph = [_get(c, "phash_binary") for c in target_chunks] - def _get_end(chunk) -> int: - if isinstance(chunk, dict): - return chunk["end_time_ms"] - return chunk.end_time_ms + # Step 1: 全量距离矩阵。每个 query chunk 保留所有 <= 阈值的候选 target, + # 按距离升序;同距时小索引优先(取最早的对齐位置,贪心连贯推进时最保守, + # 不会越过复用片段末端;重复 hash 的连续帧由 Step 2 的连贯性窗口约束)。 + candidates: list[list[tuple[int, int]]] = [] # 每 query 帧: [(target_idx, dist), ...] + for i in range(n): + dists = [hamming_distance(q_ph[i], t_ph[j]) for j in range(m)] + cand = [(j, d) for j, d in enumerate(dists) if d <= match_threshold] + cand.sort(key=lambda x: (x[1], x[0])) + candidates.append(cand) - # Step 1: 逐帧匹配 - frame_matches: list[tuple[bool, int, int]] = [] # (is_match, min_dist, best_target_idx) - for qc in query_chunks: - qc_phash = _get_phash(qc) - best_dist = 64 - best_idx = 0 - for j, tc in enumerate(target_chunks): - d = hamming_distance(qc_phash, _get_phash(tc)) - if d < best_dist: - best_dist = d - best_idx = j - frame_matches.append((best_dist <= match_threshold, best_dist, best_idx)) + # 短视频自适应连续匹配门槛(Issue #1702 工单公式): + # MIN_CONSECUTIVE_MATCHES = min(5, max(2, 分片数//2))。 + # n=1 时门槛为 2 不形成片段,由 _evaluate_candidate 的匹配帧回退 + # (temporal_coverage 按匹配帧占比估计)兜底检出,不回归。 + if min_consecutive is None: + min_consecutive = min(MIN_CONSECUTIVE_MATCHES, max(2, n // 2)) - # Step 2: 找连续匹配的 runs - runs: list[tuple[int, int]] = [] # list of (start_idx, end_idx) - run_start = None + # Step 2: 时序一致贪心对齐。 + # run 内偏好与上一匹配帧目标序号连贯(|delta| <= neighbor_window+1, + # 支持 ±1 邻接窗口/时序偏移对齐,正反向抖动均允许)的候选; + # 无连贯候选时关闭旧 run。 + # 这天然支持局部片段复用:同一 query 视频中多个复用片段各自形成独立 run。 + frame_matches: list[tuple[bool, int, int]] = [] + runs: list[tuple[int, int]] = [] + run_start: Optional[int] = None + run_last_t: Optional[int] = None gap_count = 0 - for i, (is_match, _dist, _idx) in enumerate(frame_matches): - if is_match: + def _matching_count(a: int, b: int) -> int: + return sum(1 for k in range(a, b + 1) if frame_matches[k][0]) + + def _close_run(a: int, b: int) -> None: + if b >= a and _matching_count(a, b) >= min_consecutive: + runs.append((a, b)) + + for i in range(n): + cand = candidates[i] + if run_last_t is None: + chosen = cand[0] if cand else None + else: + chosen = next( + (c for c in cand if abs(c[0] - run_last_t) <= neighbor_window + 1), + None, + ) + + if chosen is not None: + tidx, dist = chosen + frame_matches.append((True, dist, tidx)) if run_start is None: run_start = i - gap_count = 0 # 重置间隙 + gap_count = 0 + run_last_t = tidx else: + frame_matches.append((False, match_threshold + 1, -1)) if run_start is not None: gap_count += 1 if gap_count > max_gap: - # 中断当前 run - run_end = i - gap_count # 最后一个匹配帧的索引 - # 计算 run 内的实际匹配帧数(总跨度 - 间隙数) - total_gaps = sum(1 for k in range(run_start, run_end + 1) if not frame_matches[k][0]) - matching_count = (run_end - run_start + 1) - total_gaps - if matching_count >= min_consecutive: - runs.append((run_start, run_end)) - run_start = None - gap_count = 0 + # 非匹配帧从 i-gap_count+1 开始,run 结束于其前一帧 + _close_run(run_start, i - gap_count) + run_start, run_last_t, gap_count = None, None, 0 - # 处理末尾 run if run_start is not None: - last_idx = len(frame_matches) - 1 - # 回退找到最后一个匹配帧的位置(跳过尾部非匹配帧) + last_idx = n - 1 while last_idx >= run_start and not frame_matches[last_idx][0]: last_idx -= 1 - if last_idx >= run_start: - # 计算 run 内的总间隙数 - total_gaps = sum(1 for k in range(run_start, last_idx + 1) if not frame_matches[k][0]) - matching_count = (last_idx - run_start + 1) - total_gaps - if matching_count >= min_consecutive: - runs.append((run_start, last_idx)) + _close_run(run_start, last_idx) # Step 3: 构建 DuplicateSegment segments: list[DuplicateSegment] = [] for start, end in runs: - query_start = _get_start(query_chunks[start]) - query_end = _get_end(query_chunks[end]) - - # 取目标范围(按最佳匹配的目标 chunk 时间范围) target_indices = [frame_matches[k][2] for k in range(start, end + 1) if frame_matches[k][0]] - if target_indices: - t_min = min(target_indices) - t_max = max(target_indices) - target_start = _get_start(target_chunks[t_min]) - target_end = _get_end(target_chunks[t_max]) - else: - target_start = _get_start(target_chunks[0]) - target_end = _get_end(target_chunks[-1]) - - avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1)) / (end - start + 1) + t_min, t_max = min(target_indices), max(target_indices) + avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1) if frame_matches[k][0]) / len(target_indices) segments.append( DuplicateSegment( - query_start_ms=query_start, - query_end_ms=query_end, - target_start_ms=target_start, - target_end_ms=target_end, + query_start_ms=_get(query_chunks[start], "start_time_ms"), + query_end_ms=_get(query_chunks[end], "end_time_ms"), + target_start_ms=_get(target_chunks[t_min], "start_time_ms"), + target_end_ms=_get(target_chunks[t_max], "end_time_ms"), avg_distance=avg_dist, ) ) @@ -423,7 +502,9 @@ def find_duplicate_segments( class VideoDeduplicator: """Video deduplication using multiple fingerprint methods.""" - PHASH_THRESHOLD = 8 # Issue #1658: pHash 汉明距离阈值由 10 收紧到 8,降低不同视频误判率 + # Issue #1702: 阈值统一来源为模块常量 PHASH_THRESHOLD(#1658 曾收紧到 8, + # 后经 staging 真实同源/异源指纹分布重新校准,见 test_phash_threshold_calibration_1702)。 + PHASH_THRESHOLD = PHASH_THRESHOLD HISTOGRAM_THRESHOLD = 0.85 @staticmethod @@ -447,27 +528,36 @@ class VideoDeduplicator: # 单帧不视为坏指纹(短视频或抽帧不足) if len(phashes) == 1: return False - # 多帧但所有 phash 完全相同 → 黑屏/纯色视频 + # Issue #1702: 旧逻辑"所有 phash 完全相同即判黑屏"会误杀短视频—— + # 11s 视频只有几个不同镜头时,相邻 1s 采样帧可能 phash 完全一致(内容 + # 连续但非黑屏)。黑屏的特征是「大量帧全部无内容」,要求至少 8 帧 + # 且相同帧占比 >=80% 才判坏;短视频(<8 帧)只有真正单值时交给 + # _bhattacharyya/融合分兜底,不因"帧都一样"直接跳过。 + if len(phashes) < 8: + return False unique = set(phashes) - if len(unique) == 1: + same_ratio = sum(1 for x in phashes if x == phashes[0]) / len(phashes) + if len(unique) == 1 and same_ratio >= 0.8: return True - # 多帧但所有 phash 之间的汉明距离都极小(<3)→ 近似黑屏 + # 多帧但所有唯一 phash 之间的汉明距离都极小(<3)且占比 >=80% → 近似黑屏 phash_list = list(unique) - if len(phash_list) >= 2: - all_distances = [] - for i in range(len(phash_list)): - for j in range(i + 1, len(phash_list)): - all_distances.append(hamming_distance(phash_list[i], phash_list[j])) + if len(phash_list) >= 2 and same_ratio >= 0.8: + all_distances = [ + hamming_distance(phash_list[i], phash_list[j]) + for i in range(len(phash_list)) + for j in range(i + 1, len(phash_list)) + ] if all_distances and max(all_distances) < 3: return True return False def compute_fingerprint(self, video_path: str) -> VideoFingerprint: - """Compute video fingerprint using dynamic keyframe detection. + """Compute video fingerprint using dense uniform sampling. - 使用 detect_keyframe_timestamps() 检测内容感知关键帧, - 在每个关键帧处取帧计算 pHash + color_histogram。 - 同时保留 MD5 计算和分片数据结构。 + Issue #1702: 使用 sample_fingerprint_timestamps() 固定 1s 间隔密集均匀 + 采样(替代动态场景检测抽帧),保证两个同源视频复用片段的帧时刻天然 + 对齐;每帧取中心 90% 区域(center_crop_frame)计算 pHash + color_histogram, + 绕开 random_edge_crop 降重裁剪污染;MD5 仍基于原始帧。 """ cap = cv2.VideoCapture(video_path) if not cap.isOpened(): @@ -481,8 +571,8 @@ class VideoDeduplicator: cap.release() - # 1. 检测关键帧时间戳 - keyframe_times = detect_keyframe_timestamps(video_path) + # 1. 固定间隔密集采样(Issue #1702:替代动态场景检测,保证跨视频时序对齐) + keyframe_times = sample_fingerprint_timestamps(duration) if not keyframe_times: return VideoFingerprint( @@ -506,12 +596,15 @@ class VideoDeduplicator: if not ret: continue - # MD5 计算 + # MD5 计算(基于原始帧,指纹文件级去重不受裁剪影响) _, buffer = cv2.imencode(".jpg", frame) md5_hash.update(buffer) - phash = compute_phash(frame) - hist = compute_color_histogram(frame) + # Issue #1702: pHash / 颜色直方图基于中心 90% 区域,绕开 random_edge_crop + # 降重裁剪对指纹的污染(降重只服务外部平台,不污染自查重)。 + fp_frame = center_crop_frame(frame) + phash = compute_phash(fp_frame) + hist = compute_color_histogram(fp_frame) # 计算分片时间范围(从前一个关键帧到下一个关键帧的中点) prev_boundary = keyframe_times[i - 1] * 1000 if i > 0 else 0 @@ -564,12 +657,22 @@ class VideoDeduplicator: @staticmethod def _bhattacharyya_coefficient(hist_a: list[float], hist_b: list[float]) -> float: - """Bhattacharyya 系数:Σ √(a[i] * b[i]),范围 [0, 1],1=完全相同。""" + """Bhattacharyya 系数(概率分布版,范围 [0,1],1=完全相同)。 + + Issue #1702: compute_color_histogram 输出 3 通道拼接、每通道独立 NORM_L1 + (单通道 Σ=1,三通道拼接向量 Σ=3)。旧实现直接 Σ√(a*b) 对三通道拼接向量 + 算出 ~3(旧 L2 归一化更是算出 ~14.9),不是合法的概率系数。 + 这里按两个直方图各自的总量归一:BC = Σ√(a*b) / √(Σa·Σb)。 + - 单通道概率分布(Σa=Σb=1):分母 1,与旧测试/教科书定义一致; + - 三通道拼接(Σa=Σb=3):分母 3,结果在 [0,1]。 + """ min_len = min(len(hist_a), len(hist_b)) - a = hist_a[:min_len] - b = hist_b[:min_len] - # 纯标准库计算(不依赖 numpy);max(0.0, ...) 防御上游异常负值导致 sqrt domain error - return float(sum(math.sqrt(max(0.0, ai * bi)) for ai, bi in zip(a, b, strict=False))) + a = [max(0.0, float(x)) for x in hist_a[:min_len]] + b = [max(0.0, float(x)) for x in hist_b[:min_len]] + # max(0.0, ...) 防御上游异常负值导致 sqrt domain error + coeff = sum(math.sqrt(ai * bi) for ai, bi in zip(a, b, strict=False)) + norm = math.sqrt(sum(a) * sum(b)) + return float(coeff / norm) if norm > 0 else 0.0 @staticmethod def _compute_histogram_similarity( @@ -611,6 +714,72 @@ class VideoDeduplicator: hist_similarity = VideoDeduplicator._compute_histogram_similarity(hist_a, hist_b) if hist_b else 0.5 return PHASH_WEIGHT * phash_similarity + HISTOGRAM_WEIGHT * hist_similarity + @staticmethod + def _evaluate_candidate( + fingerprint: VideoFingerprint, + existing_phashes: list[str], + existing_histograms: list, + existing_chunk_objects: list, + *, + query_duration_sec: float, + ) -> dict: + """评估新视频指纹与单个候选视频的相似度(Issue #1702 共享逻辑)。 + + 指标: + - min_distances / frame_match_rate:每个新分片到候选视频全局最近邻的汉明距离, + 分母取两视频分片数的较小值(支持局部片段复用:短视频复用长视频片段时不被长视频分母稀释)。 + - temporal_coverage:时序一致连续匹配片段总时长 / 新视频时长(局部复用主指标)。 + - fusion:pHash 中位数距离 + 颜色直方图的加权融合分。 + + Returns: + {frame_match_rate, temporal_coverage, segments, median_distance, + fusion, matching_frames, min_distances} + """ + query_phashes = fingerprint.keyframe_phashes or [] + if not query_phashes or not existing_phashes: + return { + "frame_match_rate": 0.0, + "temporal_coverage": 0.0, + "segments": [], + "median_distance": 64, + "fusion": 0.0, + "matching_frames": 0, + "min_distances": [], + } + + min_distances = [min(hamming_distance(ph, ep) for ep in existing_phashes) for ph in query_phashes] + matching_frames = sum(1 for d in min_distances if d <= PHASH_THRESHOLD) + # 分母取 min(两视频分片数):局部复用时(如 B 的 5 片复用 A 9 片中的若干片) + # 命中帧占比不因候选视频更长而被稀释。 + frame_match_rate = matching_frames / min(len(query_phashes), len(existing_phashes)) + + segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects) + duration_ms = query_duration_sec * 1000 if query_duration_sec else 0 + if duration_ms > 0 and segments: + covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments) + temporal_coverage = min(covered_ms / duration_ms, 1.0) + elif matching_frames > 0: + # 无连续片段(时序连贯性不足)时,按匹配帧占比估计覆盖: + # 密集 1s 采样下每个分片≈1s 等权时间片,匹配帧数≈命中秒数。 + temporal_coverage = min(frame_match_rate, 1.0) + else: + temporal_coverage = 0.0 + + median_distance = statistics.median(min_distances) if min_distances else 64 + fusion = VideoDeduplicator._compute_fusion_score( + median_distance, fingerprint.color_histograms, existing_histograms + ) + + return { + "frame_match_rate": frame_match_rate, + "temporal_coverage": temporal_coverage, + "segments": segments, + "median_distance": median_distance, + "fusion": fusion, + "matching_frames": matching_frames, + "min_distances": min_distances, + } + def check_duplicate( self, fingerprint: VideoFingerprint, @@ -620,6 +789,7 @@ class VideoDeduplicator: scope: str = "project", user_id: str = "", duration_sec: float = 0, + exclude_video_id: str | None = None, ) -> Optional[dict]: """检查视频是否与已有视频重复。 @@ -636,6 +806,9 @@ class VideoDeduplicator: scope: "project" 项目内查重(默认),"user" 跨项目全局查重 user_id: 用户 ID(scope="user" 时使用) duration_sec: 视频时长(秒),用于时长预过滤 ±15% + exclude_video_id: 排除的视频 ID(查重自身时用)。recompute-dedup + 重算时视频记录已存在,不排除会自匹配(距离 0 分最高)导致 + duplicate_of 指向自己(Issue #1702 连带修复)。 Returns: 重复信息字典(含 duplicate, duplicate_of, reason, similarity, duplicate_segments), @@ -643,13 +816,22 @@ class VideoDeduplicator: """ video_repo = SQLAlchemyGeneratedVideoRepository(session) if scope == "user" and user_id: - dur_min = duration_sec * 0.85 if duration_sec > 0 else 0 - dur_max = duration_sec * 1.15 if duration_sec > 0 else 0 - existing_videos = video_repo.list_by_user(user_id, duration_min=dur_min, duration_max=dur_max) + # Issue #1702: 不做 ±15% 时长预过滤。旧逻辑按 duration_sec 缩小候选窗口, + # 但局部片段复用的两个视频时长必然不同(证据视频 20s vs 11s,差 42%), + # ±15% 窗口让同源视频互相不可见 → is_duplicate 恒 False。 + # 全量遍历同用户视频(与 compute_duplicate_rate 口径一致),异源视频由 + # fusion/temporal_coverage 阈值天然过滤(校准:异源最小汉明距离 24)。 + existing_videos = video_repo.list_by_user(user_id) else: existing_videos = video_repo.list_by_project(project_id) + best_score = 0.0 + best_result: Optional[dict] = None + for existing in existing_videos: + # 排除自身(recompute 时当前视频已在候选列表里,否则自匹配距离 0 必最高分) + if exclude_video_id and existing.id == exclude_video_id: + continue if not existing.video_fingerprint: continue @@ -677,61 +859,70 @@ class VideoDeduplicator: if not existing_phashes: continue - # 计算每个新关键帧到已有关键帧的最小汉明距离 - min_distances = [] - for phash in fingerprint.keyframe_phashes: - distances = [hamming_distance(phash, ep) for ep in existing_phashes] - min_distances.append(min(distances)) - - # 帧匹配比例检查 - matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD) - match_ratio = matching_frames / len(min_distances) if min_distances else 0 - if match_ratio < MATCH_RATIO_THRESHOLD: - continue - - # 中位数距离 - median_distance = statistics.median(min_distances) if min_distances else 64 - if median_distance >= self.PHASH_THRESHOLD: - continue - - # 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表) + # 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表) if chunk_data: existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")] + existing_chunk_objects = chunk_data else: existing_histograms = ef.get("color_histograms") or [] + existing_chunk_objects = [ + {"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes + ] - combined_score = self._compute_fusion_score( - median_distance, fingerprint.color_histograms, existing_histograms + # Issue #1702: 统一评估每个候选(含局部片段复用),不再用 + # "frame_match_rate<0.7 整条跳过" 的硬门槛——局部复用(如 B 结尾 2s + # ≈ A 中间 2s)帧比例天然低,但 coverage 能检出。 + ev = self._evaluate_candidate( + fingerprint, + existing_phashes, + existing_histograms, + existing_chunk_objects, + query_duration_sec=fingerprint.duration, + ) + logger.debug( + "check_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f " + "temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d", + existing.id, + ev["min_distances"], + ev["frame_match_rate"], + ev["temporal_coverage"], + ev["median_distance"], + ev["fusion"], + len(ev["segments"]), ) - if combined_score < DUPLICATE_THRESHOLD: - continue - - # 滑动窗口时序匹配:获取具体重复片段 - existing_chunk_objects = ( - chunk_data - if chunk_data - else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes] + # 全片重复判定:融合分过阈 且(帧匹配比例 >=70% 或 局部覆盖 >=50%) + is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and ( + ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD ) - segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects) - - return { - "duplicate": True, - "duplicate_of": existing.id, - "reason": "phash_histogram_fusion", - "similarity": combined_score, - "duplicate_segments": [ - { - "query_start_ms": s.query_start_ms, - "query_end_ms": s.query_end_ms, - "target_start_ms": s.target_start_ms, - "target_end_ms": s.target_end_ms, - "avg_distance": round(s.avg_distance, 2), - } - for s in segments - ], - } + if is_full_duplicate and ev["fusion"] > best_score: + best_score = ev["fusion"] + best_result = { + "duplicate": True, + "duplicate_of": existing.id, + "reason": "phash_histogram_fusion", + "similarity": ev["fusion"], + "duplicate_segments": [ + { + "query_start_ms": s.query_start_ms, + "query_end_ms": s.query_end_ms, + "target_start_ms": s.target_start_ms, + "target_end_ms": s.target_end_ms, + "avg_distance": round(s.avg_distance, 2), + } + for s in ev["segments"] + ], + } + if best_result: + return best_result + logger.info( + "check_duplicate no match (project=%s scope=%s): %d candidates evaluated, best_fusion=%.3f", + project_id, + scope, + len(existing_videos), + best_score, + ) return None def check_batch_duplicate( @@ -763,6 +954,9 @@ class VideoDeduplicator: video_repo = SQLAlchemyGeneratedVideoRepository(session) batch_videos = video_repo.list_by_batch(batch_id) + best_score = 0.0 + best_result: Optional[dict] = None + for existing in batch_videos: if existing.id == current_video_id: continue @@ -796,59 +990,59 @@ class VideoDeduplicator: if not existing_phashes: continue - min_distances = [] - for phash in fingerprint.keyframe_phashes: - distances = [hamming_distance(phash, ep) for ep in existing_phashes] - min_distances.append(min(distances)) - - # 帧匹配比例检查 - matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD) - match_ratio = matching_frames / len(min_distances) if min_distances else 0 - if match_ratio < MATCH_RATIO_THRESHOLD: - continue - - median_distance = statistics.median(min_distances) if min_distances else 64 - if median_distance >= self.PHASH_THRESHOLD: - continue - - # 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表) if chunk_data: existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")] + existing_chunk_objects = chunk_data else: existing_histograms = ef.get("color_histograms") or [] + existing_chunk_objects = [ + {"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes + ] - combined_score = self._compute_fusion_score( - median_distance, fingerprint.color_histograms, existing_histograms + ev = self._evaluate_candidate( + fingerprint, + existing_phashes, + existing_histograms, + existing_chunk_objects, + query_duration_sec=fingerprint.duration, + ) + logger.debug( + "check_batch_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f " + "temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d", + existing.id, + ev["min_distances"], + ev["frame_match_rate"], + ev["temporal_coverage"], + ev["median_distance"], + ev["fusion"], + len(ev["segments"]), ) - if combined_score < DUPLICATE_THRESHOLD: - continue - - # 滑动窗口时序匹配 - existing_chunk_objects = ( - chunk_data - if chunk_data - else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes] + is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and ( + ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD ) - segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects) - - return { - "duplicate": True, - "duplicate_of": existing.id, - "reason": "batch_phash_histogram_fusion", - "similarity": combined_score, - "duplicate_segments": [ - { - "query_start_ms": s.query_start_ms, - "query_end_ms": s.query_end_ms, - "target_start_ms": s.target_start_ms, - "target_end_ms": s.target_end_ms, - "avg_distance": round(s.avg_distance, 2), - } - for s in segments - ], - } + if is_full_duplicate and ev["fusion"] > best_score: + best_score = ev["fusion"] + best_result = { + "duplicate": True, + "duplicate_of": existing.id, + "reason": "batch_phash_histogram_fusion", + "similarity": ev["fusion"], + "duplicate_segments": [ + { + "query_start_ms": s.query_start_ms, + "query_end_ms": s.query_end_ms, + "target_start_ms": s.target_start_ms, + "target_end_ms": s.target_end_ms, + "avg_distance": round(s.avg_distance, 2), + } + for s in ev["segments"] + ], + } + if best_result: + return best_result + logger.info("check_batch_duplicate no match (batch=%s): best_fusion=%.3f", batch_id, best_score) return None def compute_duplicate_rate( @@ -897,8 +1091,7 @@ class VideoDeduplicator: max_duplicate_rate = 0.0 max_visual_similarity = 0.0 match_count = 0 - - total_duration_ms = fingerprint.duration if fingerprint.duration else 0 + evaluated = 0 for existing in existing_videos: if current_video_id and existing.id == current_video_id: @@ -933,57 +1126,63 @@ class VideoDeduplicator: if not existing_phashes or not fingerprint.keyframe_phashes: continue - min_distances = [] - for phash in fingerprint.keyframe_phashes: - distances = [hamming_distance(phash, ep) for ep in existing_phashes] - min_distances.append(min(distances)) - - # frame_match_rate - total_frames = len(min_distances) - if total_frames == 0: - continue - matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD) - frame_match_rate = matching_frames / total_frames - - # 帧匹配比例太低则跳过 - if frame_match_rate < 0.3: - continue - - # temporal_coverage_rate via find_duplicate_segments - existing_chunk_objects = ( - chunk_data - if chunk_data - else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes] - ) - segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects) - - if total_duration_ms > 0 and segments: - covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments) - temporal_coverage_rate = min(covered_ms / total_duration_ms, 1.0) - else: - temporal_coverage_rate = 0.0 - - # duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage_rate - dup_rate = (frame_match_rate * 0.4 + temporal_coverage_rate * 0.6) * 100 - - # visual_similarity (融合相似度,归一化 0~1) - median_distance = statistics.median(min_distances) if min_distances else 64 + # 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表) if chunk_data: existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")] + existing_chunk_objects = chunk_data else: - # JSON NULL 显式回退空列表 existing_histograms = ef.get("color_histograms") or [] + existing_chunk_objects = [ + {"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes + ] - visual_sim = self._compute_fusion_score(median_distance, fingerprint.color_histograms, existing_histograms) + # Issue #1702: 统一评估;frame_match_rate 分母为 min(两视频分片数), + # temporal_coverage 时长量纲在 _evaluate_candidate 内统一为毫秒。 + ev = self._evaluate_candidate( + fingerprint, + existing_phashes, + existing_histograms, + existing_chunk_objects, + query_duration_sec=fingerprint.duration, + ) + evaluated += 1 + logger.debug( + "compute_duplicate_rate candidate=%s min_distances=%s frame_match_rate=%.3f " + "temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d", + existing.id, + ev["min_distances"], + ev["frame_match_rate"], + ev["temporal_coverage"], + ev["median_distance"], + ev["fusion"], + len(ev["segments"]), + ) - # 判定是否为重复(融合分数超过阈值) - if visual_sim >= DUPLICATE_THRESHOLD: + # Issue #1702: 去掉 "frame_match_rate<0.3 整条跳过" 硬门槛—— + # 局部片段复用帧比例天然低;coverage 为主指标,0 匹配自然得 0 分。 + # duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage + dup_rate = (min(ev["frame_match_rate"], 1.0) * 0.4 + ev["temporal_coverage"] * 0.6) * 100 + + # 全片重复计数与 check_duplicate 判定口径一致 + if ev["fusion"] >= DUPLICATE_THRESHOLD and ( + ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD + ): match_count += 1 if dup_rate > max_duplicate_rate: max_duplicate_rate = dup_rate - max_visual_similarity = visual_sim + max_visual_similarity = ev["fusion"] + logger.info( + "compute_duplicate_rate done (project=%s scope=%s): evaluated=%d max_rate=%.2f%% " + "max_visual_sim=%.3f matches=%d", + project_id, + scope, + evaluated, + max_duplicate_rate, + max_visual_similarity, + match_count, + ) return { "duplicate_rate": round(max(max_duplicate_rate, 0.0), 2), "visual_similarity": round(max_visual_similarity, 4), @@ -999,18 +1198,20 @@ def _save_fingerprint_chunks( session: Session, ) -> None: """将指纹分片数据批量写入 video_fingerprint_chunks 表。幂等:已有数据时跳过。""" - # 幂等检查:已有分片数据则跳过 - existing_count = ( - session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count() - ) - if existing_count > 0: - logger.debug("Fingerprint chunks already exist for video %s (%d chunks), skipping", video_id, existing_count) - return - if not fingerprint.chunks: logger.warning("No chunks in fingerprint for video %s, skipping chunk save", video_id) return + # Issue #1702: recompute-dedup 重算时指纹算法已变(中心裁剪 + 新阈值), + # 旧分片必须替换而非跳过(旧实现"有数据就跳过"导致重算不刷新分片表)。 + deleted = ( + session.query(VideoFingerprintChunkModel) + .filter(VideoFingerprintChunkModel.video_id == video_id) + .delete(synchronize_session=False) + ) + if deleted: + logger.info("Replaced %d stale fingerprint chunks for video %s", deleted, video_id) + chunk_models = fingerprint.to_chunk_models(video_id, project_id, user_id) session.bulk_save_objects(chunk_models) logger.info("Saved %d fingerprint chunks for video %s", len(chunk_models), video_id) @@ -1032,9 +1233,17 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict: raise ValueError(f"Generated video {generated_video_id} not found") local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4") - storage_service.download_file( - f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path - ) + # Issue #1702: recompute 走的是 OSS 重新下载路径(正常生成流程用本地渲染文件, + # 不经此任务)。成片真实 OSS key 是生成时的 + # generated/projects/{pid}/tasks/{task_id}/rendered_*.mp4(见 generation.py + # _upload_and_record),旧代码硬编码 projects/{pid}/generated/{vid}/{vid}.mp4 + # 这个从不存在的 key,导致所有 recompute 任务下载 404、查重数据永远无法重算。 + # 优先从 file_url 解析真实 key,旧 key 模式仅作回退。 + download_key = getattr(video, "file_url", "") or "" + if not download_key: + download_key = f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4" + logger.warning("video %s has no file_url, falling back to legacy key %s", generated_video_id, download_key) + storage_service.download_file(download_key, local_path) fingerprint = deduplicator.compute_fingerprint(local_path) @@ -1045,7 +1254,10 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict: session, scope="user", user_id=video.user_id, - duration_sec=fingerprint.duration / 1000 if fingerprint.duration else 0, + # Issue #1702: fingerprint.duration 单位已经是秒,旧代码 /1000 导致 + # ±15% 时长预过滤窗口缩到 ~0.013s,scope=user 的跨项目查重永远返回 None。 + duration_sec=fingerprint.duration if fingerprint.duration else 0, + exclude_video_id=generated_video_id, ) video.video_fingerprint = fingerprint.to_dict() diff --git a/apps/worker/video_processing/dedup_helpers.py b/apps/worker/video_processing/dedup_helpers.py index b7e6d4965..d0ee22da8 100755 --- a/apps/worker/video_processing/dedup_helpers.py +++ b/apps/worker/video_processing/dedup_helpers.py @@ -92,7 +92,8 @@ def create_video_record_and_dedup( logger.warning("Failed to save fingerprint chunks for %s: %s", video_id, chunk_err) # (a) 历史成片查重(跨项目全局 + 时长预过滤) - duration_sec = fingerprint.duration / 1000 if fingerprint.duration else 0 + # Issue #1702: fingerprint.duration 单位是秒,旧代码 /1000 让时长预过滤失效 + duration_sec = fingerprint.duration if fingerprint.duration else 0 duplicate_result = deduplicator.check_duplicate( fingerprint, project_id, @@ -100,6 +101,7 @@ def create_video_record_and_dedup( scope="user", user_id=user_id, duration_sec=duration_sec, + exclude_video_id=video_id, ) # (b) 批次内查重(仅当有 batch_id 时) diff --git a/apps/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py index d953c23f4..34d40ad09 100755 --- a/apps/worker/worker_app/celery_app.py +++ b/apps/worker/worker_app/celery_app.py @@ -6,6 +6,24 @@ celery_app = Celery(settings.worker_name) celery_app.conf.broker_url = settings.broker_url celery_app.conf.result_backend = settings.result_backend celery_app.conf.broker_connection_retry_on_startup = True + +# #1714 队列隔离:generation(高优,独占 worker)/ transcode(素材转码)/ celery(默认) +from packages.shared.celery_queues import ( # noqa: E402 + GENERATION_WORKER_PREFETCH_MULTIPLIER, + apply_queue_settings, +) + +apply_queue_settings(celery_app) +# 长渲染任务预取 1,避免任务被预取占住导致调度不均 +celery_app.conf.worker_prefetch_multiplier = GENERATION_WORKER_PREFETCH_MULTIPLIER +celery_app.conf.task_acks_late = True # worker 崩溃时未完成任务重回队列,由执行前守卫丢弃作废消息 +# worker 进程被 OOM/容器硬杀时拒绝 ack,消息留在队列由其他 worker 接手 +celery_app.conf.task_reject_on_worker_lost = True +# Redis broker 消息可见性超时(#1714):acks_late 下,消息被预取后 visibility_timeout +# 内未 ack 才会重投。长任务(ingest HEVC 转码 20-30 分钟、生成硬超时 11 分钟) +# 必须远大于最长执行时间,否则正常任务会在执行中被误重投;4 小时覆盖最长转码 + 余量。 +celery_app.conf.broker_transport_options = {"visibility_timeout": 4 * 60 * 60} + celery_app.conf.imports = ( "worker_app.tasks.health", "worker_app.tasks.ingest", @@ -22,10 +40,25 @@ celery_app.conf.imports = ( ) # Celery Beat 定时任务调度 +# 注:worker 单实例内嵌 beat(entrypoint-worker.sh -B),定时任务不会重复执行 celery_app.conf.beat_schedule = { + # pending 任务超时清理:worker 停止消费后,卡 pending 的任务 15 分钟内释放限流名额 "cleanup-stale-pending-tasks": { "task": "worker.cleanup_stale_pending_tasks", + "schedule": 300.0, # 每 5 分钟(秒) + "options": {"expires": 240}, # 4 分钟过期,避免堆积 + }, + # running 孤儿任务巡检:容器重启/进程被杀后卡 running 的任务,20 分钟无更新则判失败 + "cleanup-stale-running-tasks": { + "task": "worker.cleanup_stale_running_tasks", + "schedule": 300.0, # 每 5 分钟(秒) + "options": {"expires": 240}, + }, + # 上传/转码链路孤儿巡检:worker 重启丢 prefetch 消息后,卡 pending/processing + # 的 ingest_job + asset 占位超时标终态(#1714)。转码任务较长,10 分钟一轮 + "cleanup-stale-ingest-jobs": { + "task": "worker.cleanup_stale_ingest_jobs", "schedule": 600.0, # 每 10 分钟(秒) - "options": {"expires": 300}, # 5 分钟过期,避免堆积 + "options": {"expires": 540}, }, } diff --git a/apps/worker/worker_app/tasks/_startup.py b/apps/worker/worker_app/tasks/_startup.py index c841f4406..afe329704 100644 --- a/apps/worker/worker_app/tasks/_startup.py +++ b/apps/worker/worker_app/tasks/_startup.py @@ -7,11 +7,83 @@ from worker_app.db import SessionLocal logger = logging.getLogger(__name__) -# 孤儿任务超时阈值:渲染任务超过此时间未更新则视为卡死 -ORPHAN_TASK_TIMEOUT_MINUTES = 10 -# Pending 任务超时阈值:pending 任务在队列中等待超过此时间则自动清理 -PENDING_TASK_TIMEOUT_MINUTES = 30 +def cleanup_stale_running_with_session(repo, timeout_minutes: int) -> int: + """清理超时未更新的 running GenerationTask(可注入 repo 的纯核心,便于单测)。 + + Returns: + 清理的任务数量 + """ + return len(cleanup_stale_running_with_session_ids(repo, timeout_minutes)) + + +def cleanup_stale_running_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]: + """同 cleanup_stale_running_with_session,返回 [(task_id, celery_task_id), ...]。""" + fn = getattr(repo, "cleanup_stale_running_with_ids", None) + if fn is not None: + return fn(timeout_minutes) + # 旧仓储无 _with_ids 方法:降级为计数,无法撤销消息(执行前状态守卫兜底) + count = repo.cleanup_stale_running(timeout_minutes) + return [("", "") for _ in range(count)] + + +def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int: + """清理超时 pending GenerationTask(可注入 repo 的纯核心,便于单测)。 + + Returns: + 清理的任务数量 + """ + return len(cleanup_stale_pending_with_session_ids(repo, timeout_minutes)) + + +def cleanup_stale_pending_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]: + """同 cleanup_stale_pending_with_session,返回 [(task_id, celery_task_id), ...]。""" + fn = getattr(repo, "cleanup_stale_pending_with_ids", None) + if fn is not None: + return fn(timeout_minutes) + count = repo.cleanup_stale_pending(timeout_minutes) + return [("", "") for _ in range(count)] + + +def _revoke_and_purge_stale_messages(items: list[tuple[str, str]]) -> int: + """把清理掉的任务对应的 Celery 消息撤销并从 Redis 队列清除(#1714)。 + + 防止「DB 已标 failed,但队列消息还在 → 重投执行 → 非法状态转换 → 半成品」。 + 失败不阻断清理流程(执行前状态守卫是第二道防线)。 + """ + biz_ids = [tid for tid, _ in items if tid] + celery_ids = [cid for _, cid in items if cid] + if not biz_ids and not celery_ids: + return 0 + try: + from worker_app.celery_app import celery_app as app + from worker_app.core.config import get_settings + + from packages.shared.celery_orphan_guard import revoke_and_purge + + broker_url = get_settings().broker_url + return revoke_and_purge( + app, + broker_url, + business_task_ids=biz_ids, + celery_task_ids=celery_ids, + ) + except Exception as e: # noqa: BLE001 + logger.error("撤销作废任务队列消息失败(执行前守卫仍会兜底): %s", e, exc_info=True) + return 0 + + +# 孤儿任务超时阈值:running 任务超过此时间无进度更新则视为卡死。 +# 依据:worker.generate_video 硬超时 time_limit=11 分钟,正常任务不可能超过; +# 20 分钟阈值覆盖硬超时 + 重试 + 余量,绝不误杀正常任务。 +ORPHAN_TASK_TIMEOUT_MINUTES = 20 + +# Pending 任务超时阈值:任务创建后超过此时间仍未开始执行则判死。 +# 注意区分 running 孤儿阈值(20 分钟):pending 是「排队等待」时间, +# 队列积压(如 20+ 转码任务)时视频生成可能正常排队较久,阈值必须放宽, +# 避免正常排队任务被误杀。队列隔离(#1714)后 generation 队列独占 worker, +# 理论上排队极短;保留 45 分钟作为兜底,覆盖 worker 短暂停止消费的场景。 +PENDING_TASK_TIMEOUT_MINUTES = 45 def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover @@ -32,11 +104,16 @@ def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> try: session = SessionLocal() - repo = SQLAlchemyGenerationTaskRepository(session) - count = repo.cleanup_stale_running(timeout_minutes) - session.close() + try: + repo = SQLAlchemyGenerationTaskRepository(session) + items = cleanup_stale_running_with_session_ids(repo, timeout_minutes) + finally: + session.close() + count = len(items) if count > 0: logger.warning("清理了 %d 个超时的孤儿 GenerationTask(超过 %d 分钟未更新)", count, timeout_minutes) + purged = _revoke_and_purge_stale_messages(items) + logger.info("孤儿任务对应队列消息撤销/清除完成: %d 条", purged) else: logger.info("无孤儿 GenerationTask 需要清理") return count @@ -70,7 +147,9 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in .all() ) count = 0 + stale_items: list[tuple[str, str]] = [] for model in stale_jobs: + stale_items.append((model.id, getattr(model, "celery_task_id", "") or "")) model.status = JobStatus.FAILED.value model.error_message = f"任务执行中断(超过 {timeout_minutes} 分钟未更新)" count += 1 @@ -80,6 +159,8 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in else: logger.info("无孤儿 Job 需要清理") session.close() + if count > 0: + _revoke_and_purge_generation(stale_items) return count except Exception as e: logger.error("清理孤儿 Job 失败: %s", e, exc_info=True) @@ -105,9 +186,12 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU session = SessionLocal() try: repo = SQLAlchemyGenerationTaskRepository(session) - count = repo.cleanup_stale_pending(timeout_minutes) + items = cleanup_stale_pending_with_session_ids(repo, timeout_minutes) + count = len(items) if count > 0: logger.warning("清理了 %d 个超时的 pending GenerationTask(超过 %d 分钟未处理)", count, timeout_minutes) + purged = _revoke_and_purge_stale_messages(items) + logger.info("超时 pending 任务对应队列消息撤销/清除完成: %d 条", purged) else: logger.info("无超时 pending GenerationTask 需要清理") return count @@ -118,6 +202,31 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU session.close() +def _revoke_and_purge_generation(items: list[tuple[str, str]]) -> int: + """撤销 Job 表孤儿任务(TTS/配音等)的队列消息,队列覆盖全部已知队列。""" + biz_ids = [tid for tid, _ in items if tid] + celery_ids = [cid for _, cid in items if cid] + if not biz_ids and not celery_ids: + return 0 + try: + from worker_app.celery_app import celery_app as app + from worker_app.core.config import get_settings + + from packages.shared.celery_orphan_guard import revoke_and_purge + + broker_url = get_settings().broker_url + return revoke_and_purge( + app, + broker_url, + business_task_ids=biz_ids, + celery_task_ids=celery_ids, + queue_names=("generation", "transcode", "celery"), + ) + except Exception as e: # noqa: BLE001 + logger.error("撤销 Job 队列消息失败: %s", e, exc_info=True) + return 0 + + def cleanup_all_stale_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict: # pragma: no cover """统一清理所有超时的孤儿任务。 @@ -148,3 +257,32 @@ def _on_worker_ready(sender, **kwargs): # pragma: no cover result = cleanup_all_stale_tasks() total = result["generation_tasks"] + result["jobs"] logger.info("Worker 启动清理完成,共清理 %d 个孤儿任务", total) + + +@worker_ready.connect +def _recover_stuck_ingest_jobs_on_ready(sender, **kwargs): # pragma: no cover + """Worker 启动完成后恢复卡死在 processing 的 ingest_job(#1714)。 + + 容器重启/进程 OOM 导致 transcode 队列 unacked 消息未重投时,processing + ingest_job 会永久卡死。启动时扫描 processing 超 10 分钟的 job,CAS 重置 + pending 并重新派单;Redis 锁保证同容器 generation/transcode 双 worker + 只有一个执行恢复。旧消息若后来重投,ingest_asset 执行前守卫会丢弃。 + """ + try: + from packages.application.ingest_orphan_cleanup import ( + make_redis_recovery_lock, + recover_stuck_ingest_jobs_on_startup, + ) + + session = SessionLocal() + try: + recovered = recover_stuck_ingest_jobs_on_startup( + session, + lock_acquire=make_redis_recovery_lock(), + stuck_minutes=10, + ) + finally: + session.close() + logger.info("Worker 启动 ingest 恢复完成,共重新派单 %d 个卡死任务", recovered) + except Exception as e: # noqa: BLE001 — 启动恢复失败不能阻断 worker 起服 + logger.error("启动 ingest 恢复扫描失败(beat 巡检仍会兜底标 failed): %s", e, exc_info=True) diff --git a/apps/worker/worker_app/tasks/cleanup.py b/apps/worker/worker_app/tasks/cleanup.py index 28d5e7e22..ff29840dc 100644 --- a/apps/worker/worker_app/tasks/cleanup.py +++ b/apps/worker/worker_app/tasks/cleanup.py @@ -1,17 +1,27 @@ """定期清理任务 — Celery Beat 调度。 包含: -- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks +- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks(worker 停止消费时占位) +- cleanup_stale_running_tasks: 定期清理卡在 running 超时的 generation_tasks(容器重启/进程被杀后的孤儿) """ import logging from celery import shared_task from worker_app.tasks._startup import ( + ORPHAN_TASK_TIMEOUT_MINUTES, PENDING_TASK_TIMEOUT_MINUTES, + cleanup_orphan_tasks, + cleanup_stale_jobs, cleanup_stale_pending_tasks, ) +from packages.application.ingest_orphan_cleanup import ( + ASSET_ORPHAN_TIMEOUT_MINUTES, + INGEST_PENDING_TIMEOUT_MINUTES, + INGEST_PROCESSING_TIMEOUT_MINUTES, +) + logger = logging.getLogger(__name__) @@ -19,12 +29,14 @@ logger = logging.getLogger(__name__) def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINUTES) -> dict: """Celery Beat 调度的定期任务:清理超时的 pending 任务。 - 每 10 分钟执行一次(由 celery_app.py 的 beat_schedule 配置), + 每 5 分钟执行一次(由 celery_app.py 的 beat_schedule 配置), 查找所有 status='pending' 且 created_at < NOW() - timeout_minutes - 的 generation_tasks,批量更新为 failed。 + 的 generation_tasks,批量更新为 failed,释放限流名额;同时 revoke 并清除 + Redis 队列中对应的 Celery 消息,杜绝作废消息重投执行(#1714)。 Args: - timeout_minutes: 超时时间(分钟),默认 30 分钟 + timeout_minutes: 超时时间(分钟),默认 45 分钟(pending 排队阈值放宽, + 与 running 孤儿 20 分钟区分,避免正常排队任务被误杀) Returns: {"cleaned": int} @@ -33,3 +45,83 @@ def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_ if count > 0: logger.info("[Beat] 清理了 %d 个超时 pending 任务(超时阈值 %d 分钟)", count, timeout_minutes) return {"cleaned": count} + + +@shared_task(name="worker.cleanup_stale_running_tasks") +def scheduled_cleanup_stale_running(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict: + """Celery Beat 调度的定期任务:清理超时的 running 孤儿任务。 + + 每 5 分钟执行一次。worker_ready 信号只在 worker 启动时清一次, + 若 worker 没重启但任务卡死(上传挂起、进程 OOM 被内核杀掉等), + 任务会永久卡在 running 占位。此任务做持续兜底: + 查找 status='running' 且 updated_at < NOW() - timeout_minutes 的任务, + 标记为 failed(原因:容器重启/超时中断),同时清理 Job 表孤儿。 + + Args: + timeout_minutes: 超时时间(分钟),默认 20 分钟 + (worker.generate_video 硬超时 11 分钟,正常任务不可能超过 20 分钟) + + Returns: + {"generation_tasks": int, "jobs": int} + """ + gen_count = cleanup_orphan_tasks(timeout_minutes) + job_count = cleanup_stale_jobs(timeout_minutes) + total = gen_count + job_count + if total > 0: + logger.warning( + "[Beat] 清理孤儿任务: running GenerationTask=%d, Job=%d(超时阈值 %d 分钟)", + gen_count, + job_count, + timeout_minutes, + ) + return {"generation_tasks": gen_count, "jobs": job_count} + + +@shared_task(name="worker.cleanup_stale_ingest_jobs") +def scheduled_cleanup_stale_ingest_jobs( + processing_timeout_minutes: int = INGEST_PROCESSING_TIMEOUT_MINUTES, + pending_timeout_minutes: int = INGEST_PENDING_TIMEOUT_MINUTES, + orphan_asset_timeout_minutes: int = ASSET_ORPHAN_TIMEOUT_MINUTES, +) -> dict: + """Celery Beat 调度:清理上传/转码链路(IngestJob + Asset)孤儿记录。 + + 每 10 分钟执行一次。worker 容器重启/进程 OOM 时,已 prefetch 的 transcode + celery 消息会丢失(队列里也不存在),ingest_job 永久卡 pending/processing、 + asset 永久卡 processing/uploading,没有兜底永远不会恢复(#1714)。 + + - ingest_job processing > processing_timeout_minutes / pending > pending_timeout_minutes + → 标 failed;关联 asset 占位(processing/uploading)联动标 error + - 无 ingest_job 关联、created_at > orphan_asset_timeout_minutes 的占位 asset + → 标 error + - 作废 celery 消息 revoke + 物理清除(防重投,执行前守卫是第二道防线) + """ + from worker_app.db import SessionLocal + + from packages.application.ingest_orphan_cleanup import ( + cleanup_orphan_processing_assets, + cleanup_stale_ingest_jobs, + revoke_stale_ingest_messages, + ) + + session = SessionLocal() + try: + job_items, asset_ids = cleanup_stale_ingest_jobs( + session, + processing_timeout_minutes=processing_timeout_minutes, + pending_timeout_minutes=pending_timeout_minutes, + ) + orphan_asset_ids = cleanup_orphan_processing_assets(session, timeout_minutes=orphan_asset_timeout_minutes) + finally: + session.close() + + purged = revoke_stale_ingest_messages(job_items) if job_items else 0 + total_jobs = len(job_items) + total_assets = len(set(asset_ids) | set(orphan_asset_ids)) + if total_jobs or total_assets: + logger.warning( + "[Beat] 清理 ingest 链路孤儿: stale_jobs=%d, assets→error=%d, 队列清除消息=%d", + total_jobs, + total_assets, + purged, + ) + return {"stale_jobs": total_jobs, "assets_to_error": total_assets, "purged_messages": purged} diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 35e195bef..9a0a755e7 100644 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -22,6 +22,8 @@ from worker_app.celery_app import celery_app from worker_app.db import SessionLocal from worker_app.tasks.generation_plan_builder import build_error_info as _build_error_info +from packages.shared.celery_orphan_guard import TERMINAL_STATUS_VALUES + OUTPUT_WIDTH = 1280 OUTPUT_HEIGHT = 720 OUTPUT_FPS = 25.0 @@ -658,6 +660,34 @@ def generate_video(self, task_id: str) -> dict: finally: _session.close() + # ── 0. 执行前状态守卫(#1714):任务已被超时清理/孤儿恢复标记为终态时, + # 这是作废消息(worker 崩溃前未 ack 的旧消息重投/重复投递),直接丢弃, + # 不进入渲染,杜绝 failed→running 非法转换后继续跑产出半成品。 + if gen_task is not None and gen_task.status.value in TERMINAL_STATUS_VALUES: + logger.warning( + "[task_id=%s] 任务状态已为 %s,丢弃作废消息,不执行渲染", + task_id, + gen_task.status.value, + ) + return { + "status": "discarded", + "task_id": task_id, + "reason": f"task already terminal: {gen_task.status.value}", + } + + # 标记任务为 running —— 必须成功:状态机非法转换(如 failed→running)说明 + # 任务已被作废,安全中止,禁止继续执行。 + if not _update_task_status(task_id, "mark_processing"): + logger.error( + "[task_id=%s] 标记 running 失败(任务可能已被作废/取消),安全中止,不执行渲染", + task_id, + ) + return { + "status": "discarded", + "task_id": task_id, + "reason": "claim failed (invalid state transition)", + } + # 记录接收任务日志 if gen_task: gen_task.append_log( @@ -669,8 +699,6 @@ def generate_video(self, task_id: str) -> dict: ) _flush_logs(task_id, gen_task) - # 标记任务为 running - _update_task_status(task_id, "mark_processing") _update_task_progress(task_id, 10, "任务启动") try: diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index 7b48f0f45..4d7094952 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -359,6 +359,54 @@ def validate_transcode_output( return True +def _original_key_from_storage_key(storage_key: str) -> str: + """从可能被 HEVC 转码改写的 storage_key 还原原始 key。 + + 转码成功后 key 形如 uploads//IMG_2282_h264.MOV, + 占位 asset 以原始 key uploads//IMG_2282.MOV 创建。 + """ + if not storage_key: + return storage_key + _p = Path(storage_key) + if _p.stem.endswith("_h264"): + return str(_p.parent / (_p.stem[: -len("_h264")] + _p.suffix)) + return storage_key + + +def _resolve_placeholder_asset(asset_repo, job, original_storage_key): + """找到 complete 阶段创建的 PROCESSING 占位 asset(Issue #1714)。 + + HEVC 转码成功后 job.storage_key 会被改写为 *_h264 新 key,旧实现用新 key + 回查占位必然落空,进而兜底新建一条 READY 记录,导致原占位永久卡 processing。 + + 查找优先级: + 1. job.asset_id(complete 派单时透传的占位 id,最可靠,不依赖 key); + 2. 原始 storage_key(占位记录以原始 key 创建); + 3. 当前 job.storage_key(未转码/降级场景与原始 key 相同)。 + + 找不到返回 None(旧链路兼容,由调用方兜底新建并告警)。 + """ + asset_id = getattr(job, "asset_id", "") or "" + if asset_id: + try: + found = asset_repo.find_by_id(asset_id) + if found is not None: + return found + except Exception as find_err: + logger.warning("占位 asset 按 id 查询失败 asset_id=%s: %s", asset_id, find_err) + for key in (original_storage_key, getattr(job, "storage_key", "")): + if not key: + continue + try: + found = asset_repo.find_by_storage_key(key) + except Exception: + logger.warning("find_by_storage_key not available, trying fallback lookup") + found = None + if found is not None: + return found + return None + + @celery_app.task(name="worker.ingest_asset") def ingest_asset(job_id: str) -> dict: """ @@ -380,6 +428,23 @@ def ingest_asset(job_id: str) -> dict: if job is None: return {"status": "failed", "error": "job not found"} + # ── 执行前状态守卫(#1714):job 已终态(失败/完成)说明这是作废消息 + # (超时清理标记 failed 后旧消息重投、或重复投递),直接丢弃不执行, + # 避免重复转码、重复回写。processing 是本任务自己第一次置位前的旧消息 + # 极少见,保守起见也丢弃(processing 的占位由恢复流程处理)。 + current_status = job.status.value if hasattr(job.status, "value") else str(job.status) + if current_status in ("failed", "completed"): + logger.warning( + "[ingest job_id=%s] 任务状态已为 %s,丢弃作废消息,不执行转码", + job_id, + current_status, + ) + return {"status": "discarded", "job_id": job_id, "reason": f"job already terminal: {current_status}"} + + # 记录原始 storage_key:HEVC 转码成功后 job.storage_key 会改写为 *_h264, + # 而 complete 阶段的占位 asset 始终以原始 key 创建,关联回写必须保留它。 + original_storage_key = job.storage_key + # Update job status to PROCESSING job.status = IngestJobStatus.PROCESSING job.updated_at = datetime.now(timezone.utc) @@ -624,22 +689,41 @@ def ingest_asset(job_id: str) -> dict: error_reason, ) - asset = Asset.create( - project_id=job.project_id, - library_id=job.library_id, - name=filename, - storage_key=job.storage_key, - mime_type=mime_type, - metadata={"source": "upload", "ingest_error": error_reason}, - file_size=int(metadata.get("size_bytes", 0)), - duration=float(metadata.get("duration", 0)), - width=int(metadata.get("width", 0)), - height=int(metadata.get("height", 0)), - codec=metadata.get("codec") or None, - status=AssetStatus.ERROR, - file_hash=job.file_hash, - ) - asset_repo.create(asset) + placeholder = _resolve_placeholder_asset(asset_repo, job, original_storage_key) + if placeholder is not None: + # 回写占位记录:标 ERROR(Issue #1714:禁止新建第二条导致占位孤儿) + asset = placeholder + asset.mime_type = mime_type + asset.metadata = {"source": "upload", "ingest_error": error_reason} + asset.file_size = int(metadata.get("size_bytes", 0)) + asset.duration = float(metadata.get("duration", 0)) or None + asset.width = int(metadata.get("width", 0)) or None + asset.height = int(metadata.get("height", 0)) or None + codec_val = metadata.get("codec") + if codec_val: + asset.codec = str(codec_val) + asset.status = AssetStatus.ERROR + asset.updated_at = datetime.now(timezone.utc) + asset_repo.update(asset) + else: + # 旧链路兜底:无占位记录(如历史 job 重跑)才新建 + logger.warning("无效素材且未找到占位记录,兜底新建 ERROR asset: job_id=%s", job_id) + asset = Asset.create( + project_id=job.project_id, + library_id=job.library_id, + name=filename, + storage_key=job.storage_key, + mime_type=mime_type, + metadata={"source": "upload", "ingest_error": error_reason}, + file_size=int(metadata.get("size_bytes", 0)), + duration=float(metadata.get("duration", 0)), + width=int(metadata.get("width", 0)), + height=int(metadata.get("height", 0)), + codec=metadata.get("codec") or None, + status=AssetStatus.ERROR, + file_hash=job.file_hash, + ) + asset_repo.create(asset) # Update job status to FAILED job.status = IngestJobStatus.FAILED @@ -656,16 +740,19 @@ def ingest_asset(job_id: str) -> dict: "error": error_reason, } - # 查找已存在的 Asset 记录(由 API 端在上传完成时立即创建为 PROCESSING 状态) - existing_asset = None - try: - existing_asset = asset_repo.find_by_storage_key(job.storage_key) - except Exception: - logger.warning("find_by_storage_key not available, trying fallback lookup") + # 查找 complete 阶段创建的占位 Asset 记录(Issue #1714)。 + # 必须用原始 storage_key / job.asset_id 关联——HEVC 转码后 job.storage_key + # 已改写为 *_h264,用新 key 回查占位必然落空,旧实现因此兜底新建 READY 记录, + # 导致原 PROCESSING 占位永久卡住(每个 HEVC 视频产生两条记录)。 + existing_asset = _resolve_placeholder_asset(asset_repo, job, original_storage_key) if existing_asset is None: - # 兜底:如果 API 端没有预先创建 Asset(旧版本兼容),则创建新记录 - logger.info("No pre-created asset found for storage_key=%s, creating new", job.storage_key) + # 兜底:仅当确实没有占位记录(旧版本 API / 历史 job 重跑)才新建。 + logger.warning( + "No placeholder asset found for job_id=%s original_key=%s, creating new", + job_id, + original_storage_key, + ) metadata["source"] = "upload" asset = Asset.create( project_id=job.project_id, @@ -685,8 +772,13 @@ def ingest_asset(job_id: str) -> dict: ) asset_repo.create(asset) else: - # 更新已有的 Asset 记录,补充元数据并将状态改为 READY + # 更新占位记录:补充元数据、置 READY。转码成功时 storage_key 同步改写为 + # *_h264(播放/下载走转码产物),原始 key 记入 metadata 可溯源。 asset = existing_asset + if job.storage_key != asset.storage_key: + metadata["original_storage_key"] = asset.storage_key + metadata["hevc_transcoded"] = True + asset.storage_key = job.storage_key asset.mime_type = mime_type metadata["source"] = "upload" asset.metadata = metadata @@ -737,9 +829,28 @@ def ingest_asset(job_id: str) -> dict: job_repo.update(job) # 将上传时创建的占位 Asset(PROCESSING/UPLOADING)标记为 ERROR, - # 避免素材永远卡在中间状态 + # 避免素材永远卡在中间状态。转码可能已把 job.storage_key 改写为 + # *_h264,需用 asset_id / 原始 key 多路径关联占位(Issue #1714)。 try: - existing = asset_repo.find_by_storage_key(job.storage_key) + existing = None + _asset_id = getattr(job, "asset_id", "") or "" + if _asset_id: + try: + existing = asset_repo.find_by_id(_asset_id) + except Exception: + existing = None + if existing is None: + _candidate_keys = [ + _original_key_from_storage_key(job.storage_key), + job.storage_key, + ] + for _key in _candidate_keys: + try: + existing = asset_repo.find_by_storage_key(_key) + except Exception: + existing = None + if existing is not None: + break if existing and existing.status in ( AssetStatus.PROCESSING, AssetStatus.UPLOADING, diff --git a/deploy/configs/.env.production b/deploy/configs/.env.production index 55ab74ff3..f14279411 100644 --- a/deploy/configs/.env.production +++ b/deploy/configs/.env.production @@ -214,3 +214,10 @@ MEDIAKIT_TIMEOUT=60 # ==================== 监控(可选)==================== # Sentry DSN(取消注释并填入实际值以启用错误追踪) # SENTRY_DSN=${SENTRY_DSN} + + +# ==================== 微信开放平台 OAuth(网页扫码登录)==================== +# 回调域名:xiaoxiajianji.com(微信开放平台已配置) +WECHAT_OPEN_APP_ID=${WECHAT_APP_ID} +WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET} +WECHAT_OPEN_REDIRECT_URI=https://saas.xiaoxiajianji.com/auth/wechat/callback diff --git a/deploy/configs/.env.staging b/deploy/configs/.env.staging index 538417191..9e09b91eb 100644 --- a/deploy/configs/.env.staging +++ b/deploy/configs/.env.staging @@ -231,3 +231,10 @@ DASHSCOPE_API_KEY=${DASHSCOPE_API_KEY} MEDIAKIT_API_KEY=${MEDIAKIT_API_KEY} MEDIAKIT_BASE_URL=https://mediakit.cn-beijing.volces.com/api/v1 MEDIAKIT_TIMEOUT=60 + + +# ==================== 微信开放平台 OAuth(网页扫码登录)==================== +# 回调域名:xiaoxiajianji.com(微信开放平台已配置) +WECHAT_OPEN_APP_ID=${WECHAT_APP_ID} +WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET} +WECHAT_OPEN_REDIRECT_URI=https://staging.xiaoxiajianji.com/auth/wechat/callback diff --git a/deploy/configs/nginx-production.conf b/deploy/configs/nginx-production.conf index 70b4b1a02..1944f8f1f 100644 --- a/deploy/configs/nginx-production.conf +++ b/deploy/configs/nginx-production.conf @@ -14,6 +14,10 @@ server { # SPA routing - index.html 禁止缓存,确保每次获取最新版本 location / { try_files $uri /index.html; + # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache, + # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用; + # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable + add_header Cache-Control "no-cache" always; } # API proxy — Production 环境代理到 production API 容器 diff --git a/deploy/configs/nginx-staging.conf b/deploy/configs/nginx-staging.conf index cc6cc4ab9..9521dbb42 100644 --- a/deploy/configs/nginx-staging.conf +++ b/deploy/configs/nginx-staging.conf @@ -21,6 +21,10 @@ server { # SPA fallback location / { try_files $uri /index.html; + # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache, + # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用; + # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable + add_header Cache-Control "no-cache" always; } # API proxy — Staging 环境代理到 staging API 容器 diff --git a/infra/docker/compose.yml b/infra/docker/compose.yml index 91297553b..154e9fad1 100755 --- a/infra/docker/compose.yml +++ b/infra/docker/compose.yml @@ -115,6 +115,8 @@ services: APP_VERSION: ${APP_VERSION:-unknown} WORKER_CONCURRENCY: ${WORKER_CONCURRENCY:-4} WORKER_MAX_TASKS_PER_CHILD: ${WORKER_MAX_TASKS_PER_CHILD:-100} + # #1714 队列隔离:generation 队列独占 worker(默认并发 2),其余并发给转码 + GENERATION_CONCURRENCY: ${GENERATION_CONCURRENCY:-2} GENERATED_FILES_DIR: /app/generated GENERATED_FILES_URL_PREFIX: /generated-files PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com} @@ -128,11 +130,11 @@ services: # 健康检查配置 # 注:celery inspect ping 依赖 broker 连接,在容器内不可靠,改用进程检查 healthcheck: - test: ["CMD-SHELL", "grep -q celery /proc/1/cmdline || exit 1"] + test: ["CMD-SHELL", "pgrep -f 'celery.*worker' | head -n1 >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1"] interval: 30s timeout: 10s retries: 3 - start_period: 30s + start_period: 40s logging: *default-logging @@ -140,7 +142,8 @@ services: # 资源限制建议(生产环境建议启用) # ========================================= # 注意: Worker 需要处理视频,建议分配更多资源 - # 并发 4 时需要 4C8G 以上,确保视频渲染不 OOM + # #1714 队列隔离后容器内运行 generation + transcode 两个 worker 进程, + # 总并发 = WORKER_CONCURRENCY(默认 4),4C8G 以上确保视频渲染不 OOM deploy: resources: limits: diff --git a/infra/docker/deploy-production-registry.sh b/infra/docker/deploy-production-registry.sh index f06faec2a..664e71f6e 100755 --- a/infra/docker/deploy-production-registry.sh +++ b/infra/docker/deploy-production-registry.sh @@ -146,6 +146,7 @@ docker run -d \ -e APP_ENV=production \ -e APP_VERSION="$IMAGE_TAG" \ -e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \ + -e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \ -e WORKER_MAX_TASKS_PER_CHILD=100 \ -e GENERATED_FILES_DIR=/app/generated \ -e GENERATED_FILES_URL_PREFIX=/generated-files \ @@ -154,7 +155,7 @@ docker run -d \ --restart unless-stopped \ --cpus 2 \ --memory 2g \ - --health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \ + --health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \ --health-interval 30s \ --health-timeout 10s \ --health-retries 3 \ diff --git a/infra/docker/deploy-staging-registry.sh b/infra/docker/deploy-staging-registry.sh index 9d3f5c729..b444b9ad9 100755 --- a/infra/docker/deploy-staging-registry.sh +++ b/infra/docker/deploy-staging-registry.sh @@ -109,13 +109,14 @@ docker run -d \ -e APP_ENV=staging \ -e APP_VERSION="$IMAGE_TAG" \ -e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \ + -e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \ -e WORKER_MAX_TASKS_PER_CHILD=100 \ -e GENERATED_FILES_DIR=/app/generated \ -e GENERATED_FILES_URL_PREFIX=/generated-files \ -v "$GENERATED_DIR:/app/generated" \ --restart unless-stopped \ --label com.centurylinklabs.watchtower.enable=true \ - --health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \ + --health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \ --health-interval 30s \ --health-timeout 10s \ --health-retries 3 \ diff --git a/infra/docker/entrypoint-worker.sh b/infra/docker/entrypoint-worker.sh index f2e208d96..0a74b474a 100755 --- a/infra/docker/entrypoint-worker.sh +++ b/infra/docker/entrypoint-worker.sh @@ -1,18 +1,64 @@ #!/bin/bash -# Worker 启动脚本 — 支持 WORKER_CONCURRENCY 环境变量 -# 未设置时默认 2(保持向后兼容) +# Worker 启动脚本 — #1714 队列隔离 +# +# 部署约束:worker 容器单实例(replicas=1),容器内启动两个 celery 进程: +# 1. generation-worker:独占消费 generation 队列(用户视频生成,高优先级), +# 内嵌 celery beat(-B),定时清理任务只在一个进程里跑,避免重复执行; +# 2. transcode-worker:消费 transcode + celery 默认队列(素材转码/分类/查重/ +# 配音/下载等后台任务)。 +# 转码队列积压时,generation 队列仍有独立 worker 立即领取视频生成任务。 +# +# 环境变量: +# WORKER_CONCURRENCY 总并发槽参考(默认 4);生成 worker 并发默认 2, +# 可用 GENERATION_CONCURRENCY 覆盖 +# GENERATION_CONCURRENCY generation worker 并发(默认 2) +# TRANSCODE_CONCURRENCY transcode worker 并发(默认 = WORKER_CONCURRENCY - 2,最小 1) +# WORKER_MAX_TASKS_PER_CHILD 每个子进程最大任务数(默认 100) set -e -CONCURRENCY="${WORKER_CONCURRENCY:-2}" +CONCURRENCY="${WORKER_CONCURRENCY:-4}" +MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}" -# ⚠️ 部署约束:此 Worker 必须且只能运行单实例(replicas=1) -# -B 标志嵌入 celery beat,beat 负责定期触发 pending 超时清理等定时任务 -# 多实例部署会导致每个 Worker 独立运行 Beat,造成定时任务重复执行 -# 若需横向扩展 Worker,必须将 Beat 拆分为独立服务(celery beat -A worker_app.celery_app) -exec celery \ +GEN_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" +if [ -z "$TRANSCODE_CONCURRENCY" ]; then + TRANS_CONCURRENCY=$((CONCURRENCY - GEN_CONCURRENCY)) + if [ "$TRANS_CONCURRENCY" -lt 1 ]; then + TRANS_CONCURRENCY=1 + fi +else + TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY" +fi + +echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY, beat embedded)" +celery \ -A worker_app.celery_app \ worker \ --loglevel=info \ "-B" \ - "--concurrency=${CONCURRENCY}" + -s /tmp/celerybeat-schedule \ + -Q generation \ + "--concurrency=${GEN_CONCURRENCY}" \ + "--max-tasks-per-child=${MAX_TASKS}" \ + -n generation@%h & +GEN_PID=$! + +echo "Starting transcode worker (queues=transcode,celery, concurrency=$TRANS_CONCURRENCY)" +celery \ + -A worker_app.celery_app \ + worker \ + --loglevel=info \ + -Q transcode,celery \ + "--concurrency=${TRANS_CONCURRENCY}" \ + "--max-tasks-per-child=${MAX_TASKS}" \ + -n transcode@%h & +TRANS_PID=$! + +# 任一进程退出则终止另一个,让容器整体重启(restart: unless-stopped) +trap 'echo "Shutting down workers..."; kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true' TERM INT + +wait -n $GEN_PID $TRANS_PID +EXIT_CODE=$? +echo "One worker exited (code=$EXIT_CODE), stopping the other..." +kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true +exit $EXIT_CODE diff --git a/infra/docker/nginx-production.conf b/infra/docker/nginx-production.conf index c80cfa7b5..c5269f003 100755 --- a/infra/docker/nginx-production.conf +++ b/infra/docker/nginx-production.conf @@ -16,6 +16,10 @@ server { # 注意:不能加 $uri/,否则 /assets 等与构建产物目录同名的路由会被当成目录访问,返回 403 location / { try_files $uri /index.html; + # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache, + # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用; + # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable + add_header Cache-Control "no-cache" always; } # API proxy diff --git a/infra/docker/nginx-staging.conf b/infra/docker/nginx-staging.conf index d92cdb789..f6ec40cda 100755 --- a/infra/docker/nginx-staging.conf +++ b/infra/docker/nginx-staging.conf @@ -23,6 +23,10 @@ server { location / { try_files $uri /index.html; + # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache, + # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用; + # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable + add_header Cache-Control "no-cache" always; } # API proxy diff --git a/infra/docker/nginx.conf b/infra/docker/nginx.conf index a5b581f7b..bd8a1fe68 100755 --- a/infra/docker/nginx.conf +++ b/infra/docker/nginx.conf @@ -33,6 +33,10 @@ server { location / { try_files $uri /index.html; + # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache, + # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用; + # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable + add_header Cache-Control "no-cache" always; } # API proxy diff --git a/packages/adapters/in_memory/asset_repository.py b/packages/adapters/in_memory/asset_repository.py index cc5bd76b2..3e6486db8 100755 --- a/packages/adapters/in_memory/asset_repository.py +++ b/packages/adapters/in_memory/asset_repository.py @@ -146,3 +146,44 @@ class InMemoryAssetRepository: if asset.library_id == library_id and asset.file_hash == file_hash: return asset return None + + def find_by_library_and_client_upload_id( + self, + library_id: str, + client_upload_id: str, + ) -> Asset | None: + """按素材库 + 客户端幂等 token 查找已有素材。""" + if not client_upload_id: + return None + for asset in self._assets.values(): + if asset.library_id == library_id and getattr(asset, "client_upload_id", "") == client_upload_id: + return asset + return None + + def find_recent_active_by_library_and_name( + self, + library_id: str, + name: str, + within_minutes: int = 30, + file_size: int = 0, + ) -> Asset | None: + """兜底去重:同库 + 同文件名(+同大小)且近期活动状态的素材。""" + from datetime import datetime, timedelta, timezone + + if not name: + return None + from packages.domain import AssetStatus + + cutoff = datetime.now(timezone.utc) - timedelta(minutes=within_minutes) + candidates = [ + a + for a in self._assets.values() + if a.library_id == library_id + and a.name == name + and a.status in (AssetStatus.UPLOADING, AssetStatus.PROCESSING) + and a.created_at >= cutoff + and (not file_size or file_size <= 0 or a.file_size == file_size) + ] + if not candidates: + return None + return max(candidates, key=lambda a: a.created_at) diff --git a/packages/adapters/sqlalchemy_impl/asset_repository.py b/packages/adapters/sqlalchemy_impl/asset_repository.py index 146ab706d..fad219727 100755 --- a/packages/adapters/sqlalchemy_impl/asset_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_repository.py @@ -134,6 +134,7 @@ class SQLAlchemyAssetRepository: quality_score=asset.quality_score, uploaded_by_user_id=asset.uploaded_by_user_id or "system", file_hash=asset.file_hash or None, + client_upload_id=asset.client_upload_id or None, created_at=asset.created_at, updated_at=now, ) @@ -163,6 +164,8 @@ class SQLAlchemyAssetRepository: model.quality_score = asset.quality_score model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id model.file_hash = asset.file_hash or model.file_hash + if getattr(model, "client_upload_id", None) is None and asset.client_upload_id: + model.client_upload_id = asset.client_upload_id model.updated_at = datetime.now(timezone.utc) self.session.flush() self._sync_asset_tags(asset.id, asset.tag_ids) @@ -388,6 +391,7 @@ class SQLAlchemyAssetRepository: quality_score=model.quality_score, uploaded_by_user_id=model.uploaded_by_user_id, file_hash=model.file_hash or "", + client_upload_id=getattr(model, "client_upload_id", None) or "", metadata=metadata, tag_ids=tag_ids, created_at=model.created_at, @@ -452,3 +456,58 @@ class SQLAlchemyAssetRepository: if model is None: return None return self._to_domain(model) + + def find_by_library_and_client_upload_id( + self, + library_id: str, + client_upload_id: str, + ) -> Asset | None: + """按素材库 + 客户端幂等 token 查找已有素材(complete 幂等)。""" + if not client_upload_id: + return None + model = ( + self.session.query(AssetModel) + .filter( + AssetModel.asset_library_id == library_id, + AssetModel.client_upload_id == client_upload_id, + ) + .first() + ) + if model is None: + return None + return self._to_domain(model) + + def find_recent_active_by_library_and_name( + self, + library_id: str, + name: str, + within_minutes: int = 30, + file_size: int = 0, + ) -> Asset | None: + """兜底去重:同库 + 同文件名(+同大小)且近期仍处活动状态(uploading/processing)的素材。 + + 用于旧客户端未传 file_hash/client_upload_id 时,防止 complete 超时重试 + 反复创建 PROCESSING 占位记录。只命中"活动中"的近期记录,READY 历史素材不拦。 + + 严格模式(#1714 误杀修复):file_size 必须 > 0 且与记录大小严格一致; + file_size=0(大小未知)时直接返回 None——宁可漏判(极端情况下多建一条 + 占位)也不可仅凭同名 + processing 误杀内容全新的视频。 + """ + from datetime import datetime, timedelta, timezone + + if not name: + return None + if not file_size or file_size <= 0: + return None + cutoff = datetime.now(timezone.utc) - timedelta(minutes=within_minutes) + query = self.session.query(AssetModel).filter( + AssetModel.asset_library_id == library_id, + AssetModel.name == name, + AssetModel.status.in_([AssetStatus.UPLOADING.value, AssetStatus.PROCESSING.value]), + AssetModel.created_at >= cutoff, + AssetModel.file_size == file_size, + ) + model = query.order_by(AssetModel.created_at.desc()).first() + if model is None: + return None + return self._to_domain(model) diff --git a/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py index 38aa44f9d..893e6a582 100755 --- a/packages/adapters/sqlalchemy_impl/generation_task_repository.py +++ b/packages/adapters/sqlalchemy_impl/generation_task_repository.py @@ -38,6 +38,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask: bgm_config=dict(getattr(model, "bgm_config", {}) or {}), is_preview=bool(getattr(model, "is_preview", False)), source_task_id=getattr(model, "source_task_id", "") or "", + celery_task_id=getattr(model, "celery_task_id", "") or "", output_width=getattr(model, "output_width", 1280) or 1280, output_height=getattr(model, "output_height", 720) or 720, cover_url=getattr(model, "cover_url", "") or "", @@ -82,6 +83,7 @@ class SQLAlchemyGenerationTaskRepository: bgm_config=task.bgm_config or {}, is_preview=task.is_preview or False, source_task_id=task.source_task_id or "", + celery_task_id=getattr(task, "celery_task_id", "") or "", output_width=task.output_width, output_height=task.output_height, cover_url=task.cover_url or "", @@ -138,6 +140,52 @@ class SQLAlchemyGenerationTaskRepository: .count() ) + def count_running_by_user(self, user_id: str) -> int: + """统计指定用户处于 running 状态的任务数(用于限流提示展示)。""" + return ( + self.session.query(GenerationTaskModel) + .filter( + GenerationTaskModel.created_by_user_id == user_id, + GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value, + ) + .count() + ) + + def count_running_total(self) -> int: + """统计全局处于 running 状态的任务数(worker 实际在执行的任务数)。""" + return ( + self.session.query(GenerationTaskModel) + .filter(GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value) + .count() + ) + + def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float: + """估算最近完成任务的平均耗时(秒),用于 429 限流提示的等待预估。 + + 取最近 N 条 completed 任务的 (completed_at - started_at) 平均值; + 无足够历史数据时返回 default_seconds。 + 用 Python 侧计算差值,避免 SQLite/PostgreSQL 方言差异。 + """ + rows = ( + self.session.query(GenerationTaskModel.started_at, GenerationTaskModel.completed_at) + .filter( + GenerationTaskModel.status == GenerationTaskStatus.COMPLETED.value, + GenerationTaskModel.started_at.isnot(None), + GenerationTaskModel.completed_at.isnot(None), + ) + .order_by(GenerationTaskModel.completed_at.desc()) + .limit(limit) + .all() + ) + durations = [ + (completed - started).total_seconds() + for started, completed in rows + if completed and started and (completed - started).total_seconds() > 0 + ] + if not durations: + return default_seconds + return sum(durations) / len(durations) + def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: models = ( self.session.query(GenerationTaskModel) @@ -269,6 +317,7 @@ class SQLAlchemyGenerationTaskRepository: if hasattr(model, "is_preview"): model.is_preview = task.is_preview or False model.source_task_id = task.source_task_id or "" + model.celery_task_id = getattr(task, "celery_task_id", "") or model.celery_task_id or "" model.output_width = task.output_width model.output_height = task.output_height model.cover_url = task.cover_url or "" @@ -280,12 +329,14 @@ class SQLAlchemyGenerationTaskRepository: def cleanup_stale_running(self, timeout_minutes: int = 10) -> int: """清理超时未更新的 running 任务(孤儿任务)。 - 将 status=running 且 updated_at 超过 timeout_minutes 分钟未更新的任务 - 标记为 failed,error_message 标记为任务执行中断。 - Returns: - 清理的任务数量 + 清理的任务数量(仅计数,保持旧签名兼容) """ + items = self.cleanup_stale_running_with_ids(timeout_minutes) + return len(items) + + def cleanup_stale_running_with_ids(self, timeout_minutes: int = 10) -> list[tuple[str, str]]: + """同 cleanup_stale_running,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。""" from datetime import timedelta cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes) @@ -298,8 +349,10 @@ class SQLAlchemyGenerationTaskRepository: .all() ) if not models: - return 0 + return [] + result: list[tuple[str, str]] = [] for model in models: + result.append((model.id, getattr(model, "celery_task_id", "") or "")) model.status = GenerationTaskStatus.FAILED.value model.error_message = "任务执行中断(worker重启/超时)" model.error_info = { @@ -309,43 +362,43 @@ class SQLAlchemyGenerationTaskRepository: } model.completed_at = datetime.now(timezone.utc) self.session.commit() - return len(models) + return result def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int: """清理超时的 pending 任务(未被 Worker 拉取的任务)。 - 全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列, - 导致新用户无法创建任务。将超时的 pending 任务标记为 failed。 - - Args: - timeout_minutes: 超时时间(分钟),默认 30 分钟 - Returns: - 清理的任务数量 + 清理的任务数量(仅计数,保持旧签名兼容) """ + items = self.cleanup_stale_pending_with_ids(timeout_minutes) + return len(items) + + def cleanup_stale_pending_with_ids(self, timeout_minutes: int = 30) -> list[tuple[str, str]]: + """同 cleanup_stale_pending,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。""" from datetime import timedelta cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes) - error_info = { - "error_type": "PendingTimeout", - "message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理", - "failed_at": datetime.now(timezone.utc).isoformat(), - } - count = ( + models = ( self.session.query(GenerationTaskModel) .filter( GenerationTaskModel.status == GenerationTaskStatus.PENDING.value, GenerationTaskModel.created_at < cutoff, ) - .update( - { - GenerationTaskModel.status: GenerationTaskStatus.FAILED.value, - GenerationTaskModel.error_message: "pending timeout: auto cleanup", - GenerationTaskModel.error_info: error_info, - GenerationTaskModel.completed_at: datetime.now(timezone.utc), - }, - synchronize_session=False, - ) + .all() ) + if not models: + return [] + error_info = { + "error_type": "PendingTimeout", + "message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理", + "failed_at": datetime.now(timezone.utc).isoformat(), + } + result: list[tuple[str, str]] = [] + for model in models: + result.append((model.id, getattr(model, "celery_task_id", "") or "")) + model.status = GenerationTaskStatus.FAILED.value + model.error_message = "pending timeout: auto cleanup" + model.error_info = error_info + model.completed_at = datetime.now(timezone.utc) self.session.commit() - return count + return result diff --git a/packages/adapters/sqlalchemy_impl/ingest_job_repository.py b/packages/adapters/sqlalchemy_impl/ingest_job_repository.py index f16dc735d..7c42fdeee 100644 --- a/packages/adapters/sqlalchemy_impl/ingest_job_repository.py +++ b/packages/adapters/sqlalchemy_impl/ingest_job_repository.py @@ -18,6 +18,8 @@ class SQLAlchemyIngestJobRepository: error_message=job.error_message, result_asset_id=job.result_asset_id, file_hash=job.file_hash, + asset_id=job.asset_id or "", + celery_task_id=getattr(job, "celery_task_id", "") or "", created_at=job.created_at, updated_at=job.updated_at, ) @@ -38,6 +40,8 @@ class SQLAlchemyIngestJobRepository: error_message=model.error_message, result_asset_id=model.result_asset_id, file_hash=model.file_hash or "", + asset_id=getattr(model, "asset_id", "") or "", + celery_task_id=getattr(model, "celery_task_id", "") or "", created_at=model.created_at, updated_at=model.updated_at, ) @@ -54,6 +58,12 @@ class SQLAlchemyIngestJobRepository: model.error_message = job.error_message model.result_asset_id = job.result_asset_id model.file_hash = job.file_hash + model.storage_key = job.storage_key + if job.asset_id: + model.asset_id = job.asset_id + celery_tid = getattr(job, "celery_task_id", "") + if celery_tid: + model.celery_task_id = celery_tid model.updated_at = job.updated_at self.session.commit() return job diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 8e71fcb32..04ff2381d 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -38,6 +38,7 @@ class UserModel(Base): phone = Column(String(32), nullable=True, unique=True, index=True) phone_verified = Column(Boolean, nullable=False, default=False) binding_completed_at = Column(DateTime, nullable=True) + profile_completed = Column(Boolean, nullable=False, default=True, server_default="true") created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -94,6 +95,7 @@ class AssetModel(Base): quality_score = Column(Float, nullable=True) uploaded_by_user_id = Column(String(36), nullable=False) file_hash = Column(String(64), nullable=True, index=True) + client_upload_id = Column(String(64), nullable=True, index=True) extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc), index=True) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -244,6 +246,8 @@ class IngestJobModel(Base): error_message = Column(Text, nullable=False, default="") result_asset_id = Column(String(36), nullable=False, default="") file_hash = Column(String(64), nullable=True, index=True) + asset_id = Column(String(36), nullable=False, default="", index=True) + celery_task_id = Column(String(64), nullable=False, default="", server_default="") created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -296,6 +300,7 @@ class GenerationTaskModel(Base): resolution = Column(String(20), nullable=False, default="") is_preview = Column(Boolean, nullable=False, default=False, index=True) source_task_id = Column(String(32), nullable=False, default="", index=True) + celery_task_id = Column(String(64), nullable=False, default="", server_default="") output_width = Column(Integer, nullable=False, default=1280) output_height = Column(Integer, nullable=False, default=720) cover_url = Column(String(1000), nullable=False, default="") diff --git a/packages/adapters/sqlalchemy_impl/user_repository.py b/packages/adapters/sqlalchemy_impl/user_repository.py index ac9631aa9..4bafb2b60 100755 --- a/packages/adapters/sqlalchemy_impl/user_repository.py +++ b/packages/adapters/sqlalchemy_impl/user_repository.py @@ -38,6 +38,7 @@ class SQLAlchemyUserRepository(UserRepository): model.phone = user.phone model.phone_verified = user.phone_verified model.binding_completed_at = user.binding_completed_at + model.profile_completed = user.profile_completed model.created_at = user.created_at self.session.commit() @@ -113,5 +114,6 @@ class SQLAlchemyUserRepository(UserRepository): phone=model.phone, phone_verified=model.phone_verified or False, binding_completed_at=model.binding_completed_at, + profile_completed=model.profile_completed if model.profile_completed is not None else True, created_at=model.created_at, ) diff --git a/packages/application/auth/wechat_bind_use_case.py b/packages/application/auth/wechat_bind_use_case.py new file mode 100644 index 000000000..4b9cce0bb --- /dev/null +++ b/packages/application/auth/wechat_bind_use_case.py @@ -0,0 +1,115 @@ +""" +微信账号绑定/解绑 Use Case(已登录用户场景) + +与 wechat_sync_use_case(登录/注册,系统级)不同: +- bind:把微信 openid/unionid 绑定到【当前登录账号】,不创建新用户; + 微信身份若已绑定其他账号则冲突(409)。 +- unbind:解除当前账号的微信绑定;若账号没有其他登录方式(手机/邮箱/密码), + 解绑后将无法登录,因此拒绝解绑。 +""" + +from __future__ import annotations + +from typing import Optional + +from packages.domain.entities import User + + +class WechatBindRequest: + """微信绑定请求""" + + def __init__(self, user_id: str, openid: str, unionid: str = ""): + self.user_id = user_id + self.openid = (openid or "").strip() + self.unionid = (unionid or "").strip() + + +class WechatBindResult: + """微信绑定/解绑结果""" + + def __init__(self, user: User): + self.user = user + + +class WechatBindUseCase: + """已登录用户绑定微信用例""" + + def __init__(self, user_repository): + self.user_repository = user_repository + + def bind(self, request: WechatBindRequest) -> tuple[Optional[WechatBindResult], Optional[str], int]: + """ + 绑定微信到当前登录账号。 + + Returns: + (结果, 错误信息, http状态码) - 成功时错误信息为 None、状态码为 200; + 冲突返回 409,客户端/服务端错误返回 400/404。 + """ + if not request.openid: + return None, "缺少微信 openid", 400 + + user = self.user_repository.find_by_id(request.user_id) + if user is None: + return None, "当前用户不存在", 404 + + # 已绑定同一个微信:幂等成功 + if user.wechat_openid == request.openid: + return WechatBindResult(user=user), None, 200 + + # 当前账号已绑定其他微信 + if user.wechat_openid: + return None, "当前账号已绑定微信,请先解绑", 409 + + # openid 已被其他账号占用 + existing = self.user_repository.find_by_wechat_openid(request.openid) + if existing is not None and existing.id != user.id: + return None, "该微信已绑定其他账号,请先在原账号解绑", 409 + + # unionid 冲突:同主体微信已绑其他账号 + if request.unionid: + existing_union = self.user_repository.find_by_wechat_unionid(request.unionid) + if existing_union is not None and existing_union.id != user.id: + return None, "该微信主体已绑定其他账号,请先在原账号解绑", 409 + + user.wechat_openid = request.openid + if request.unionid and not user.wechat_unionid: + user.wechat_unionid = request.unionid + self.user_repository.save(user) + + return WechatBindResult(user=user), None, 200 + + +class WechatUnbindUseCase: + """已登录用户解绑微信用例""" + + def __init__(self, user_repository): + self.user_repository = user_repository + + def unbind(self, user_id: str) -> tuple[Optional[WechatBindResult], Optional[str], int]: + """ + 解除当前账号的微信绑定。 + + 解绑前置条件:账号必须还有其他登录方式(密码 / 已验证手机 / 真实邮箱), + 否则解绑后将永远无法登录。 + """ + user = self.user_repository.find_by_id(user_id) + if user is None: + return None, "当前用户不存在", 404 + + if not user.wechat_openid: + return None, "当前账号未绑定微信", 400 + + # 守卫:解绑后账号必须仍有可实际使用的登录方式。 + # 注意:微信注册用户带的是【随机密码】(用户不知道、无法用密码登录, + # 且 @wechat.local 占位邮箱收不到重置邮件),故 password_hash 不作为兜底依据, + # 口径与 /auth/me 的 binding_complete 一致。 + has_phone = bool(user.phone and user.phone_verified) + has_real_email = bool(user.email and user.email_verified and "@wechat.local" not in user.email) + if not (has_phone or has_real_email): + return None, "账号需要至少一种其他登录方式(已验证手机或真实邮箱)后才能解绑微信", 400 + + user.wechat_openid = None + user.wechat_unionid = None + self.user_repository.save(user) + + return WechatBindResult(user=user), None, 200 diff --git a/packages/application/auth/wechat_oauth_service.py b/packages/application/auth/wechat_oauth_service.py index 16cdf5d47..08e4416ce 100755 --- a/packages/application/auth/wechat_oauth_service.py +++ b/packages/application/auth/wechat_oauth_service.py @@ -20,6 +20,7 @@ import requests logger = logging.getLogger(__name__) STATE_TTL_SECONDS = 600 # state 有效期 10 分钟 +STATE_KEY_PREFIX = "wechat:state:" # Redis key 前缀(独立逻辑命名空间) class MemoryStateStore: @@ -53,6 +54,80 @@ class MemoryStateStore: del self._states[s] +class RedisStateStore: + """Redis state 存储(多实例/容器重启安全)。 + + 复用现有 Redis(celery broker 同实例),key 前缀 wechat:state:, + TTL 10 分钟,SET NX EX + GETDEL 保证一次性消费。 + Redis 不可用时降级为内存存储,保证登录流程不中断(单节点场景)。 + """ + + def __init__( + self, + redis_url: str = "", + ttl_seconds: int = STATE_TTL_SECONDS, + key_prefix: str = STATE_KEY_PREFIX, + client=None, + ): + self._ttl = ttl_seconds + self._prefix = key_prefix + self._fallback = MemoryStateStore(ttl_seconds=ttl_seconds) + self._redis = None + if client is not None: + # 测试/显式注入 + self._redis = client + return + try: + import redis + + self._redis = redis.Redis.from_url(redis_url, decode_responses=True) + self._redis.ping() + logger.info( + "微信 state 存储使用 Redis: %s db=%s", + self._redis.connection_pool.connection_kwargs.get("host"), + self._redis.connection_pool.connection_kwargs.get("db"), + ) + except Exception as e: # noqa: BLE001 — Redis 不可用降级内存,登录流程不中断 + logger.warning("微信 state Redis 不可用,降级为内存存储: %s", e) + self._redis = None + + def _key(self, state: str) -> str: + return f"{self._prefix}{state}" + + def put(self, state: str) -> None: + if self._redis is None: + self._fallback.put(state) + return + try: + # SET key 1 NX EX ttl:不存在才写入,自带过期 + self._redis.set(self._key(state), "1", nx=True, ex=self._ttl) + except Exception as e: # noqa: BLE001 + logger.warning("微信 state 写入 Redis 失败,降级内存: %s", e) + self._fallback.put(state) + + # Lua:原子读取并删除(单线程执行),兼容所有 Redis 版本(GETDEL 需 6.2+) + _CONSUME_LUA = """ +local v = redis.call('GET', KEYS[1]) +if v then redis.call('DEL', KEYS[1]) end +return v +""" + + def verify_and_consume(self, state: str) -> bool: + if self._redis is None: + return self._fallback.verify_and_consume(state) + try: + try: + val = self._redis.eval(self._CONSUME_LUA, 1, self._key(state)) + except Exception: # noqa: BLE001 — eval 不可用时退化 GET+DELETE + val = self._redis.get(self._key(state)) + if val is not None: + self._redis.delete(self._key(state)) + return val is not None + except Exception as e: # noqa: BLE001 + logger.warning("微信 state 校验 Redis 失败,降级内存: %s", e) + return self._fallback.verify_and_consume(state) + + @dataclass class WechatUserInfo: """微信用户信息""" @@ -158,6 +233,8 @@ class WechatOAuthService: "grant_type": "authorization_code", } token_resp = requests.get(token_url, params=token_params, timeout=10) + # 微信响应头不带 charset,requests 默认按 ISO-8859-1 解码会导致中文乱码 + token_resp.encoding = "utf-8" token_data = token_resp.json() if "errcode" in token_data and token_data["errcode"] != 0: @@ -176,6 +253,8 @@ class WechatOAuthService: "lang": "zh_CN", } user_resp = requests.get(user_url, params=user_params, timeout=10) + # 同上:显式 UTF-8 解码,保证中文昵称/unionid 等不乱码 + user_resp.encoding = "utf-8" user_data = user_resp.json() if "errcode" in user_data and user_data["errcode"] != 0: @@ -200,7 +279,29 @@ class WechatOAuthService: return None, "微信登录处理失败" +# 模块级单例:state 存储必须跨请求共享,否则 /wechat/url 生成的 state +# 与 /wechat/callback 校验时不在同一个 MemoryStateStore,回调必然 400。 +# 多实例部署时应替换为 Redis state store(单容器多 worker 也需如此)。 +_oauth_service_singleton: WechatOAuthService | None = None + + +def _build_default_state_store(): + """默认 state 存储:优先 Redis(多实例/重启安全),不可用由 store 内部降级内存。""" + redis_url = "" + try: + from app.config import get_settings + + redis_url = get_settings().CELERY_BROKER_URL or get_settings().REDIS_URL + except Exception: # noqa: BLE001 — API 配置不可用时退回环境变量 + redis_url = os.environ.get("CELERY_BROKER_URL", "") or os.environ.get("REDIS_URL", "") + if redis_url: + return RedisStateStore(redis_url) + return MemoryStateStore() + + def get_wechat_oauth_service() -> WechatOAuthService: - """获取微信 OAuth 服务单例""" - # TODO: 可替换为 Redis state store - return WechatOAuthService() + """获取微信 OAuth 服务单例(state store 跨请求共享)""" + global _oauth_service_singleton + if _oauth_service_singleton is None: + _oauth_service_singleton = WechatOAuthService(state_store=_build_default_state_store()) + return _oauth_service_singleton diff --git a/packages/application/auth/wechat_sync_use_case.py b/packages/application/auth/wechat_sync_use_case.py index 7f343eb7b..bb96dc1d5 100644 --- a/packages/application/auth/wechat_sync_use_case.py +++ b/packages/application/auth/wechat_sync_use_case.py @@ -200,6 +200,8 @@ class WechatSyncUseCase: email_verified=True, # 微信登录视为已验证 wechat_openid=request.openid, wechat_unionid=request.unionid or None, + # 微信新建用户首次登录需引导设置昵称 + profile_completed=False, ) self.user_repository.save(user) diff --git a/packages/application/ingest_jobs.py b/packages/application/ingest_jobs.py index 75a708de7..92f062ac2 100644 --- a/packages/application/ingest_jobs.py +++ b/packages/application/ingest_jobs.py @@ -12,6 +12,8 @@ class SubmitIngestJobCommand: library_id: str storage_key: str file_hash: str = "" + asset_id: str = "" + celery_task_id: str = "" class SubmitIngestJobUseCase: @@ -24,5 +26,7 @@ class SubmitIngestJobUseCase: library_id=command.library_id, storage_key=command.storage_key, file_hash=command.file_hash, + asset_id=command.asset_id, + celery_task_id=command.celery_task_id, ) return self.ingest_job_repository.create(job) diff --git a/packages/application/ingest_orphan_cleanup.py b/packages/application/ingest_orphan_cleanup.py new file mode 100644 index 000000000..22fbb09a4 --- /dev/null +++ b/packages/application/ingest_orphan_cleanup.py @@ -0,0 +1,310 @@ +"""上传/转码链路(IngestJob + Asset)孤儿清理核心逻辑。 + +#1714:generation 链路有 cleanup_stale_running/pending 兜底,但上传链路 +(ingest_jobs + assets)没有。worker 容器重启/进程 OOM 时,已 prefetch 的 +celery 消息会丢失(transcode 队列 worker_prefetch_multiplier=1,消息预取后 +宕机即丢失,Redis 队列里也不再存在),导致: + +- ingest_jobs.status 永久卡 pending/processing +- assets.status 永久卡 processing/uploading(complete 阶段预建的占位) + +本模块提供纯核心(session 注入,便于单测):超时阈值内无更新的记录 +批量标终态(job→failed、asset→error),并返回 (job_id, celery_task_id) +列表供调用方 revoke + purge 残留队列消息。 +""" + +from __future__ import annotations + +import logging +from datetime import datetime, timedelta, timezone +from typing import Any, Callable + +logger = logging.getLogger(__name__) + +# ingest_job PROCESSING 超时阈值:ingest 任务包含下载 + ffprobe + HEVC 转码 +# (1GB 视频约 10-20 分钟)+ 回传 OSS,正常任务可能跑 20-30 分钟; +# 60 分钟阈值覆盖大文件转码 + 抖动,绝不误杀正常任务。 +INGEST_PROCESSING_TIMEOUT_MINUTES = 60 + +# ingest_job PENDING 超时阈值:transcode 队列 concurrency=1,队列积压时 +# 正常排队可能较久;90 分钟覆盖 worker 短暂停消费 + 排队。 +INGEST_PENDING_TIMEOUT_MINUTES = 90 + +# Asset 占位超时阈值:无关联 ingest_job 的孤儿占位(complete 预建后派单失败等), +# 阈值放宽到 120 分钟,避免与 ingest_job 生命周期错杀。 +ASSET_ORPHAN_TIMEOUT_MINUTES = 120 + +_TERMINAL_JOB_STATUSES = ("failed", "completed") +_TERMINAL_ASSET_STATUSES = ("ready", "error", "deleted") + + +def _now() -> datetime: + return datetime.now(timezone.utc) + + +def cleanup_stale_ingest_jobs( + session: Any, + *, + processing_timeout_minutes: int = INGEST_PROCESSING_TIMEOUT_MINUTES, + pending_timeout_minutes: int = INGEST_PENDING_TIMEOUT_MINUTES, + commit: bool = True, +) -> tuple[list[tuple[str, str]], list[str]]: + """清理超时卡 pending/processing 的 ingest_jobs,并联动关联 asset。 + + Args: + session: SQLAlchemy session(或提供 query/commit 的鸭子类型) + processing_timeout_minutes: processing 状态超时阈值 + pending_timeout_minutes: pending 状态超时阈值 + commit: 是否提交事务 + + Returns: + (job_items, asset_ids) + - job_items: [(job_id, celery_task_id), ...] 供 revoke/purge + - asset_ids: 被联动标记为 error 的 asset id 列表 + """ + from packages.adapters.sqlalchemy_impl.models import AssetModel, IngestJobModel + + now = _now() + processing_cutoff = now - timedelta(minutes=processing_timeout_minutes) + pending_cutoff = now - timedelta(minutes=pending_timeout_minutes) + + stale_jobs = ( + session.query(IngestJobModel) + .filter( + IngestJobModel.status.in_(["pending", "processing"]), + ( + (IngestJobModel.status == "processing") & (IngestJobModel.updated_at < processing_cutoff) + | (IngestJobModel.status == "pending") & (IngestJobModel.created_at < pending_cutoff) + ), + ) + .all() + ) + + job_items: list[tuple[str, str]] = [] + asset_ids: list[str] = [] + stale_asset_models: list[Any] = [] + for job_model in stale_jobs: + ref_time = job_model.updated_at or job_model.created_at + if ref_time.tzinfo is None: # SQLite 读回 naive datetime 的防御 + ref_time = ref_time.replace(tzinfo=timezone.utc) + stale_minutes = int((now - ref_time).total_seconds() // 60) + job_model.status = "failed" + job_model.error_message = ( + f"转码任务执行中断(超过超时阈值未更新,疑似 worker 重启/进程退出,已卡死 {stale_minutes} 分钟)" + ) + job_model.updated_at = now + job_items.append((job_model.id, getattr(job_model, "celery_task_id", "") or "")) + if job_model.asset_id: + asset_ids.append(job_model.asset_id) + + if asset_ids: + stale_asset_models = ( + session.query(AssetModel) + .filter( + AssetModel.id.in_(asset_ids), + AssetModel.status.in_(["processing", "uploading"]), + ) + .all() + ) + for asset_model in stale_asset_models: + asset_model.status = "error" + asset_model.updated_at = now + + if commit and (job_items or stale_asset_models): + session.commit() + + if job_items: + logger.warning( + "[ingest-cleanup] 清理 %d 个超时 ingest_job(processing>%dm / pending>%dm),联动 %d 个 asset 标 error", + len(job_items), + processing_timeout_minutes, + pending_timeout_minutes, + len(stale_asset_models), + ) + return job_items, [a.id for a in stale_asset_models] + + +def cleanup_orphan_processing_assets( + session: Any, + *, + timeout_minutes: int = ASSET_ORPHAN_TIMEOUT_MINUTES, + commit: bool = True, +) -> list[str]: + """清理无 ingest_job 关联、超时卡 processing/uploading 的孤儿 asset 占位。 + + complete 阶段预建 asset 后若派单失败(或 direct 上传 complete 后 + 未触发 ingest),占位会永久卡住。这类 asset 没有对应 ingest_job, + 只能按 created_at 超时兜底标 error。 + """ + from packages.adapters.sqlalchemy_impl.models import AssetModel, IngestJobModel + + cutoff = _now() - timedelta(minutes=timeout_minutes) + orphan_assets = ( + session.query(AssetModel) + .outerjoin(IngestJobModel, IngestJobModel.asset_id == AssetModel.id) + .filter( + AssetModel.status.in_(["processing", "uploading"]), + AssetModel.created_at < cutoff, + IngestJobModel.id.is_(None), + ) + .all() + ) + for asset_model in orphan_assets: + asset_model.status = "error" + asset_model.updated_at = _now() + if commit and orphan_assets: + session.commit() + logger.warning("[ingest-cleanup] 清理 %d 个无 job 关联的超时孤儿 asset 占位", len(orphan_assets)) + return [a.id for a in orphan_assets] + + +def revoke_stale_ingest_messages( + job_items: list[tuple[str, str]], + *, + celery_app_factory: Callable[[], Any] | None = None, + broker_url_factory: Callable[[], str] | None = None, +) -> int: + """revoke + 物理清理 ingest 作废消息(transcode/celery 队列)。 + + 消息可能已在 worker 宕机时丢失(队列里查不到),那也无害; + 若消息还在(极端重复投递),物理清除防止重投执行。 + 失败不阻断清理(ingest_asset 的执行前状态守卫是第二道防线)。 + """ + biz_ids = [jid for jid, _ in job_items if jid] + celery_ids = [cid for _, cid in job_items if cid] + if not biz_ids and not celery_ids: + return 0 + try: + from packages.shared.celery_orphan_guard import revoke_and_purge + + app = celery_app_factory() if celery_app_factory else None + broker_url = broker_url_factory() if broker_url_factory else "" + if app is None or not broker_url: + from worker_app.celery_app import celery_app as _app + from worker_app.core.config import get_settings + + app = _app + broker_url = get_settings().broker_url + return revoke_and_purge( + app, + broker_url, + business_task_ids=biz_ids, + celery_task_ids=celery_ids, + queue_names=("transcode", "celery"), + ) + except Exception as e: # noqa: BLE001 + logger.error("撤销作废 ingest 队列消息失败(执行前守卫仍会兜底): %s", e, exc_info=True) + return 0 + + +# ── worker 启动恢复(#1714)────────────────────────────────────────────── +# +# task_acks_late=True 下,worker 崩溃/容器重启时未 ack 的消息理论上会在 +# visibility_timeout 到期后重新投递;但 prefork 进程异常、部署窗口跨 +# visibility 配置边界等场景仍可能留下卡在 processing 的 ingest_job +# (staging 实证:03:16 派单、03:45 置 processing 后 worker 重启, +# unacked 消息未重投,任务永久卡死)。启动时做一次显式恢复扫描兜底。 +# +# 恢复策略:processing 超过 stuck_minutes(默认 10 分钟,部署中跨进程 +# 交接的正常窗口 < 10 分钟,不会误抢别的 worker 正在执行的任务)的 job, +# CAS 重置为 pending 并重新 send_task;旧消息若后来重投,ingest_asset +# 的执行前守卫会把状态不匹配的旧 celery 消息丢弃。 + + +def recover_stuck_ingest_jobs_on_startup( + session: Any, + *, + send_task: Callable[..., Any] | None = None, + update_celery_task_id: Callable[[str, str], None] | None = None, + lock_acquire: Callable[[], bool] | None = None, + stuck_minutes: int = 10, + commit: bool = True, +) -> int: + """worker 启动时把卡在 processing 超时的 ingest_job 重新派单。 + + Args: + session: SQLAlchemy session + send_task: celery send_task 可调用(注入便于测试);不传则用 worker celery_app + update_celery_task_id: 回写新 celery task id 的回调(job_id, new_task_id) + lock_acquire: 分布式锁获取回调(多 worker 进程同时启动时只允许一个恢复); + 返回 False 表示未抢到锁,本次跳过 + stuck_minutes: processing 超过该分钟数视为卡死 + + Returns: + 重新派单的 job 数 + """ + if lock_acquire is not None and not lock_acquire(): + logger.info("[ingest-recover] 未抢到恢复锁,跳过(另一进程正在恢复)") + return 0 + + from packages.adapters.sqlalchemy_impl.models import IngestJobModel + + cutoff = _now() - timedelta(minutes=stuck_minutes) + stuck_jobs = ( + session.query(IngestJobModel) + .filter(IngestJobModel.status == "processing", IngestJobModel.updated_at < cutoff) + .order_by(IngestJobModel.updated_at.asc()) + .all() + ) + + if not stuck_jobs: + logger.info("[ingest-recover] 无卡死 processing ingest_job 需要恢复") + return 0 + + if send_task is None: + from worker_app.celery_app import celery_app as _app + + send_task = _app.send_task + + recovered = 0 + for job_model in stuck_jobs: + # CAS:只有仍是 processing 才重置(并发/旧消息已回写终态时不碰) + updated = ( + session.query(IngestJobModel) + .filter(IngestJobModel.id == job_model.id, IngestJobModel.status == "processing") + .update({"status": "pending", "error_message": "", "updated_at": _now()}) + ) + if not updated: + continue + try: + result = send_task("worker.ingest_asset", args=[job_model.id]) + new_task_id = getattr(result, "id", "") or "" + except Exception as e: # noqa: BLE001 + logger.error("[ingest-recover] 重新派单失败 job_id=%s: %s", job_model.id, e) + continue + if new_task_id: + job_model.celery_task_id = new_task_id + if update_celery_task_id is not None: + update_celery_task_id(job_model.id, new_task_id) + logger.warning( + "[ingest-recover] 卡死 ingest_job %s 已重置 pending 并重新派单 (new celery task=%s)", + job_model.id, + new_task_id, + ) + recovered += 1 + + if commit and recovered: + session.commit() + logger.warning("[ingest-recover] 启动恢复完成,共重新派单 %d 个卡死 ingest_job", recovered) + return recovered + + +def make_redis_recovery_lock(lock_key: str = "ingest:recover:startup", ttl_seconds: int = 300): + """构造基于 Redis SET NX 的恢复锁工厂(多 worker 进程互斥)。 + + 返回一个无参 callable,调用时尝试抢锁:抢到返回 True,未抢到返回 False。 + Redis 不可用时不阻断启动恢复(返回 True,恢复逻辑自身有 CAS 幂等保护)。 + """ + + def _acquire() -> bool: + try: + import redis as redis_lib + from worker_app.core.config import get_settings + + client = redis_lib.Redis.from_url(get_settings().broker_url) + return bool(client.set(lock_key, "1", nx=True, ex=ttl_seconds)) + except Exception as e: # noqa: BLE001 + logger.warning("[ingest-recover] Redis 锁不可用,降级为无锁执行(CAS 兜底): %s", e) + return True + + return _acquire diff --git a/packages/domain/entities.py b/packages/domain/entities.py index 4df342e4d..b03997a16 100755 --- a/packages/domain/entities.py +++ b/packages/domain/entities.py @@ -57,6 +57,8 @@ class User: phone: str | None = None phone_verified: bool = False binding_completed_at: datetime | None = None + # 资料是否已完善(微信新用户首次设置昵称后置 True;邮箱注册默认 True) + profile_completed: bool = True created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) @@ -174,6 +176,7 @@ class Asset: quality_score: float | None = None uploaded_by_user_id: str = "" file_hash: str = "" + client_upload_id: str = "" metadata: dict[str, Any] = field(default_factory=dict) tag_ids: list[str] = field(default_factory=list) created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) @@ -208,6 +211,7 @@ class Asset: quality_score: float | None = None, uploaded_by_user_id: str = "", file_hash: str = "", + client_upload_id: str = "", ) -> "Asset": clean_name = name.strip() if not clean_name: @@ -235,6 +239,7 @@ class Asset: quality_score=quality_score, uploaded_by_user_id=uploaded_by_user_id.strip(), file_hash=file_hash.strip(), + client_upload_id=client_upload_id.strip(), metadata=metadata or {}, tag_ids=[], ) @@ -266,6 +271,8 @@ class IngestJob: error_message: str = "" result_asset_id: str = "" file_hash: str = "" + asset_id: str = "" + celery_task_id: str = "" created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) @@ -276,6 +283,8 @@ class IngestJob: library_id: str, storage_key: str, file_hash: str = "", + asset_id: str = "", + celery_task_id: str = "", ) -> "IngestJob": if not project_id.strip(): raise ValueError("project_id 不能为空") @@ -289,4 +298,6 @@ class IngestJob: library_id=library_id.strip(), storage_key=storage_key.strip(), file_hash=file_hash.strip(), + asset_id=asset_id.strip(), + celery_task_id=celery_task_id.strip(), ) diff --git a/packages/domain/generation_task.py b/packages/domain/generation_task.py index a0e46ac37..ad00c3a70 100755 --- a/packages/domain/generation_task.py +++ b/packages/domain/generation_task.py @@ -117,6 +117,7 @@ class GenerationTask: bgm_config: dict = field(default_factory=dict) is_preview: bool = False source_task_id: str = "" + celery_task_id: str = "" output_width: int = 1280 output_height: int = 720 cover_url: str = "" diff --git a/packages/ports/asset_repository.py b/packages/ports/asset_repository.py index 9a9c830ad..b92c19fa4 100755 --- a/packages/ports/asset_repository.py +++ b/packages/ports/asset_repository.py @@ -125,3 +125,23 @@ class AssetRepository(ABC): ) -> Asset | None: """按素材库 + 文件哈希查找已有素材(去重检测)。""" pass + + @abstractmethod + def find_by_library_and_client_upload_id( + self, + library_id: str, + client_upload_id: str, + ) -> Asset | None: + """按素材库 + 客户端幂等 token 查找已有素材(complete 幂等)。""" + pass + + @abstractmethod + def find_recent_active_by_library_and_name( + self, + library_id: str, + name: str, + within_minutes: int = 30, + file_size: int = 0, + ) -> Asset | None: + """兜底去重:同库 + 同文件名(+同大小)且近期仍在 uploading/processing 的素材。""" + pass diff --git a/packages/ports/generation_task_repository.py b/packages/ports/generation_task_repository.py index 5c21200e5..cc79f0370 100755 --- a/packages/ports/generation_task_repository.py +++ b/packages/ports/generation_task_repository.py @@ -20,6 +20,12 @@ class GenerationTaskRepository(Protocol): def count_pending_total(self) -> int: ... + def count_running_by_user(self, user_id: str) -> int: ... + + def count_running_total(self) -> int: ... + + def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float: ... + def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: ... def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: ... diff --git a/packages/shared/celery_orphan_guard.py b/packages/shared/celery_orphan_guard.py new file mode 100644 index 000000000..123404f06 --- /dev/null +++ b/packages/shared/celery_orphan_guard.py @@ -0,0 +1,222 @@ +"""孤儿任务消息撤销与执行前状态守卫(API / Worker 共享)。 + +#1714 / #1710 缺陷修复:超时清理/孤儿恢复把 DB 任务标记为 failed/cancelled +后,Redis 队列里对应的 Celery 消息仍然存在;worker 重启或重新拉取时该消息 +被再次执行,状态机抛「非法状态转换: failed → running」,旧实现打印 ERROR 后 +继续跑,最终产出半成品。 + +防御两道: +1. 清理任务标 failed 时,调用 revoke_and_purge() 撤销(celery revoke 广播, + 通知在线 worker 丢弃)并直接扫描 Redis 队列移除消息体(worker 下线期间 + 队列中的消息 revoke 广播收不到,必须物理移除); +2. 任务真正开始业务逻辑前,调用 ensure_task_claimable() 校验 DB 状态, + 非 pending 的消息直接丢弃(抛 StaleTaskDiscarded,task 捕获后安全返回, + 不进入渲染/转码,不产出半成品)。 +""" + +from __future__ import annotations + +import base64 +import json +import logging +from collections.abc import Callable, Iterable +from typing import Any + +logger = logging.getLogger(__name__) + + +class StaleTaskDiscarded(Exception): + """任务消息已作废(DB 中任务已是终态),应安全中止、丢弃消息。""" + + def __init__(self, task_id: str, status: str): + self.task_id = task_id + self.status = status + super().__init__(f"任务 {task_id} 已是终态 {status},丢弃重复/作废消息") + + +# 终态状态值集合:处于这些状态的任务消息一律不执行 +TERMINAL_STATUS_VALUES = frozenset({"failed", "cancelled", "completed"}) + + +def ensure_task_claimable( + task_id: str, + get_status: Callable[[str], str | None], + *, + task_label: str = "任务", +) -> str: + """执行前守卫:任务必须处于可领取状态(pending)。 + + Args: + task_id: 业务任务 ID + get_status: 回调,返回 DB 中任务当前状态字符串;返回 None 表示任务不存在 + task_label: 日志用任务类型名 + + Returns: + 当前状态字符串(pending);任务不存在时返回空串(由调用方处理 not found) + + Raises: + StaleTaskDiscarded: 任务已是终态(failed/cancelled/completed),消息必须丢弃 + """ + status = get_status(task_id) + if status is None: + return "" + if status in TERMINAL_STATUS_VALUES: + logger.warning("[%s] task_id=%s 状态已为 %s,消息作废,丢弃不执行", task_label, task_id, status) + raise StaleTaskDiscarded(task_id, status) + return status + + +def _extract_business_ids(raw: bytes) -> tuple[str | None, str | None]: + """从 Redis 中的 Celery 消息提取 (celery 消息 ID, 业务任务 ID)。 + + Redis transport 存储格式为 JSON 信封: + {"body": base64(json), "headers": {"id": , "task": , ...}, ...} + body 解码后 Celery task 协议为 [args, kwargs, embed]; + generate_video / ingest_asset 均以 args=[业务任务ID] 投递。 + + 无法解析时返回 (None, None)(保守保留该消息,绝不误删)。 + """ + try: + envelope = json.loads(raw) + celery_id = None + headers = envelope.get("headers") or {} + if isinstance(headers, dict): + celery_id = headers.get("id") + body = envelope.get("body") + if not body: + return celery_id, None + decoded = base64.b64decode(body) + payload = json.loads(decoded) + # 两种 body 形态: + # 1. 标准 Celery task 消息:[args, kwargs, embed] 三元组 → 业务 ID 在 payload[0][0] + # 2. 裸 producer 发布:body 即 args 数组 ["biz-id"] → 业务 ID 在 payload[0] + args = None + if isinstance(payload, dict): + args = payload.get("args") + elif isinstance(payload, (list, tuple)) and payload: + first = payload[0] + if isinstance(first, (list, tuple)): + args = first # 三元组:[args, kwargs, embed] + else: + args = payload # body 本身就是 args + if isinstance(args, (list, tuple)) and args and args[0] is not None: + return celery_id, str(args[0]) + return celery_id, None + except Exception: + return None, None + + +def purge_stale_messages_from_queues( + broker_url: str, + queue_names: Iterable[str], + business_task_ids: Iterable[str] = (), + celery_task_ids: Iterable[str] = (), +) -> int: + """扫描 Redis 队列,移除作废任务的待消费消息。 + + 同时按业务任务 ID(消息 args[0])和 celery 消息 ID(headers.id)匹配, + 任一命中即移除。未命中或无法解析的消息原样保留(保持相对顺序)。 + + Returns: + 实际移除的消息条数 + """ + biz_ids = {bid for bid in business_task_ids if bid} + msg_ids = {mid for mid in celery_task_ids if mid} + if not biz_ids and not msg_ids: + return 0 + + try: + import redis + except ImportError: + logger.warning("redis-py 不可用,跳过队列消息清理") + return 0 + + try: + client = redis.Redis.from_url(broker_url) + client.ping() + except Exception as e: + logger.warning("连接 Redis 清理作废消息失败: %s", e) + return 0 + + removed_total = 0 + try: + for queue in queue_names: + removed_total += _purge_one_queue(client, queue, biz_ids, msg_ids) + finally: + try: + client.close() + except Exception: + pass + if removed_total: + logger.info( + "从 Redis 队列移除 %d 条作废消息(biz=%s, celery=%s)", + removed_total, + sorted(biz_ids), + sorted(msg_ids), + ) + return removed_total + + +def _purge_one_queue(client: Any, queue_name: str, biz_ids: set[str], msg_ids: set[str]) -> int: + try: + raw_messages = client.lrange(queue_name, 0, -1) + except Exception as e: + logger.warning("读取队列 %s 失败: %s", queue_name, e) + return 0 + if not raw_messages: + return 0 + + keep: list[bytes] = [] + removed = 0 + for raw in raw_messages: + celery_id, biz_id = _extract_business_ids(raw) + hit = (biz_id is not None and biz_id in biz_ids) or (celery_id is not None and celery_id in msg_ids) + if hit: + removed += 1 + continue + keep.append(raw) + + if removed: + try: + pipe = client.pipeline() + pipe.delete(queue_name) + if keep: + pipe.rpush(queue_name, *keep) + pipe.execute() + except Exception as e: + logger.warning("重写队列 %s 失败: %s", queue_name, e) + return 0 + return removed + + +def revoke_and_purge( + celery_app: Any, + broker_url: str, + business_task_ids: Iterable[str] = (), + celery_task_ids: Iterable[str] = (), + *, + queue_names: Iterable[str] = ("generation", "transcode", "celery"), +) -> int: + """撤销作废任务:revoke 广播(在线 worker)+ 物理清理 Redis 队列消息。 + + Args: + celery_app: Celery app 实例(worker 端 worker_app.celery_app.celery_app) + broker_url: Redis broker URL + business_task_ids: 业务任务 ID(generation_tasks.id / ingest_jobs.id) + celery_task_ids: 入队时记录的 celery 消息 ID + queue_names: 需要扫描清理的队列名 + + Returns: + 从队列中实际移除的消息条数 + """ + for tid in celery_task_ids: + if not tid: + continue + try: + celery_app.control.revoke(tid) + except Exception as e: + logger.warning("revoke celery 消息 %s 失败: %s", tid, e) + + return purge_stale_messages_from_queues( + broker_url, queue_names, business_task_ids=business_task_ids, celery_task_ids=celery_task_ids + ) diff --git a/packages/shared/celery_queues.py b/packages/shared/celery_queues.py new file mode 100644 index 000000000..c9f4e45e6 --- /dev/null +++ b/packages/shared/celery_queues.py @@ -0,0 +1,58 @@ +"""Celery 队列定义与路由配置(API / Worker 共享)。 + +#1714 队列隔离:用户等待的视频生成任务路由到高优先级 `generation` 队列, +由专用 worker 进程独占消费;素材入库/转码等后台批量任务路由到 `transcode` +队列;其余杂项任务走默认 `celery` 队列。转码队列积压时,视频生成任务 +仍能被 generation worker 立即领取执行,不会排队。 + +队列说明: +- generation: 用户提交的视频生成/预览渲染(延迟敏感,资源消耗大) +- transcode: 素材入库(HEVC 转码)、AI 分类、素材查重(批量、可排队) +- celery(默认): 配音、语音、下载缩略图、定时清理等杂项 +""" + +from __future__ import annotations + +from kombu import Queue + +# ── 队列名常量(生产端与消费端共用,禁止拼写漂移) ── +QUEUE_GENERATION = "generation" +QUEUE_TRANSCODE = "transcode" +QUEUE_DEFAULT = "celery" + +# Worker 消费的队列列表(顺序即优先级:高优队列排在前面) +WORKER_QUEUES = (QUEUE_GENERATION, QUEUE_TRANSCODE, QUEUE_DEFAULT) + +# 队列声明:持久化队列,broker 重启不丢消息 +task_queues = ( + Queue(QUEUE_GENERATION, routing_key=QUEUE_GENERATION, durable=True), + Queue(QUEUE_TRANSCODE, routing_key=QUEUE_TRANSCODE, durable=True), + Queue(QUEUE_DEFAULT, routing_key=QUEUE_DEFAULT, durable=True), +) + +# ── 任务路由表:task name → 队列 ── +# 键支持 celery 标准通配符。 +task_routes = { + # 高优先级:用户等待的视频生成 + "worker.generate_video": {"queue": QUEUE_GENERATION}, + # 后台批量:素材入库/转码 + AI 分类 + 素材查重,积压不影响生成 + "worker.ingest_asset": {"queue": QUEUE_TRANSCODE}, + "worker.classify_asset": {"queue": QUEUE_TRANSCODE}, + "worker.process_duplication_check": {"queue": QUEUE_TRANSCODE}, + "worker.check_duplicate": {"queue": QUEUE_TRANSCODE}, +} + +# 生成任务的预取数:渲染是长任务,预取 1 避免任务被某个 worker 占住不调度 +GENERATION_WORKER_PREFETCH_MULTIPLIER = 1 + + +def apply_queue_settings(app) -> None: + """把队列隔离配置应用到 Celery app(API 生产端与 Worker 消费端都要调用)。 + + 配置 task_queues / task_routes / task_default_queue。生产端靠 task_routes + 把消息投递到对应队列;消费端靠 task_queues 声明自己消费哪些队列 + (实际消费集由启动参数 -Q 控制)。 + """ + app.conf.task_queues = task_queues + app.conf.task_routes = task_routes + app.conf.task_default_queue = QUEUE_DEFAULT diff --git a/scripts/render_env.sh b/scripts/render_env.sh index 42930f6e1..a5f128045 100644 --- a/scripts/render_env.sh +++ b/scripts/render_env.sh @@ -57,7 +57,7 @@ if [ "$TARGET_ENV" = "staging" ]; then fi # 共用 secrets 直接导出(如果存在) -SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY" +SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY WECHAT_APP_ID WECHAT_APP_SECRET" for var in $SHARED_SECRETS; do value="${!var:-}" # 已经在环境中了,无需额外操作 diff --git a/start-worker.ps1 b/start-worker.ps1 index a63e9e659..86f1e7197 100644 --- a/start-worker.ps1 +++ b/start-worker.ps1 @@ -14,4 +14,4 @@ Write-Host "`n启动 Celery Worker..." -ForegroundColor Yellow Write-Host "监听任务队列: Redis (47.98.113.167:6379)" -ForegroundColor Cyan Write-Host "`n按 Ctrl+C 停止服务`n" -ForegroundColor Gray -celery -A celery_app worker --loglevel=info --pool=solo +celery -A celery_app worker --loglevel=info --pool=solo -Q generation,transcode,celery diff --git a/tests/unit/test_1677_batch_variants.py b/tests/unit/test_1677_batch_variants.py new file mode 100644 index 000000000..bf67b5b62 --- /dev/null +++ b/tests/unit/test_1677_batch_variants.py @@ -0,0 +1,507 @@ +"""Issue #1677 多视频批量生成 — 变体独立配置与批量预览/批量生成测试。 + +覆盖: +- 批量预览:preview_count=N 一次创建 N 个独立任务,返回变体数组 +- 变体克隆链路:N 个预览/正式任务各自关联独立克隆 plan +- 变体独立配置:titles[]/voice_library_ids[]/cover_urls[] 按变体注入 +- 长度校验:数组长度必须为 1 或 N(共用或独立),非法长度报错 +- N=1 向后兼容:旧字段单值行为不变 +""" + +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch + +import pytest +from app.core.task_enqueue import GlobalQueueFull, UserPendingLimitExceeded +from app.schemas.generation_task import ( + BatchPreviewGenerationTaskResponse, + CreateGenerationTaskRequest, + CreatePreviewGenerationTaskRequest, +) + +from packages.domain import GenerationTask +from packages.domain.generation_task import GenerationTaskStatus + +# ════════════════════════════════════════════════════════════════════════════ +# 辅助构造 +# ════════════════════════════════════════════════════════════════════════════ + + +def _make_user(user_id="test_user_001"): + mock_user = MagicMock() + mock_user.id = user_id + auth = MagicMock() + auth.user = mock_user + return auth + + +def _make_task(task_id=None, status=GenerationTaskStatus.PENDING, source_plan_id=None): + task = GenerationTask.create( + project_id="", + asset_library_id="", + template_id="tpl_001", + asset_ids=["asset_1"], + ) + if task_id: + task.id = task_id + task.status = status + task.is_preview = True + task.source_edit_plan_id = source_plan_id or "" + task.voice_library_id = "" + task.title_config = {} + task.cover_url = "" + return task + + +def _make_preview_request(**kwargs): + defaults = { + "template_id": "tpl_001", + "asset_ids": ["asset_1", "asset_2"], + } + defaults.update(kwargs) + return CreatePreviewGenerationTaskRequest(**defaults) + + +def _repo_mock(): + repo = MagicMock() + repo.count_pending_by_user.return_value = 0 + repo.count_pending_total.return_value = 0 + repo.get.side_effect = lambda tid: None + return repo + + +# ════════════════════════════════════════════════════════════════════════════ +# Schema 校验:变体数组长度 +# ════════════════════════════════════════════════════════════════════════════ + + +class TestVariantArrayValidation: + """变体数组字段长度校验。""" + + def test_preview_titles_length_matches_count(self): + """titles 长度 = preview_count 合法""" + req = _make_preview_request(preview_count=3, titles=["标题A", "标题B", "标题C"]) + assert len(req.titles) == 3 + + def test_preview_titles_single_shared(self): + """titles 长度 1 = 所有变体共用,合法""" + req = _make_preview_request(preview_count=3, titles=["共用标题"]) + assert req.titles == ["共用标题"] + + def test_preview_titles_wrong_length_raises(self): + """titles 长度 2 与 preview_count=3 不匹配 → 报错""" + with pytest.raises(ValueError, match="titles"): + _make_preview_request(preview_count=3, titles=["A", "B"]) + + def test_preview_voice_ids_wrong_length_raises(self): + """voice_library_ids 长度非法 → 报错""" + from pydantic import ValidationError + + with pytest.raises(ValidationError, match="voice_library_ids"): + _make_preview_request(preview_count=4, voice_library_ids=["v1", "v2"]) + + def test_preview_empty_arrays_ok(self): + """空数组(回退单值字段)合法""" + req = _make_preview_request(preview_count=3) + assert req.titles == [] + assert req.voice_library_ids == [] + assert req.cover_urls == [] + + def test_generation_titles_length_matches_count(self): + """正式生成 titles 长度 = count 合法""" + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + titles=["A", "B", "C"], + ) + assert len(req.titles) == 3 + + def test_generation_arrays_wrong_length_raises(self): + """正式生成 cover_urls 长度与 count 不匹配 → 报错""" + from pydantic import ValidationError + + with pytest.raises(ValidationError, match="cover_urls"): + CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + cover_urls=["c1", "c2"], + ) + + def test_generation_single_count_no_arrays(self): + """N=1 且不传数组:完全旧行为""" + req = CreateGenerationTaskRequest(template_id="tpl_1", asset_ids=["a1"]) + assert req.count == 1 + assert req.titles == [] + assert req.voice_library_ids == [] + assert req.cover_urls == [] + + +# ════════════════════════════════════════════════════════════════════════════ +# 批量预览路由 +# ════════════════════════════════════════════════════════════════════════════ + + +class TestBatchPreviewRoute: + """POST /preview 批量变体。""" + + def test_preview_count_1_returns_single_item_array(self): + """N=1 返回 items 长度 1 的批量响应(结构统一)""" + from app.api.routes.generation_preview import create_preview_generation_task + + task = _make_task(task_id="task_1") + repo = _repo_mock() + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.return_value = task + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + resp = create_preview_generation_task( + _make_preview_request(preview_count=1), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert isinstance(resp, BatchPreviewGenerationTaskResponse) + assert resp.total == 1 + assert len(resp.items) == 1 + assert resp.items[0].task_id == "task_1" + assert resp.items[0].variant_index == 0 + + def test_preview_count_3_creates_three_independent_tasks(self): + """N=3 创建 3 个独立任务,返回 3 个变体,task_id 各不相同""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.side_effect = tasks + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + resp = create_preview_generation_task( + _make_preview_request(preview_count=3), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert resp.total == 3 + task_ids = [item.task_id for item in resp.items] + assert task_ids == ["task_0", "task_1", "task_2"] + assert len(set(task_ids)) == 3 + for i, item in enumerate(resp.items): + assert item.variant_index == i + + def test_preview_count_3_clones_three_variant_plans(self): + """有源 plan 时,N=3 克隆 3 个独立变体 plan(预览全部克隆,不用源 plan)""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)] + repo = _repo_mock() + cloned_plan_ids = ["clone_1", "clone_2", "clone_3"] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.side_effect = tasks + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc: + clone_results = [MagicMock(id=pid) for pid in cloned_plan_ids] + MockPlanSvc.return_value.clone_plan_for_variant.side_effect = clone_results + create_preview_generation_task( + _make_preview_request(preview_count=3), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + # 克隆被调用 3 次 + assert MockPlanSvc.return_value.clone_plan_for_variant.call_count == 3 + # 每个任务关联到不同的克隆 plan + for i, task in enumerate(tasks): + assert task.source_edit_plan_id == cloned_plan_ids[i] + + def test_preview_variant_titles_injected_per_variant(self): + """titles[] 按变体注入 title_config.text""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + captured_commands = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured_commands.append(cmd) + return tasks[len(captured_commands) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request( + preview_count=3, + title_config={"font": "黑体", "position": "bottom"}, + titles=["标题A", "标题B", "标题C"], + ), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert len(captured_commands) == 3 + assert captured_commands[0].title_config["text"] == "标题A" + assert captured_commands[1].title_config["text"] == "标题B" + assert captured_commands[2].title_config["text"] == "标题C" + # 样式全局共用 + assert all(c.title_config["font"] == "黑体" for c in captured_commands) + + def test_preview_shared_title_when_single_length(self): + """titles 长度 1 = 所有变体共用同一标题""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + captured = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return tasks[len(captured) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request(preview_count=3, titles=["共用标题"]), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert all(c.title_config["text"] == "共用标题" for c in captured) + + def test_preview_independent_voice_per_variant(self): + """voice_library_ids[] 按变体注入独立配音""" + from app.api.routes.generation_preview import create_preview_generation_task + + tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] + repo = _repo_mock() + captured = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return tasks[len(captured) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request( + preview_count=3, + voice_library_ids=["voice_a", "voice_b", "voice_c"], + ), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert [c.voice_library_id for c in captured] == ["voice_a", "voice_b", "voice_c"] + + def test_preview_voice_fallback_to_single_field(self): + """voice_library_ids 为空时回退 voice_library_id 单值字段(向后兼容)""" + from app.api.routes.generation_preview import create_preview_generation_task + + task = _make_task(task_id="task_1") + repo = _repo_mock() + captured = [] + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return task + + MockUC.return_value.execute.side_effect = _execute + with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): + create_preview_generation_task( + _make_preview_request(voice_library_id="legacy_voice"), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert captured[0].voice_library_id == "legacy_voice" + + def test_preview_queue_limit_checks_total_count(self): + """限流预检查按变体总数计:用户 pending + N 超限 → 429""" + from app.api.routes.generation_preview import create_preview_generation_task + from fastapi import HTTPException + + repo = MagicMock() + repo.count_pending_by_user.return_value = 3 + repo.count_pending_total.return_value = 0 + with pytest.raises(HTTPException) as exc: + create_preview_generation_task( + _make_preview_request(preview_count=5), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert exc.value.status_code == 429 + + def test_preview_clone_failure_marks_all_failed(self): + """克隆变体 plan 失败 → 已创建任务全部标记 failed 并 500""" + from app.api.routes.generation_preview import create_preview_generation_task + from fastapi import HTTPException + + tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)] + repo = _repo_mock() + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: + MockUC.return_value.execute.side_effect = tasks + with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc: + MockPlanSvc.return_value.clone_plan_for_variant.side_effect = RuntimeError("db down") + with pytest.raises(HTTPException) as exc: + create_preview_generation_task( + _make_preview_request(preview_count=3), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + assert exc.value.status_code == 500 + # 所有已创建任务都被标记 failed + assert all(t.status == GenerationTaskStatus.FAILED for t in tasks) + + +# ════════════════════════════════════════════════════════════════════════════ +# 批量正式生成:变体配置注入 +# ════════════════════════════════════════════════════════════════════════════ + + +class TestBatchGenerationVariantConfig: + """POST /tasks count=N 时变体独立配置。""" + + def _call_create_tasks(self, request, repo=None): + from app.api.routes.generation_tasks import create_generation_task + + repo = repo or MagicMock() + repo.count_pending_by_user.return_value = 0 + repo.count_pending_total.return_value = 0 + repo.update.return_value = None + + # 模板模式:asset_repository.find_by_id 返回 None(无 project 关联, + # 纯模板模式 project_id/library_id 都为空),避免 MagicMock 属性污染 + asset_repo = MagicMock() + asset_repo.find_by_id.return_value = None + + # db.query().filter()...first() 返回 None:不走兜底关联编辑计划 + db = MagicMock() + db.query.return_value.filter.return_value.order_by.return_value.first.return_value = None + + return create_generation_task( + request, + authenticated_user=_make_user(), + generation_task_repository=repo, + project_repository=MagicMock(), + asset_library_repository=MagicMock(), + asset_repository=asset_repo, + db=db, + ) + + def test_count_3_variant_titles_voices_covers_injected(self): + """count=3:titles/voice_library_ids/cover_urls 按变体注入""" + from app.api.routes import generation_tasks as routes + + tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)] + captured = [] + with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + t = tasks[len(captured) - 1] + t.title_config = cmd.title_config + t.voice_library_id = cmd.voice_library_id + t.cover_url = cmd.cover_url + return t + + MockUC.return_value.execute.side_effect = _execute + with patch.object(routes, "safe_enqueue_generation_task", return_value=True): + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + title_config={"font": "宋体"}, + titles=["成片标题1", "成片标题2", "成片标题3"], + voice_library_ids=["v1", "v2", "v3"], + cover_urls=["http://c1", "http://c2", "http://c3"], + ) + resp = self._call_create_tasks(req) + assert resp.total == 3 + assert [c.title_config["text"] for c in captured] == ["成片标题1", "成片标题2", "成片标题3"] + assert [c.voice_library_id for c in captured] == ["v1", "v2", "v3"] + assert [c.cover_url for c in captured] == ["http://c1", "http://c2", "http://c3"] + # 样式共用 + assert all(c.title_config["font"] == "宋体" for c in captured) + + def test_count_1_legacy_fields_unchanged(self): + """N=1 不传数组:旧字段 voice_library_id/cover_url/title_config 行为不变""" + from app.api.routes import generation_tasks as routes + + task = _make_task(task_id="gen_1") + task.is_preview = False + captured = [] + with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return task + + MockUC.return_value.execute.side_effect = _execute + with patch.object(routes, "safe_enqueue_generation_task", return_value=True): + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=1, + voice_library_id="legacy_voice", + cover_url="http://legacy-cover", + title_config={"text": "旧标题", "font": "黑体"}, + ) + resp = self._call_create_tasks(req) + assert resp.total == 1 + assert captured[0].voice_library_id == "legacy_voice" + assert captured[0].cover_url == "http://legacy-cover" + assert captured[0].title_config["text"] == "旧标题" + + def test_count_3_shared_single_value_arrays(self): + """数组长度 1:3 个变体共用同一配音/封面""" + from app.api.routes import generation_tasks as routes + + tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)] + captured = [] + with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: + + def _execute(cmd): + captured.append(cmd) + return tasks[len(captured) - 1] + + MockUC.return_value.execute.side_effect = _execute + with patch.object(routes, "safe_enqueue_generation_task", return_value=True): + req = CreateGenerationTaskRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + voice_library_ids=["shared_voice"], + cover_urls=["http://shared"], + ) + self._call_create_tasks(req) + assert all(c.voice_library_id == "shared_voice" for c in captured) + assert all(c.cover_url == "http://shared" for c in captured) + + +class TestVariantValueHelper: + """_variant_value 取值逻辑。""" + + def test_empty_returns_fallback(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value([], 0, fallback="fb") == "fb" + + def test_single_length_shared(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value(["only"], 5) == "only" + + def test_indexed_access(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value(["a", "b", "c"], 1) == "b" + + def test_index_out_of_range_fallback(self): + from app.api.routes.generation_preview import _variant_value + + assert _variant_value(["a", "b"], 9, fallback="x") == "x" diff --git a/tests/unit/test_asset_repo_fallback_dedup_1714.py b/tests/unit/test_asset_repo_fallback_dedup_1714.py new file mode 100644 index 000000000..d68890773 --- /dev/null +++ b/tests/unit/test_asset_repo_fallback_dedup_1714.py @@ -0,0 +1,93 @@ +"""#1714 find_recent_active_by_library_and_name 严格模式测试。 + +file_size=0(未知)时必须返回 None(宁可漏判不可误杀); +大小严格匹配;只命中近期 UPLOADING/PROCESSING 记录。 +""" + +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from sqlalchemy import create_engine # noqa: E402 +from sqlalchemy.orm import sessionmaker # noqa: E402 + +from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository # noqa: E402 +from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402 +from packages.domain import Asset, AssetStatus # noqa: E402 + + +def _repository(): + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(engine) + session = sessionmaker(bind=engine)() + return SQLAlchemyAssetRepository(session) + + +def _mk_asset(name="IMG_2285.MOV", file_size=5_000_000, status=AssetStatus.PROCESSING, minutes_ago=5): + asset = Asset.create( + project_id="proj-1", + library_id="lib-1", + name=name, + storage_key=f"uploads/x/{name}", + mime_type="video/quicktime", + file_size=file_size, + ) + asset.status = status + asset.created_at = datetime.now(timezone.utc) - timedelta(minutes=minutes_ago) + return asset + + +def test_returns_none_when_file_size_zero(): + """file_size=0(大小未知)直接返回 None——不许仅凭同名 + processing 判重。""" + repo = _repository() + repo.create(_mk_asset(file_size=0)) + + result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=0) + assert result is None + + +def test_matches_when_name_size_strict_equal(): + """同名 + 同大小 + processing 近期记录 → 命中。""" + repo = _repository() + repo.create(_mk_asset(file_size=5_000_000)) + + result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=5_000_000) + assert result is not None + assert result.name == "IMG_2285.MOV" + + +def test_no_match_when_same_name_but_different_size(): + """同名但大小不同 → 不命中(内容全新的视频不能误杀)。""" + repo = _repository() + repo.create(_mk_asset(file_size=5_000_000)) + + result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=9_999_999) + assert result is None + + +def test_no_match_ready_history_even_with_same_size(): + """READY 历史同名素材不命中(允许再次上传同名文件)。""" + repo = _repository() + repo.create(_mk_asset(file_size=5_000_000, status=AssetStatus.READY)) + + result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=5_000_000) + assert result is None + + +def test_no_match_when_window_expired(): + """超过 30 分钟窗口的活动记录不命中。""" + repo = _repository() + repo.create(_mk_asset(file_size=5_000_000, minutes_ago=45)) + + result = repo.find_recent_active_by_library_and_name( + library_id="lib-1", name="IMG_2285.MOV", within_minutes=30, file_size=5_000_000 + ) + assert result is None + + +def test_returns_none_when_name_empty(): + repo = _repository() + result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="", file_size=100) + assert result is None diff --git a/tests/unit/test_bad_fingerprint_filter.py b/tests/unit/test_bad_fingerprint_filter.py index 05a18cb08..975a1ae8a 100644 --- a/tests/unit/test_bad_fingerprint_filter.py +++ b/tests/unit/test_bad_fingerprint_filter.py @@ -85,17 +85,18 @@ class TestIsBadFingerprint: assert VideoDeduplicator._is_bad_fingerprint(["abcdef0123456789"]) is False def test_all_identical_phashes_is_bad(self): - """多帧但所有 phash 完全相同 → 黑屏/纯色视频。""" - phashes = ["aaaaaaaaaaaaaaaa"] * 5 + """>=8 帧且所有 phash 完全相同 → 黑屏/纯色视频(#1702:短帧不误杀)。""" + phashes = ["aaaaaaaaaaaaaaaa"] * 10 assert VideoDeduplicator._is_bad_fingerprint(phashes) is True - def test_two_identical_phashes_is_bad(self): - """两帧完全相同也视为坏指纹。""" - assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb", "bbbbbbbbbbbbbbbb"]) is True + def test_short_identical_phashes_not_bad(self): + """<8 帧完全相同不判坏——短视频内容连续时相邻采样帧 phash 天然相同(#1702)。""" + assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb"] * 5) is False + assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb", "bbbbbbbbbbbbbbbb"]) is False def test_all_very_similar_phashes_is_bad(self): - """多帧 phash 之间的汉明距离都 < 3 → 近似黑屏。""" - phashes = ["0000000000000000", "0000000000000001", "0000000000000002"] + """>=8 帧 phash 之间的汉明距离都 < 3 且高占比 → 近似黑屏。""" + phashes = ["0000000000000000"] * 8 + ["0000000000000001", "0000000000000002"] assert VideoDeduplicator._is_bad_fingerprint(phashes) is True def test_diverse_phashes_is_good(self): @@ -122,7 +123,9 @@ class TestIsBadFingerprint: """已知黑屏视频的 phash 特征(全零或均匀分布)。""" assert VideoDeduplicator._is_bad_fingerprint(["0000000000000000"] * 10) is True assert VideoDeduplicator._is_bad_fingerprint(["ffffffffffffffff"] * 8) is True - assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 6) is True + assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 8) is True + # <8 帧不判坏(#1702 短视频保护) + assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 5) is False # ── Helper ────────────────────────────────────────────────────── @@ -151,13 +154,13 @@ class TestCheckDuplicateBadFingerprint: deduplicator = VideoDeduplicator() mock_session = MagicMock() - black_screen = _make_existing_video("vid-black", "md5_black", ["aaaaaaaaaaaaaaaa"] * 5) + black_screen = _make_existing_video("vid-black", "md5_black", ["aaaaaaaaaaaaaaaa"] * 10) mock_repo = MagicMock() mock_repo.list_by_user.return_value = [black_screen] fingerprint = VideoFingerprint( md5="md5_normal", - keyframe_phashes=["aaaaaaaaaaaaaaaa"] * 5, + keyframe_phashes=["aaaaaaaaaaaaaaaa"] * 10, color_histograms=[], duration=10.0, resolution=(1280, 720), @@ -206,7 +209,7 @@ class TestCheckDuplicateBadFingerprint: deduplicator = VideoDeduplicator() mock_session = MagicMock() - black_screen = _make_existing_video("vid-black", "same_md5", ["aaaaaaaaaaaaaaaa"] * 5) + black_screen = _make_existing_video("vid-black", "same_md5", ["aaaaaaaaaaaaaaaa"] * 10) mock_repo = MagicMock() mock_repo.list_by_user.return_value = [black_screen] @@ -283,8 +286,8 @@ class TestComputeDuplicateRateBadFingerprint: mock_session = MagicMock() videos = [ - _make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 5), - _make_existing_video("vid-b2", "md5_b2", ["bbbbbbbbbbbbbbbb"] * 5), + _make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 10), + _make_existing_video("vid-b2", "md5_b2", ["cccccccccccccccc"] * 5), # hamming(a,c)=32 > PHASH_THRESHOLD ] mock_repo = MagicMock() mock_repo.list_by_user.return_value = videos diff --git a/tests/unit/test_celery_queue_isolation_1714.py b/tests/unit/test_celery_queue_isolation_1714.py new file mode 100644 index 000000000..1c473eada --- /dev/null +++ b/tests/unit/test_celery_queue_isolation_1714.py @@ -0,0 +1,176 @@ +"""#1714 队列隔离 + 作废消息清除 单元测试。 + +覆盖: +1. task_routes:generate_video → generation,ingest_asset/classify/duplication → transcode +2. purge_stale_messages_from_queues:Redis 队列中作废任务消息被物理移除,未命中保留 +3. revoke_and_purge:revoke 广播 + 队列清理同时生效 +4. ensure_task_claimable:终态任务抛 StaleTaskDiscarded,pending 放行 +""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest +from celery import Celery + +from packages.shared.celery_orphan_guard import ( + StaleTaskDiscarded, + _extract_business_ids, + ensure_task_claimable, + purge_stale_messages_from_queues, + revoke_and_purge, +) +from packages.shared.celery_queues import ( + QUEUE_GENERATION, + QUEUE_TRANSCODE, + apply_queue_settings, + task_routes, +) + +BROKER_URL = "redis://localhost:6379/15" +TEST_QUEUES = ("_test_gen_q", "_test_transcode_q") + + +# ── 1. 路由表 ────────────────────────────────────────────────────────── + + +def test_routes_send_generation_to_generation_queue(): + assert task_routes["worker.generate_video"]["queue"] == QUEUE_GENERATION + + +def test_routes_send_ingest_to_transcode_queue(): + assert task_routes["worker.ingest_asset"]["queue"] == QUEUE_TRANSCODE + assert task_routes["worker.classify_asset"]["queue"] == QUEUE_TRANSCODE + assert task_routes["worker.process_duplication_check"]["queue"] == QUEUE_TRANSCODE + assert task_routes["worker.check_duplicate"]["queue"] == QUEUE_TRANSCODE + + +def test_apply_queue_settings_configures_celery_app(): + app = Celery("test-routes") + apply_queue_settings(app) + queue_names = {q.name for q in app.conf.task_queues} + assert queue_names == {"generation", "transcode", "celery"} + assert app.conf.task_default_queue == "celery" + + +# ── Redis 队列消息清理(需要本地 redis;不可用时 skip) ───────────────── + + +def _redis_available() -> bool: + try: + import redis + + return bool(redis.Redis.from_url(BROKER_URL).ping()) + except Exception: + return False + + +@pytest.fixture() +def redis_client(): + import redis + + client = redis.Redis.from_url(BROKER_URL) + for q in TEST_QUEUES: + client.delete(q) + yield client + for q in TEST_QUEUES: + client.delete(q) + + +def _publish(app: Celery, queue: str, celery_id: str, business_id: str) -> None: + from kombu import Queue + from kombu.pools import producers + + with app.connection_for_write() as conn: + with producers[conn].acquire(block=True) as prod: + prod.publish( + (business_id,), + exchange="", + routing_key=queue, + serializer="json", + headers={"id": celery_id, "task": "worker.generate_video"}, + retry=False, + delivery_mode=1, + declare=[Queue(queue, routing_key=queue, durable=False)], + ) + + +@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用") +def test_purge_removes_stale_business_message_and_keeps_others(redis_client): + app = Celery("test-purge") + app.conf.broker_url = BROKER_URL + _publish(app, TEST_QUEUES[0], "celery-1", "task-KEEP-A") + _publish(app, TEST_QUEUES[0], "celery-2", "task-STALE-B") + _publish(app, TEST_QUEUES[0], "celery-3", "task-KEEP-C") + _publish(app, TEST_QUEUES[1], "celery-4", "task-STALE-B") # 同一业务任务在转码队列?不应出现但验证全队列扫描 + + removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, business_task_ids={"task-STALE-B"}) + assert removed == 2 + + remaining = [] + for raw in redis_client.lrange(TEST_QUEUES[0], 0, -1): + _celery_id, biz_id = _extract_business_ids(raw) + remaining.append(biz_id) + assert set(remaining) == {"task-KEEP-A", "task-KEEP-C"} + assert redis_client.llen(TEST_QUEUES[1]) == 0 + + +@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用") +def test_purge_matches_by_celery_message_id(redis_client): + app = Celery("test-purge-msg-id") + app.conf.broker_url = BROKER_URL + _publish(app, TEST_QUEUES[0], "celery-stale-id", "task-X") + _publish(app, TEST_QUEUES[0], "celery-good-id", "task-Y") + + removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, celery_task_ids={"celery-stale-id"}) + assert removed == 1 + assert redis_client.llen(TEST_QUEUES[0]) == 1 + + +@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用") +def test_revoke_and_purge_calls_control_revoke(redis_client): + app = Celery("test-revoke") + app.conf.broker_url = BROKER_URL + app.control = MagicMock() + _publish(app, TEST_QUEUES[0], "celery-revoke-1", "task-R") + + removed = revoke_and_purge( + app, + BROKER_URL, + business_task_ids={"task-R"}, + celery_task_ids={"celery-revoke-1"}, + queue_names=TEST_QUEUES, + ) + assert removed == 1 + app.control.revoke.assert_called_once_with("celery-revoke-1") + + +def test_purge_empty_ids_is_noop(): + assert purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES) == 0 + + +# ── 2. 执行前状态守卫 ────────────────────────────────────────────────── + + +def test_guard_allows_pending(): + status = ensure_task_claimable("t1", lambda _id: "pending", task_label="generation") + assert status == "pending" + + +def test_guard_rejects_failed(): + with pytest.raises(StaleTaskDiscarded) as exc: + ensure_task_claimable("t2", lambda _id: "failed", task_label="generation") + assert exc.value.task_id == "t2" + assert exc.value.status == "failed" + + +def test_guard_rejects_cancelled_and_completed(): + with pytest.raises(StaleTaskDiscarded): + ensure_task_claimable("t3", lambda _id: "cancelled") + with pytest.raises(StaleTaskDiscarded): + ensure_task_claimable("t4", lambda _id: "completed") + + +def test_guard_missing_task_returns_empty(): + assert ensure_task_claimable("t5", lambda _id: None) == "" diff --git a/tests/unit/test_cleanup_ingest_beat_1714.py b/tests/unit/test_cleanup_ingest_beat_1714.py new file mode 100644 index 000000000..f7c5bf11d --- /dev/null +++ b/tests/unit/test_cleanup_ingest_beat_1714.py @@ -0,0 +1,75 @@ +"""#1714 beat 任务 scheduled_cleanup_stale_ingest_jobs 薄封装测试。 + +mock SessionLocal 和清理核心,验证 beat 任务正确串联 +cleanup_stale_ingest_jobs → cleanup_orphan_processing_assets → revoke 消息。 +""" + +from __future__ import annotations + +import os +import sys +from pathlib import Path +from unittest.mock import MagicMock, patch + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret") +os.environ.setdefault("DATABASE_URL", "sqlite:///test_beat.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker")) +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) +sys.path.insert(0, str(Path(__file__).resolve().parents[2])) + +import worker_app.tasks.cleanup as cleanup # noqa: E402 + + +def test_beat_cleanup_calls_core_and_revokes(): + """beat 任务串联三个核心步骤,返回汇总计数。""" + fake_session = MagicMock() + + with ( + patch("worker_app.db.SessionLocal", return_value=fake_session) as m_db, + patch( + "packages.application.ingest_orphan_cleanup.cleanup_stale_ingest_jobs", + return_value=([("job-1", "cel-1"), ("job-2", "")], ["a-1"]), + ) as m_jobs, + patch( + "packages.application.ingest_orphan_cleanup.cleanup_orphan_processing_assets", + return_value=["a-2"], + ) as m_assets, + patch( + "packages.shared.celery_orphan_guard.revoke_and_purge", + return_value=1, + ) as m_revoke, + ): + result = cleanup.scheduled_cleanup_stale_ingest_jobs() + + m_db.assert_called_once() + m_jobs.assert_called_once() + assert m_jobs.call_args.kwargs["processing_timeout_minutes"] == 60 + m_assets.assert_called_once() + m_revoke.assert_called_once() + # 队列名只传 transcode/celery(不传 generation) + assert m_revoke.call_args.kwargs["queue_names"] == ("transcode", "celery") + fake_session.close.assert_called_once() + assert result == {"stale_jobs": 2, "assets_to_error": 2, "purged_messages": 1} + + +def test_beat_cleanup_no_op_when_nothing_stale(): + """无孤儿时不调 revoke,返回全 0。""" + fake_session = MagicMock() + + with ( + patch("worker_app.db.SessionLocal", return_value=fake_session), + patch( + "packages.application.ingest_orphan_cleanup.cleanup_stale_ingest_jobs", + return_value=([], []), + ), + patch( + "packages.application.ingest_orphan_cleanup.cleanup_orphan_processing_assets", + return_value=[], + ), + patch("packages.shared.celery_orphan_guard.revoke_and_purge") as m_revoke, + ): + result = cleanup.scheduled_cleanup_stale_ingest_jobs() + + m_revoke.assert_not_called() + assert result == {"stale_jobs": 0, "assets_to_error": 0, "purged_messages": 0} diff --git a/tests/unit/test_dedup_1702_zero_rate_fix.py b/tests/unit/test_dedup_1702_zero_rate_fix.py new file mode 100644 index 000000000..49f4e4871 --- /dev/null +++ b/tests/unit/test_dedup_1702_zero_rate_fix.py @@ -0,0 +1,548 @@ +"""Issue #1702 — 查重率恒为 0% 修复:单测. + +覆盖验收要求: +1. 同源不同裁剪的两个视频能检出非 0 相似度(指纹中心裁剪绕开降重 + 阈值校准) +2. 局部片段复用(B 结尾 2s ≈ A 中间 2s)能检出 +3. 异源视频不误报(相似度接近 0) +4. N=1 现有流程不回归 +5. P1 确定性 bug:时长预过滤单位 /1000、直方图归一化、temporal_coverage 量纲、阈值比较统一 +6. P0:±1 邻接对齐、短视频自适应连续门槛 +7. P2:0 匹配也要落日志 +""" + +from __future__ import annotations + +import logging +import sys +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +sys.modules.setdefault("cv2", MagicMock()) + +ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(ROOT / "apps" / "worker")) +sys.path.insert(0, str(ROOT / "packages")) + + +from video_processing.dedup import ( # noqa: E402 + PHASH_THRESHOLD, + SEGMENT_MATCH_THRESHOLD, + FingerprintChunk, + VideoDeduplicator, + VideoFingerprint, + find_duplicate_segments, +) + +# ── helpers ──────────────────────────────────────────────────── + + +def _h(d: int) -> str: + """64-bit phash with exactly d bits set vs zero hash.""" + bits = ["0"] * 64 + for i in range(d): + bits[i] = "1" + return f"{int(''.join(bits), 2):016x}" + + +def _chunk(phash: str, t0: float, t1: float): + + return FingerprintChunk( + start_time_ms=int(t0 * 1000), + end_time_ms=int(t1 * 1000), + phash_binary=phash, + color_histogram=[], + frame_count=1, + ) + + +def _fingerprint(phashes, duration, chunks=None, md5="fp-md5-x"): + + return VideoFingerprint( + md5=md5, + keyframe_phashes=list(phashes), + color_histograms=[], + duration=duration, + resolution=(1280, 720), + chunks=chunks or [], + ) + + +def _video(vid, phashes, duration=10.0, project_id="proj1"): + from packages.domain import GeneratedVideo + + return GeneratedVideo( + id=vid, + project_id=project_id, + generation_task_id=f"task-{vid}", + name=f"video-{vid}.mp4", + file_url=f"https://example.com/{vid}.mp4", + file_size=1000, + duration=duration, + width=1280, + height=720, + fps=25.0, + video_fingerprint={"md5": f"md5-{vid}", "keyframe_phashes": list(phashes)}, + ) + + +def _rate(deduplicator, fp, videos, session=None): + session_magic = MagicMock() + # 分片表无数据 -> 回退 JSON keyframe_phashes + session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = [] + with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo: + repo = MockRepo.return_value + repo.list_by_project.return_value = videos + repo.list_by_user.return_value = videos + return deduplicator.compute_duplicate_rate(fp, "proj1", "new-vid", session_magic, scope="project") + + +def _check(deduplicator, fp, videos, scope="project", **kw): + session_magic = MagicMock() + session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = [] + with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo: + repo = MockRepo.return_value + repo.list_by_project.return_value = videos + repo.list_by_user.return_value = videos + return deduplicator.check_duplicate(fp, "proj1", session_magic, scope=scope, **kw) + + +# ── P0-1/P0-2: 同源不同裁剪(距离 6~10)检出非 0 ────────────── + + +class TestSameSourceDifferentCrop: + """同源成片:random_edge_crop 后 pHash 距离 6~10,应检出非 0 相似度。""" + + def test_same_source_high_similarity_detected(self): + + ddp = VideoDeduplicator() + # 新视频 5 个分片,每个 phash 与已有视频对应分片距离 6(< 阈值) + base = [_h(0) for _ in range(5)] + new = [_h(6) for _ in range(5)] + existing = _video("v-old", base, duration=11.0) + chunks = [_chunk(h, i * 2.2, (i + 1) * 2.2) for i, h in enumerate(new)] + fp = _fingerprint(new, 11.0, chunks=chunks) + + result = _rate(ddp, fp, [existing], MagicMock()) + assert result["duplicate_rate"] > 0 + assert result["visual_similarity"] > 0 + + def test_same_source_distance_at_threshold_still_detected(self): + """距离正好等于阈值(<=)也要算匹配——阈值比较统一为 <=。""" + + assert PHASH_THRESHOLD <= 16, "阈值应经真实数据校准保持在能检出同源裁剪/降重对的范围(#1702 二次校准为 16)" + ddp = VideoDeduplicator() + base = [_h(0) for _ in range(6)] + new = [_h(PHASH_THRESHOLD) for _ in range(6)] + existing = _video("v-old", base, duration=12.0) + chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)] + fp = _fingerprint(new, 12.0, chunks=chunks) + + result = _rate(ddp, fp, [existing], MagicMock()) + assert result["duplicate_rate"] > 0 + + +# ── P0-2: 局部片段复用(B 结尾 2s ≈ A 中间 2s) ──────────────── + + +class TestPartialReuse: + def test_partial_reuse_tail_overlap_detected(self): + """新视频 6 片,最后 2 片命中已有视频中间 2 片(距离 4),其余不匹配。 + + 旧逻辑 frame_match_rate=2/6≈0.33(<0.3 硬跳过边界)+ MIN_CONSECUTIVE=5 + 导致完全检不出;新逻辑 coverage 为主指标 + 自适应门槛应检出。 + """ + + ddp = VideoDeduplicator() + # 已有 8 片:索引 3、4 是被复用的镜头 + old = [_h(20 + i) for i in range(8)] + # 新视频 6 片:最后 2 片对应 old[3], old[4],距离 4;其余距离 30 + new = [_h(50 + i) for i in range(4)] + [_h(4)] * 2 + # 让 new[4] 与 old[3] 距离 4、new[5] 与 old[4] 距离 4(构造近似) + new[4] = f"{int('1' * 4 + '0' * 60, 2):016x}" + new[5] = f"{int('1' * 4 + '0' * 60, 2):016x}" + old[3] = _h(0) + old[4] = _h(0) + + existing = _video("v-old", old, duration=16.0) + chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)] + fp = _fingerprint(new, 12.0, chunks=chunks) + + result = _rate(ddp, fp, [existing], MagicMock()) + # 局部复用:duplicate_rate 必须非 0 + assert result["duplicate_rate"] > 0 + + def test_short_video_adaptive_consecutive_threshold(self): + """11s/5 片短视频:MIN_CONSECUTIVE 自适应 min(5, max(2, 5//2))=2, + 2 片连续命中即报片段(旧值 5 让短视频永远无法报片段)。""" + + q = [ + FingerprintChunk(0, 2000, "f" * 16, []), + FingerprintChunk(2000, 4000, "0" * 16, []), + FingerprintChunk(4000, 6000, f"{int('11110000', 2):016x}", []), + ] + t = [ + FingerprintChunk(0, 2000, "f" * 16, []), + FingerprintChunk(2000, 4000, "0" * 16, []), + FingerprintChunk(4000, 6000, "e" * 16, []), + ] + # 3 片视频自适应门槛 = min(5, max(2, 3//2)) = 2 + segs = find_duplicate_segments(q, t) + assert len(segs) >= 1 + + +# ── P0-3: ±1 邻接窗口对齐 ───────────────────────────────────── + + +class TestNeighborAlignment: + def test_neighbor_window_absorbs_boundary_jitter(self): + """切点错位导致目标索引偏移 ±1 时,连续匹配不应被中断。""" + + q = [FingerprintChunk(i * 1000, (i + 1) * 1000, f"{i:016x}", []) for i in range(4)] + # 目标:前 3 片与 q 相同,但第 3 片最佳匹配偏移 +1(t[4]),t[3] 是无关内容 + t_hashes = [f"{i:016x}" for i in range(3)] + ["f" * 16, f"{3:016x}"] + t = [FingerprintChunk(i * 1000, (i + 1) * 1000, h, []) for i, h in enumerate(t_hashes)] + segs = find_duplicate_segments(q, t) + # q[0],q[1] 精确匹配 t[0],t[1];q[2]->t[2];q[3]->t[4](步进 2,窗口 ±1 内) + assert len(segs) >= 1 + assert segs[0].query_end_ms >= 3000 + + +# ── P0-5 / 验收:异源不误报 ─────────────────────────────────── + + +class TestDifferentSourceNoFalsePositive: + def test_unrelated_videos_near_zero(self): + + ddp = VideoDeduplicator() + # 异源:所有分片距离 >= 20 + old = [_h(40 + i * 3 % 20) for i in range(6)] + new = [_h(0 + i) for i in range(6)] + existing = _video("v-old", old, duration=12.0) + chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)] + fp = _fingerprint(new, 12.0, chunks=chunks) + + result = _rate(ddp, fp, [existing], MagicMock()) + assert result["duplicate_rate"] == 0 + assert result["visual_similarity"] < 0.7 + assert result["match_count"] == 0 + + def test_check_duplicate_returns_none_for_unrelated(self): + + ddp = VideoDeduplicator() + old = [_h(40 + i) for i in range(6)] + new = [_h(i) for i in range(6)] + existing = _video("v-old", old, duration=12.0) + fp = _fingerprint(new, 12.0) + + result = _check(ddp, fp, [existing]) + assert result is None + + +# ── N=1 不回归 ──────────────────────────────────────────────── + + +class TestSingleChunkNoRegression: + def test_single_chunk_identical_detected(self): + + ddp = VideoDeduplicator() + h = _h(2) + existing = _video("v-old", [h], duration=3.0) + chunks = [_chunk(h, 0, 3000)] + fp = _fingerprint([h], 3.0, chunks=chunks) + result = _rate(ddp, fp, [existing], MagicMock()) + assert result["duplicate_rate"] > 0 + + def test_single_chunk_md5_exact_match(self): + + ddp = VideoDeduplicator() + existing = _video("v-old", [_h(0)], duration=3.0) + existing.video_fingerprint["md5"] = "same" + fp = _fingerprint([_h(0)], 3.0, md5="same") + result = _check(ddp, fp, [existing]) + assert result is not None + assert result["reason"] == "exact_md5_match" + + +# ── P1-6: 时长预过滤单位 bug ────────────────────────────────── + + +class TestDurationPrefilterUnit: + def test_user_scope_skips_duration_prefilter(self): + """Issue #1702: scope=user 跨项目查重不做 ±15% 时长预过滤。 + + 旧逻辑 duration/1000 单位 bug 先修成秒,但 ±15% 窗口与局部片段复用 + 根本矛盾——复用片段的两个视频时长必然不同(证据视频 20s vs 11s 差 42%), + 窗口内找不到对方导致 is_duplicate 恒 False。最终口径:scope=user 全量 + 遍历同用户视频(与 compute_duplicate_rate 一致),不传 duration_min/max。 + """ + + ddp = VideoDeduplicator() + fp = _fingerprint([_h(0)], 13.5) + session_magic = MagicMock() + session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = [] + with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo: + repo = MockRepo.return_value + repo.list_by_user.return_value = [] + ddp.check_duplicate(fp, "proj1", session_magic, scope="user", user_id="u1", duration_sec=fp.duration) + args, kwargs = repo.list_by_user.call_args + # 全量查询:不带任何时长过滤参数(局部复用必须跨时长比较) + assert "duration_min" not in kwargs + assert "duration_max" not in kwargs + assert args == ("u1",) or args == () + + +# ── P1-7: 颜色直方图归一化 ──────────────────────────────────── + + +class TestHistogramNormalization: + def test_bhattacharyya_coefficient_in_unit_range(self): + """Bhattacharyya 系数必须在 [0,1](旧 L2 + 3 通道拼接算出 ~14.9)。""" + + # 3 通道拼接、每通道概率分布(Σ=1) + hist_a = [0.5, 0.5] + [0.0] * 94 + [0.5, 0.5] + [0.0] * 94 + [0.5, 0.5] + [0.0] * 94 + # 长度裁剪到 96(3 通道 × 32 bins) + hist_a = ([0.5, 0.5] + [0.0] * 30) * 3 + hist_b = ([0.5, 0.5] + [0.0] * 30) * 3 + + coeff = VideoDeduplicator._bhattacharyya_coefficient(hist_a, hist_b) + assert 0.0 <= coeff <= 1.0 + assert coeff > 0.99 # 完全相同 -> 1.0 + + def test_bhattacharyya_disjoint_hist_low(self): + + hist_a = ([1.0] + [0.0] * 31) * 3 + hist_b = ([0.0] * 31 + [1.0]) * 3 + coeff = VideoDeduplicator._bhattacharyya_coefficient(hist_a, hist_b) + assert coeff < 0.05 + + +# ── P1-8: temporal_coverage 量纲 ────────────────────────────── + + +class TestTemporalCoverageUnits: + def test_coverage_uses_milliseconds(self): + """命中片段 6s / 视频 12s -> coverage=0.5;旧 bug 把 duration(秒)当毫秒, + covered_ms(6000)/duration(12) = 500 -> min(1.0)=1.0 误判 100% 覆盖。""" + + ddp = VideoDeduplicator() + old = [_h(0) for _ in range(6)] + new = [_h(0) for _ in range(3)] + [_h(30) for _ in range(3)] + existing = _video("v-old", old, duration=12.0) + # 新视频 12s,前 6s(3 片)与 old 相同 + chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)] + fp = _fingerprint(new, 12.0, chunks=chunks) + result = _rate(ddp, fp, [existing], MagicMock()) + # coverage 应约 0.5(3 片 × 2s = 6s / 12s),duplicate_rate ≈ (0.5*0.4 + 0.5*0.6)*100 = 50 + assert 30 < result["duplicate_rate"] < 70 + + +# ── P1-9: 阈值比较统一 ──────────────────────────────────────── + + +class TestThresholdConsistency: + def test_frame_and_segment_thresholds_same_source(self): + + assert SEGMENT_MATCH_THRESHOLD == PHASH_THRESHOLD + assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD + + +# ── P2: 0 匹配也要有日志痕迹 ────────────────────────────────── + + +class TestZeroMatchLogging: + def test_no_match_emits_info_log(self, caplog): + + ddp = VideoDeduplicator() + old = [_h(40 + i) for i in range(5)] + existing = _video("v-old", old, duration=10.0) + fp = _fingerprint([_h(i) for i in range(5)], 10.0) + + with caplog.at_level(logging.INFO, logger="video_processing.dedup"): + result = _check(ddp, fp, [existing]) + assert result is None + assert any("no match" in r.message for r in caplog.records) + + +# ── recompute 任务下载路径(#1702 连带修复:旧硬编码 key 404) ───── + + +class TestRecomputeDownloadPath: + """recompute-dedup 走 check_duplicate_task,需要从 OSS 重新下载成片。 + + 旧代码硬编码 projects/{pid}/generated/{vid}/{vid}.mp4(从不存在), + 真实 key 在 file_url:generated/projects/{pid}/tasks/{tid}/rendered_*.mp4。 + """ + + def test_task_downloads_from_file_url(self): + import inspect + + import video_processing.dedup as dedup_mod + + source = inspect.getsource(dedup_mod.check_duplicate_task) + # 下载 key 必须来自 video.file_url + assert 'getattr(video, "file_url"' in source or "video.file_url" in source + # 旧的硬编码 key 只能作为回退存在,不能是主路径 + assert "falling back to legacy key" in source + # download_file 接收的是派生 key 而非硬编码 f-string + assert "storage_service.download_file(download_key" in source + assert '/generated/{generated_video_id}/{generated_video_id}.mp4"' not in source.replace( + 'download_key = f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4"', + "", + ) + + +# ── check_duplicate 排除自身(#1702 连带修复:recompute 自匹配) ───── + + +class TestCheckDuplicateExcludesSelf: + def test_exclude_video_id_skips_self_match(self): + """recompute 时当前视频已在候选列表:自匹配距离 0 分会让 duplicate_of + 指向自己。exclude_video_id 必须跳过自身,返回真实的其他匹配或 None。 + """ + + ddp = VideoDeduplicator() + h = _h(0) + # 候选列表里同时放「自己」(完全相同)和一个异源视频 + self_video = _video("v-self", [h], duration=10.0) + other_video = _video("v-other", [_h(40 + i) for i in range(3)], duration=10.0) + fp = _fingerprint([h], 10.0) + + session = MagicMock() + session.query.return_value.filter.return_value.order_by.return_value.all.return_value = [] + + # 不传 exclude → 自匹配命中(错误行为复现) + with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo: + MockRepo.return_value.list_by_project.return_value = [self_video, other_video] + result = ddp.check_duplicate(fp, "proj1", session) + assert result is not None and result["duplicate_of"] == "v-self" + + # 传 exclude_video_id → 跳过自己,异源不匹配 → None + with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo: + MockRepo.return_value.list_by_project.return_value = [self_video, other_video] + result = ddp.check_duplicate(fp, "proj1", session, exclude_video_id="v-self") + assert result is None + + # 排除自己后,真实同源其他视频仍能检出 + real_dup = _video("v-real", [h], duration=10.0) + with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo: + MockRepo.return_value.list_by_project.return_value = [self_video, real_dup] + result = ddp.check_duplicate(fp, "proj1", session, exclude_video_id="v-self") + assert result is not None and result["duplicate_of"] == "v-real" + + +# ── 阈值 16 二次校准 + 时序抖动对齐(#1702 第二轮真实数据校准) ────── + + +class TestThreshold16Calibration: + """二次校准:staging 15 个真实成片实测——同源降重对中位数距离 14、 + <=16 命中 8/11=0.73;异源 13 个候选每帧全局最近邻最小距离 18、<=16 + 命中全 0。阈值 16 检出同源且异源零误报(>=2bit 安全裕度)。""" + + def test_threshold_calibrated_to_16(self): + assert PHASH_THRESHOLD == 16 + + @staticmethod + def _variant(phash: str, d: int) -> str: + """在 phash 基础上翻转恰好 d 个低位 bit → 与原哈希汉明距离恰为 d。""" + v = int(phash, 16) + for b in range(d): + v ^= 1 << b + return f"{v:016x}" + + def test_distance_18_unrelated_not_matched(self): + """距离 18(异源实测最小最近邻距离)不判匹配,距离 16 判匹配。""" + ddp = VideoDeduplicator() + # 多样化 base(相邻帧各不相同,避免黑屏过滤器) + base = [_h(i + 4) for i in range(8)] + near = [self._variant(h, 16) for h in base] # 同源降重:每帧距离恰 16 + far = [self._variant(h, 18) for h in base] # 异源边界:每帧距离恰 18 + + fp_near = _fingerprint(near, 8.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(near)]) + fp_far = _fingerprint(far, 8.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(far)]) + + r_near = _rate(ddp, fp_near, [_video("v-base", base, duration=8.0)]) + r_far = _rate(ddp, fp_far, [_video("v-base", base, duration=8.0)]) + + assert r_near["duplicate_rate"] > 0, "距离16的同源降重对必须检出" + assert r_far["duplicate_rate"] == 0.0, "距离18的异源对不得误报" + assert r_far["match_count"] == 0 + + def test_deduped_pair_frame_match_rate_over_threshold(self): + """真实场景比例:11 帧中 8 帧距离 <=16(0.73 >= 0.7), + 其余 3 帧异源距离(>=18)——frame_match_rate 必须过 0.7 门槛。""" + ddp = VideoDeduplicator() + base = [_h(i + 4) for i in range(11)] + near = [self._variant(h, 14) for h in base[:8]] # 中位数 14 的同源降重帧 + # 异源帧用完全不同前缀(与 base 距离 >=30) + far = [_h(52 + i) for i in range(3)] + query = near + far + + fp = _fingerprint(query, 11.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(query)]) + r = _rate(ddp, fp, [_video("v-base", base, duration=11.0)]) + # frame_match_rate=8/11=0.73、时序片段覆盖 ~0.73 + # → duplicate_rate = 0.4*0.73+0.6*0.73 ≈ 73%(空直方图回退下 fusion=0.6965 + # 略低于 is_duplicate 的 0.70 判定阈值,故此处断言查重率而非 match_count; + # 真实视频带颜色直方图时 fusion≈0.80,staging A-C 实测 is_duplicate=True) + assert r["duplicate_rate"] >= 70.0 + + +class TestTemporalJitterAlignment: + """时序对齐允许目标索引正/反向 ±(neighbor_window+1) 抖动。 + + 密集 1s 采样下相邻帧 pHash 接近,全局最近邻会在目标相邻帧间 + 正负 1 跳变(场景切割/取帧错位/局部倒退);旧逻辑只允许正向 + delta,把同源连续匹配拆碎,min_consecutive 门槛够不上而漏检。 + """ + + def test_backward_jitter_keeps_run_continuous(self): + """匹配目标索引序列 0,1,2,1,2,3(含一次 -1 倒退)应保持同一 run。""" + from video_processing.dedup import find_duplicate_segments + + # 构造 target 相邻帧 pHash 相同(距离0),query 帧的最近邻在 + # target[1]/target[2] 之间抖动;全部 <= 阈值 + t_hash = _h(0) + other = _h(40) + # target: 帧0-3 相同场景,帧4+ 异源 + t_chunks = [_chunk(t_hash, i, i + 1) for i in range(4)] + [_chunk(other, i, i + 1) for i in range(4, 8)] + # query 6 帧同场景(最近邻会落到 target 0~3,索引可正可负) + q_chunks = [_chunk(t_hash, i, i + 1) for i in range(6)] + + segments = find_duplicate_segments(q_chunks, t_chunks) + assert segments, "含 ±1 时序抖动的连续匹配必须形成片段" + # 6 帧匹配 >= min_consecutive(min(5,max(2,6//2))=5),报为一个片段 + assert len(segments) == 1 + seg = segments[0] + assert seg.query_end_ms - seg.query_start_ms >= 5000 + + def test_large_backward_jump_breaks_run(self): + """目标索引倒退 > neighbor_window+1(如从 5 跳回 0)不属于抖动, + 不桥接为同一片段;孤立短匹配 < min_consecutive 不报片段。""" + from video_processing.dedup import find_duplicate_segments + + # 异源段:9-bit 不重叠段(相邻段隔 3 bit),跨段距离 18~24 > 阈值 16 + def _bit_seg(start): + bits = ["0"] * 64 + for b in range(9): + bits[start + b] = "1" + return f"{int(''.join(bits), 2):016x}" + + t_hash = _bit_seg(0) # 复用场景:bit 0-8 + t_other = [_bit_seg(22 + 4 * i) for i in range(4)] # target 异源段 + q_other = [_bit_seg(40 + 4 * i) for i in range(3)] # query 异源段 + # target: 帧0 同场景;帧1-4 异源;帧5-6 同场景 + t_chunks = ( + [_chunk(t_hash, 0, 1)] + + [_chunk(t_other[i - 1], i, i + 1) for i in range(1, 5)] + + [_chunk(t_hash, i, i + 1) for i in range(5, 7)] + ) + # query: 帧0 匹配 target[0];帧1-3 异源(与 target 任何帧距离 >16);帧4-5 匹配 target[5,6] + q_chunks = ( + [_chunk(t_hash, 0, 1)] + + [_chunk(q_other[i - 1], i, i + 1) for i in range(1, 4)] + + [_chunk(t_hash, i, i + 1) for i in range(4, 6)] + ) + segments = find_duplicate_segments(q_chunks, t_chunks) + # 两段各 1、2 帧 < min_consecutive=5 → 不报片段(大跳跃不桥接) + assert segments == [] diff --git a/tests/unit/test_dedup_engine.py b/tests/unit/test_dedup_engine.py index e9ace015d..35ce57d0e 100644 --- a/tests/unit/test_dedup_engine.py +++ b/tests/unit/test_dedup_engine.py @@ -358,11 +358,11 @@ class TestVideoDeduplicatorCheckDuplicate: finally: self._restore_repo(mod, orig) - def test_first_match_returned(self, deduplicator, mock_session): - """返回第一个通过阈值的匹配(非最优匹配)。""" - # vid-1: 距离=2 bits(0x03 XOR 0x01 = 0x02 → 1 bit),通过阈值 + def test_highest_score_match_returned(self, deduplicator, mock_session): + """Issue #1702: 遍历所有候选取融合分最高者(旧逻辑首个过阈即返回)。""" + # vid-1: 距离=1 bit(0x03 XOR 0x01 = 0x02 → 1 bit),通过阈值 vid1 = self._make_existing_video("vid-1", "md5_1", phashes=["0000000000000003"]) - # vid-2: 距离=0 bits(完全匹配) + # vid-2: 距离=0 bits(完全匹配),融合分更高 vid2 = self._make_existing_video("vid-2", "md5_2", phashes=["0000000000000001"]) mock_repo = MagicMock() @@ -380,8 +380,8 @@ class TestVideoDeduplicatorCheckDuplicate: try: result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session) assert result is not None - # 返回第一个通过阈值的匹配(vid-1 距离=1 < 10) - assert result["duplicate_of"] == "vid-1" + # 两个候选都过阈,返回融合分最高的 vid-2(距离 0 < 1) + assert result["duplicate_of"] == "vid-2" finally: self._restore_repo(mod, orig) diff --git a/tests/unit/test_dedup_pure.py b/tests/unit/test_dedup_pure.py index ab70cdb15..a80787daa 100755 --- a/tests/unit/test_dedup_pure.py +++ b/tests/unit/test_dedup_pure.py @@ -185,11 +185,14 @@ class TestBhattacharyyaCoefficient: """_bhattacharyya_coefficient Bhattacharyya 系数测试.""" def test_identical_histograms(self): - """完全相同的直方图系数为1.0.""" - hist = [0.5, 0.5, 0.0, 0.3] + """完全相同的直方图系数为1.0(#1702:按 Σ 归一,概率分布语义)。""" + hist = [0.5, 0.5, 0.0, 0.0] # Σ=1 的概率分布 bc = VideoDeduplicator._bhattacharyya_coefficient(hist, hist) - # Σ √(a[i]*a[i]) = Σ a[i] = 1.0 (normalized) - assert bc == pytest.approx(sum(h for h in hist)) + assert bc == pytest.approx(1.0) + # 非归一化输入也归一到 1.0(三通道拼接 Σ=3 的等价情形) + hist3 = [0.5, 0.5, 0.0, 0.3] + bc3 = VideoDeduplicator._bhattacharyya_coefficient(hist3, hist3) + assert bc3 == pytest.approx(1.0) def test_zero_histograms(self): """全零直方图系数为0.""" @@ -202,10 +205,10 @@ class TestBhattacharyyaCoefficient: assert bc == pytest.approx(0.0) def test_different_lengths(self): - """不同长度直方图取最小长度对齐.""" + """不同长度直方图取最小长度对齐,并按各自总量归一(#1702 概率分布语义)。""" + # 对齐到前 2 维:coeff = 2,norm = √(Σa·Σb) = √(2·2) = 2 → 1.0 bc = VideoDeduplicator._bhattacharyya_coefficient([1.0, 1.0, 0.0, 0.0], [1.0, 1.0]) - # 对齐到前2维: √(1*1) + √(1*1) = 2.0 - assert bc == pytest.approx(2.0) + assert bc == pytest.approx(1.0) def test_known_value(self): """已知值验证.""" diff --git a/tests/unit/test_dedup_v2.py b/tests/unit/test_dedup_v2.py index 8fe0e8540..d0f7c6efe 100644 --- a/tests/unit/test_dedup_v2.py +++ b/tests/unit/test_dedup_v2.py @@ -103,6 +103,7 @@ from video_processing.dedup import ( # noqa: E402 MIN_CONSECUTIVE_MATCHES, MIN_KEYFRAME_INTERVAL_SEC, MIN_KEYFRAMES, + PHASH_THRESHOLD, PHASH_WEIGHT, SCENE_CHANGE_THRESHOLD, SEGMENT_MATCH_THRESHOLD, @@ -269,22 +270,17 @@ class TestFindDuplicateSegments: 注意:使用不同的 hash 对,确保后半部分帧距离 > 阈值。 """ same_hash = "aaaaaaaaaaaaaaaa" - # 4 帧匹配,后面 6 帧各自不同(在 query 和 target 中使用不同 hash) + # 4 帧匹配,后面 6 帧用与匹配哈希距离 32 的不匹配哈希(> PHASH_THRESHOLD=16) + nomatch_hash = "cccccccccccccccc" # hamming(aaaa, cccc)=32 chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [ - _make_chunk(i * 1000, (i + 1) * 1000, "bbbbbbbbbbbbbbbb") for i in range(4, 10) + _make_chunk(i * 1000, (i + 1) * 1000, nomatch_hash) for i in range(4, 10) ] chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [ - _make_chunk(i * 1000, (i + 1) * 1000, "cccccccccccccccc") for i in range(4, 10) + _make_chunk(i * 1000, (i + 1) * 1000, nomatch_hash) for i in range(4, 10) ] - - # hamming("bbbb...", "cccc...") should be > 8 (SEGMENT_MATCH_THRESHOLD) - # b=1011, c=1100 → 4 bits differ per hex digit × 16 digits = 64 bits total? No... - # Actually: hamming_distance("bbbbbbbbbbbbbbbb", "cccccccccccccccc") - # b=0xb=1011, c=0xc=1100 → XOR=0111=0x7 → 3 bits per digit × 16 = 48 - # That's > 8 so won't match - + # hamming(aaaa..., cccc...) = 32 > PHASH_THRESHOLD(16),后半段不匹配; + # 前 4 帧匹配 < min_consecutive=5,不形成片段 segments = find_duplicate_segments(chunks_a, chunks_b) - # 只有 4 帧匹配(< min_consecutive=5),所以不报告 assert segments == [] def test_max_gap_behavior(self): @@ -293,10 +289,12 @@ class TestFindDuplicateSegments: 关键:间隙帧必须在 query 和 target 中使用不同 hash,使其真正不匹配。 """ match_hash = "aaaaaaaaaaaaaaaa" - gap_hash_a = "bbbbbbbbbbbbbbbb" # query 端 - gap_hash_b = "cccccccccccccccc" # target 端(与 query 端距离 > 8) - tail_hash_a = "dddddddddddddddd" - tail_hash_b = "eeeeeeeeeeeeeeee" + # 间隙/尾部哈希与 match_hash 及彼此之间汉明距离均 >64 (> PHASH_THRESHOLD=16), + # 确保在 ±(neighbor_window+1) 时序抖动对齐窗口内也不会误匹配 + gap_hash_a = "ffffffffffffffff" # hamming(a,f)=128 + gap_hash_b = "9999999999999999" # hamming(a,9)=128, hamming(f,9)=128 + tail_hash_a = "7777777777777777" # hamming(a,7)=192 + tail_hash_b = "1111111111111111" # hamming(a,1)=192, hamming(7,1)=128 # 5 帧匹配, 1 帧间隙, 3 帧匹配, 5 帧不匹配 hashes_a = [match_hash] * 5 + [gap_hash_a] + [match_hash] * 3 + [tail_hash_a] * 5 @@ -318,10 +316,10 @@ class TestFindDuplicateSegments: def test_max_gap_exceeded(self): """间隙超过 max_gap → 分成两段.""" match_hash = "aaaaaaaaaaaaaaaa" - gap_hash_a = "bbbbbbbbbbbbbbbb" - gap_hash_b = "cccccccccccccccc" - tail_hash_a = "dddddddddddddddd" - tail_hash_b = "eeeeeeeeeeeeeeee" + gap_hash_a = "ffffffffffffffff" # hamming(a,f)=128 + gap_hash_b = "9999999999999999" # hamming(a,9)=128 + tail_hash_a = "7777777777777777" # hamming(a,7)=192 + tail_hash_b = "1111111111111111" # hamming(a,1)=192 # 5 帧匹配, 3 帧间隙 (> max_gap=2), 5 帧匹配, 5 帧不匹配 hashes_a = [match_hash] * 5 + [gap_hash_a] * 3 + [match_hash] * 5 + [tail_hash_a] * 5 @@ -484,8 +482,10 @@ class TestBackwardCompatibility: chunks_b = [{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": 0, "end_time_ms": 5000}] segments = find_duplicate_segments(chunks_a, chunks_b) - # 1 帧 < min_consecutive=5,不会报重复 - assert segments == [] + # Issue #1702: 自适应门槛 min(5, max(2, 1//2))=2,1 帧不成段; + # N=1 的检出由 _evaluate_candidate 匹配帧回退兜底(见 test_dedup_1702)。 + # 这里只要求不崩溃。 + assert isinstance(segments, list) # ── TestConstants ─────────────────────────────────────────────── @@ -495,8 +495,11 @@ class TestConstants: """常量值验证 — 使用已在模块顶部导入的常量,避免重新 import.""" def test_segment_match_threshold(self): - # 从已导入的 find_duplicate_segments 默认参数间接验证 - assert SEGMENT_MATCH_THRESHOLD == 8 + # Issue #1702 二次校准:阈值经 staging 真实数据两轮回归—— + # 第一轮同源 4/11、异源 min=24 定 12;第二轮扩样本(15 个真实成片) + # 同源降重对中位数距离 14、<=16 命中 8/11=0.73,异源 13 个候选 + # <=16 命中全 0、最近邻最小距离 18 → 校准为 16。 + assert SEGMENT_MATCH_THRESHOLD == PHASH_THRESHOLD == 16 def test_min_consecutive_matches(self): assert MIN_CONSECUTIVE_MATCHES == 5 diff --git a/tests/unit/test_duplicate_rate_scope.py b/tests/unit/test_duplicate_rate_scope.py index 4cf6f02ce..6fd9a0d08 100644 --- a/tests/unit/test_duplicate_rate_scope.py +++ b/tests/unit/test_duplicate_rate_scope.py @@ -132,9 +132,14 @@ class TestCheckDuplicateScopeUser: class TestDurationPrefilter: - """test_duration_prefilter:时长 ±15% 过滤.""" + """Issue #1702: scope=user 跨项目查重不做时长预过滤。 - def test_duration_prefilter_passes_correct_range(self): + 局部片段复用的两个视频时长必然不同(证据视频 20s vs 11s,差 42%), + 旧的 ±15% 窗口会让同源视频互相不可见 → is_duplicate 恒 False。 + 全量遍历同用户视频,异源视频由 fusion/temporal_coverage 阈值天然过滤。 + """ + + def test_user_scope_no_duration_filter(self): from video_processing.dedup import VideoDeduplicator deduplicator = VideoDeduplicator() @@ -153,12 +158,13 @@ class TestDurationPrefilter: duration_sec=30.0, ) - # Should pass duration_min=25.5, duration_max=34.5 (30 ± 15%) + # scope=user 全量遍历:位置参数只传 user_id,kwargs 不含时长过滤 call_args = mock_repo.list_by_user.call_args - assert call_args[1]["duration_min"] == pytest.approx(25.5, abs=0.1) - assert call_args[1]["duration_max"] == pytest.approx(34.5, abs=0.1) + assert call_args[0] == ("user1",) + assert "duration_min" not in call_args[1] + assert "duration_max" not in call_args[1] - def test_no_duration_prefilter_when_zero(self): + def test_user_scope_no_duration_filter_when_zero(self): from video_processing.dedup import VideoDeduplicator deduplicator = VideoDeduplicator() @@ -178,8 +184,34 @@ class TestDurationPrefilter: ) call_args = mock_repo.list_by_user.call_args - assert call_args[1]["duration_min"] == 0 - assert call_args[1]["duration_max"] == 0 + assert "duration_min" not in call_args[1] + assert "duration_max" not in call_args[1] + + def test_project_scope_also_no_duration_filter(self): + """scope=project 走 list_by_project,本来就不做时长过滤。""" + from video_processing.dedup import VideoDeduplicator + + deduplicator = VideoDeduplicator() + fingerprint = _make_fingerprint(duration_ms=30000) + session = MagicMock() + + with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo: + mock_repo = MockRepo.return_value + mock_repo.list_by_project.return_value = [] + deduplicator.check_duplicate( + fingerprint, + "proj1", + session, + scope="project", + user_id="user1", + duration_sec=30.0, + ) + + mock_repo.list_by_project.assert_called_once() + call_args = mock_repo.list_by_project.call_args + assert call_args[0] == ("proj1",) + assert "duration_min" not in call_args[1] + assert "duration_max" not in call_args[1] class TestComputeDuplicateRateFormula: diff --git a/tests/unit/test_enqueue_persists_celery_id_1714.py b/tests/unit/test_enqueue_persists_celery_id_1714.py new file mode 100644 index 000000000..09ecd1949 --- /dev/null +++ b/tests/unit/test_enqueue_persists_celery_id_1714.py @@ -0,0 +1,57 @@ +"""#1714:入队成功后 celery 消息 ID 必须持久化到任务行(供清理时 revoke)。""" + +from __future__ import annotations + +import sys +from pathlib import Path +from unittest.mock import MagicMock + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from app.core import task_enqueue # noqa: E402 + + +class _FakeTask: + def __init__(self): + self.id = "task-enqueue-1" + self.status = "pending" + self.celery_task_id = "" + + def mark_failed(self, msg): # noqa: ARG002 + self.status = "failed" + + +class _FakeRepo: + def __init__(self): + self.updated = None + + def count_pending_total(self): + return 0 + + def count_pending_by_user(self, user_id): # noqa: ARG002 + return 0 + + def update(self, task): + self.updated = task + return task + + +def test_safe_enqueue_persists_celery_message_id(monkeypatch): + fake_result = MagicMock() + fake_result.id = "celery-msg-id-enqueue-999" + mock_celery = MagicMock() + mock_celery.send_task.return_value = fake_result + monkeypatch.setattr(task_enqueue, "celery_app", mock_celery) + + task = _FakeTask() + repo = _FakeRepo() + + ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1") + assert ok is True + # celery_task_id 已持久化 + assert task.celery_task_id == "celery-msg-id-enqueue-999" + assert repo.updated is task + mock_celery.send_task.assert_called_once() + args, kwargs = mock_celery.send_task.call_args + assert args[0] == "worker.generate_video" + assert kwargs.get("args") == [task.id] diff --git a/tests/unit/test_fingerprint_chunks.py b/tests/unit/test_fingerprint_chunks.py index 12635d72b..ccaf649b5 100644 --- a/tests/unit/test_fingerprint_chunks.py +++ b/tests/unit/test_fingerprint_chunks.py @@ -3,7 +3,7 @@ 覆盖: - 分片策略:60秒视频 → 30片,120秒视频 → 24片 - VideoFingerprint.to_chunk_models() 输出正确 -- _save_fingerprint_chunks 幂等性(已有数据跳过) +- _save_fingerprint_chunks 替换语义(Issue #1702:重算时先删旧分片再写入) - to_dict() 向后兼容 """ @@ -169,11 +169,15 @@ class TestVideoFingerprintToChunkModels: assert models == [] -class TestSaveFingerprintChunksIdempotent: - """测试 _save_fingerprint_chunks 幂等性。""" +class TestSaveFingerprintChunksReplace: + """测试 _save_fingerprint_chunks 替换语义(Issue #1702)。 - def test_save_skips_existing(self): - """已有分片数据时跳过写入。""" + 重算查重时指纹算法已升级(中心裁剪 + 新采样/阈值),旧分片必须先删除 + 再写入新分片,否则 recompute-dedup 永远读到旧指纹、修复对存量视频不生效。 + """ + + def test_save_replaces_existing(self): + """已有分片数据时:先删除旧分片,再写入新分片。""" fp = VideoFingerprint( md5="abc", keyframe_phashes=["a1b2"], @@ -186,16 +190,22 @@ class TestSaveFingerprintChunksIdempotent: ) session = MagicMock() - # Mock: 已有 1 条分片数据 - session.query.return_value.filter.return_value.count.return_value = 1 + # Mock: 删除旧分片返回 3(旧算法留下的 3 条分片) + session.query.return_value.filter.return_value.delete.return_value = 3 _save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session) - # bulk_save_objects 不应被调用 - session.bulk_save_objects.assert_not_called() + # 必须先执行删除 + session.query.return_value.filter.return_value.delete.assert_called_once() + # 新分片必须写入 + session.bulk_save_objects.assert_called_once() + saved_models = session.bulk_save_objects.call_args[0][0] + assert len(saved_models) == 1 + assert saved_models[0].video_id == "v1" + assert saved_models[0].phash_binary == "a1b2" def test_save_writes_new(self): - """无分片数据时写入。""" + """无旧分片时直接写入。""" fp = VideoFingerprint( md5="abc", keyframe_phashes=["a1b2"], @@ -208,12 +218,12 @@ class TestSaveFingerprintChunksIdempotent: ) session = MagicMock() - # Mock: 无分片数据 - session.query.return_value.filter.return_value.count.return_value = 0 + # Mock: 无旧分片 + session.query.return_value.filter.return_value.delete.return_value = 0 _save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session) - # bulk_save_objects 应被调用一次 + session.query.return_value.filter.return_value.delete.assert_called_once() session.bulk_save_objects.assert_called_once() saved_models = session.bulk_save_objects.call_args[0][0] assert len(saved_models) == 1 @@ -221,7 +231,7 @@ class TestSaveFingerprintChunksIdempotent: assert saved_models[0].phash_binary == "a1b2" def test_save_skips_no_chunks(self): - """指纹无 chunks 时跳过。""" + """指纹无 chunks 时跳过(不删不写)。""" fp = VideoFingerprint( md5="abc", keyframe_phashes=[], @@ -232,11 +242,11 @@ class TestSaveFingerprintChunksIdempotent: ) session = MagicMock() - session.query.return_value.filter.return_value.count.return_value = 0 _save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session) - # bulk_save_objects 不应被调用 + # 无 chunks:不查询、不删除、不写入 + session.query.assert_not_called() session.bulk_save_objects.assert_not_called() diff --git a/tests/unit/test_generation_preview.py b/tests/unit/test_generation_preview.py index 696e889aa..69a3444bc 100644 --- a/tests/unit/test_generation_preview.py +++ b/tests/unit/test_generation_preview.py @@ -725,8 +725,12 @@ class TestCreatePreviewRoute: generation_task_repository=repo, db=MagicMock(), ) - assert resp.task_id == "preview_task_001" - assert resp.status == "pending" + # 批量响应:N=1 时 items 长度为 1 + assert resp.total == 1 + assert len(resp.items) == 1 + assert resp.items[0].task_id == "preview_task_001" + assert resp.items[0].status == "pending" + assert resp.items[0].variant_index == 0 def test_user_pending_limit_exceeded(self): """用户待处理任务超限 → 429""" @@ -807,7 +811,13 @@ class TestCreatePreviewRoute: repo.count_pending_total.return_value = 0 task = _make_task() - from fastapi import HTTPException + + # 模拟 mark_failed 真实更新任务状态(_mark_task_failed 内部调用) + def _set_failed(error_message="", **_kwargs): + task.status = GenerationTaskStatus.FAILED + task.error_message = error_message + + task.mark_failed.side_effect = _set_failed with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task @@ -815,14 +825,15 @@ class TestCreatePreviewRoute: "app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=False, ): - with pytest.raises(HTTPException) as exc_info: - create_preview_generation_task( - self._make_request(), - authenticated_user=_make_user(), - generation_task_repository=repo, - db=MagicMock(), - ) - assert exc_info.value.status_code == 500 + resp = create_preview_generation_task( + self._make_request(), + authenticated_user=_make_user(), + generation_task_repository=repo, + db=MagicMock(), + ) + # 入队失败:任务被标记 failed(mark_failed 设置错误信息),响应正常返回 + assert resp.total == 1 + assert resp.items[0].status == "failed" def test_enqueue_raises_user_limit(self): """safe_enqueue 抛出 UserPendingLimitExceeded → 429""" @@ -833,6 +844,12 @@ class TestCreatePreviewRoute: task = _make_task() from fastapi import HTTPException + def _set_failed_limit(error_message="", **_kwargs): + task.status = GenerationTaskStatus.FAILED + task.error_message = error_message or "待处理任务超限" + + task.mark_failed.side_effect = _set_failed_limit + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task with patch( @@ -846,6 +863,7 @@ class TestCreatePreviewRoute: generation_task_repository=repo, db=MagicMock(), ) + # 全部变体入队失败且错误消息含"待处理任务" → 429 assert exc_info.value.status_code == 429 def test_enqueue_raises_global_queue_full(self): @@ -857,6 +875,12 @@ class TestCreatePreviewRoute: task = _make_task() from fastapi import HTTPException + def _set_failed_queue(error_message="", **_kwargs): + task.status = GenerationTaskStatus.FAILED + task.error_message = error_message or "系统队列已满" + + task.mark_failed.side_effect = _set_failed_queue + with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task with patch( @@ -870,6 +894,7 @@ class TestCreatePreviewRoute: generation_task_repository=repo, db=MagicMock(), ) + # 全部变体入队失败且错误消息含"队列" → 503 assert exc_info.value.status_code == 503 diff --git a/tests/unit/test_ingest_hevc_orphan_1714.py b/tests/unit/test_ingest_hevc_orphan_1714.py new file mode 100644 index 000000000..3701ea381 --- /dev/null +++ b/tests/unit/test_ingest_hevc_orphan_1714.py @@ -0,0 +1,317 @@ +"""Issue #1714:HEVC 转码后禁止兜底新建重复 READY 记录,必须回写占位 asset。 + +覆盖: +- 转码成功 + 占位 asset 存在(按原始 key 找到)→ 更新占位为 READY、 + storage_key 改写为 *_h264,绝不 create 新记录(回归 P1 孤儿 PROCESSING bug) +- job.asset_id 透传时优先按 id 关联占位(即使 key 对不上也能命中) +- 无占位记录(旧链路)→ 兜底新建(保留兼容) +- 非 HEVC:占位同样被更新为 READY,不新建 +- 无效媒体:占位标记为 ERROR,不新建 ERROR 记录 +- ingest 异常:占位(按还原后的原始 key)标记 ERROR +""" + +from __future__ import annotations + +import sys +import tempfile +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +# ── 与 test_ingest_hevc_transcode_task.py 相同的 worker 模块加载方式 ── +_SAVED_MODULES_KEYS = set(sys.modules.keys()) + +_mock_db_module = MagicMock() +_mock_db_module.SessionLocal = MagicMock() +sys.modules["worker_app.db"] = _mock_db_module +sys.modules["worker_app.core.config"] = MagicMock() + +_mock_celery_module = MagicMock() + + +def _passthrough_decorator(*args, **kwargs): + if len(args) == 1 and callable(args[0]): + return args[0] + return lambda f: f + + +_mock_celery_module.celery_app.task = MagicMock(side_effect=_passthrough_decorator) +sys.modules["worker_app.celery_app"] = _mock_celery_module + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker")) + +import pytest # noqa: E402 +from worker_app.tasks import ingest as ingest_mod # noqa: E402 + +from packages.domain import Asset, AssetStatus # noqa: E402 + +for _key in list(sys.modules.keys()): + if _key not in _SAVED_MODULES_KEYS and not _key.startswith("video_processing"): + del sys.modules[_key] +del _SAVED_MODULES_KEYS + + +# ── 假仓储 ────────────────────────────────────────────────────────────── +class _FakeJobRepo: + def __init__(self, job): + self.job = job + self.updated = None + + def get(self, job_id): + return self.job + + def update(self, job): + self.updated = job + return job + + +class _FakeAssetRepo: + """记录 create 调用;find_* 按内部 assets 列表查询。""" + + def __init__(self, assets: list[Asset] | None = None): + self.assets = list(assets or []) + self.created: list[Asset] = [] + self.updated: list[Asset] = [] + + def create(self, asset: Asset) -> Asset: + self.created.append(asset) + self.assets.append(asset) + return asset + + def update(self, asset: Asset) -> Asset: + self.updated.append(asset) + return asset + + def find_by_id(self, asset_id: str) -> Asset | None: + return next((a for a in self.assets if a.id == asset_id), None) + + def find_by_storage_key(self, storage_key: str) -> Asset | None: + return next((a for a in self.assets if a.storage_key == storage_key), None) + + +def _make_job(asset_id: str = "", storage_key: str = "uploads/proj/IMG_2282.MOV"): + return SimpleNamespace( + id="job-1", + project_id="proj-1", + library_id="lib-1", + storage_key=storage_key, + file_hash="hash-1", + asset_id=asset_id, + status=None, + error_message=None, + result_asset_id=None, + updated_at=None, + ) + + +def _make_placeholder(storage_key: str = "uploads/proj/IMG_2282.MOV", asset_id: str = "asset-ph"): + return Asset( + id=asset_id, + project_id="proj-1", + library_id="lib-1", + name="IMG_2282.MOV", + storage_key=storage_key, + mime_type="video/quicktime", + status=AssetStatus.PROCESSING, + file_hash="hash-1", + ) + + +def _video_metadata(codec="hevc"): + return { + "codec": codec, + "width": 1920, + "height": 1080, + "duration": 10.0, + "size_bytes": 5 * 1024 * 1024, + } + + +@pytest.fixture +def transcode_env(tmp_path): + """HEVC 转码成功的标准 mock 环境(同 test_ingest_hevc_transcode_task)。""" + local_file = tmp_path / "local_hevc.MOV" + local_file.write_bytes(b"fake-hevc-source") + tc_out = tmp_path / "transcode_out_h264.mp4" + + control = { + "validate_ok": True, + "tc_out": tc_out, + "local_file": local_file, + "download_ok": True, + "extract_success": True, + "codec": "hevc", + "raise_in_flow": None, + } + + def fake_ntf(*args, **kwargs): + mock_file = MagicMock() + mock_file.name = str(tc_out) if kwargs.get("suffix") == "_h264.mp4" else str(local_file) + mock_file.close = MagicMock() + mock_file.__enter__.return_value = mock_file + mock_file.__exit__.return_value = False + return mock_file + + def fake_subprocess_run(cmd, **kwargs): + if cmd and cmd[0] == "ffmpeg" and "libx264" in cmd: + Path(cmd[-1]).write_bytes(b"fake-h264-output") + return SimpleNamespace(returncode=0, stderr="") + return SimpleNamespace(returncode=0, stdout="", stderr="") + + control["patchers"] = { + "session": patch.object(ingest_mod, "SessionLocal", return_value=MagicMock()), + "download": patch.object(ingest_mod, "download_asset", side_effect=lambda *a, **kw: control["download_ok"]), + "upload": patch("video_processing.oss_helpers.upload_to_oss", return_value="https://oss/x"), + "metadata": patch.object( + ingest_mod, + "extract_media_metadata", + side_effect=lambda path, mt: ( + (_video_metadata("h264"), control["extract_success"]) + if Path(path).name == tc_out.name + else (_video_metadata(control["codec"]), control["extract_success"]) + ), + ), + "validate": patch.object( + ingest_mod, "validate_transcode_output", side_effect=lambda p, portrait: control["validate_ok"] + ), + "subprocess": patch.object(ingest_mod.subprocess, "run", side_effect=fake_subprocess_run), + "ntf": patch.object(tempfile, "NamedTemporaryFile", side_effect=fake_ntf), + "thumb": patch( + "video_processing.thumbnail_generator.extract_first_frame", + side_effect=RuntimeError("skip thumb"), + ), + } + return control + + +def _start(control, job, assets): + job_repo = _FakeJobRepo(job) + asset_repo = _FakeAssetRepo(assets) + patchers = dict(control["patchers"]) + patchers["job_repo"] = patch.object(ingest_mod, "SQLAlchemyIngestJobRepository", return_value=job_repo) + patchers["asset_repo"] = patch.object(ingest_mod, "SQLAlchemyAssetRepository", return_value=asset_repo) + started = {name: p.start() for name, p in patchers.items()} + return started, job_repo, asset_repo + + +def _stop(control): + for p in control["patchers"].values(): + p.stop() + + +class TestHEVCTranscodePlaceholderRewrite: + def test_transcode_success_updates_placeholder_no_duplicate_ready(self, transcode_env): + """转码成功 → 占位 asset 原地更新为 READY + storage_key 改写 _h264,禁止新建。""" + control = transcode_env + placeholder = _make_placeholder() + job = _make_job() # 旧 job 无 asset_id,靠原始 key 关联 + mocks, job_repo, asset_repo = _start(control, job, [placeholder]) + try: + result = ingest_mod.ingest_asset("job-1") + finally: + _stop(control) + + assert result["status"] == "completed" + # 核心断言 1:没有新建任何 READY 记录(旧 bug 会 create 一条 _h264 READY) + assert asset_repo.created == [], "转码回写不得新建 asset 记录" + # 核心断言 2:占位被更新为 READY,且 storage_key 已是 _h264 + assert len(asset_repo.updated) == 1 + updated = asset_repo.updated[0] + assert updated.id == placeholder.id + assert updated.status == AssetStatus.READY + assert updated.storage_key == "uploads/proj/IMG_2282_h264.MOV" + assert updated.metadata.get("hevc_transcoded") is True + assert updated.metadata.get("original_storage_key") == "uploads/proj/IMG_2282.MOV" + # job 关联到同一条 asset + assert job_repo.updated.result_asset_id == placeholder.id + assert job_repo.updated.storage_key == "uploads/proj/IMG_2282_h264.MOV" + + def test_placeholder_resolved_by_job_asset_id(self, transcode_env): + """job.asset_id 透传时优先按 id 关联(即使 storage_key 对不上也命中)。""" + control = transcode_env + placeholder = _make_placeholder(storage_key="uploads/different/key.MOV", asset_id="asset-by-id") + job = _make_job(asset_id="asset-by-id") + _, _, asset_repo = _start(control, job, [placeholder]) + try: + result = ingest_mod.ingest_asset("job-1") + finally: + _stop(control) + + assert result["status"] == "completed" + assert asset_repo.created == [] + assert len(asset_repo.updated) == 1 + assert asset_repo.updated[0].id == "asset-by-id" + assert asset_repo.updated[0].status == AssetStatus.READY + + def test_no_placeholder_fallback_creates_ready(self, transcode_env): + """旧链路无占位记录 → 兜底新建 READY(兼容保留,但必须是唯一一条)。""" + control = transcode_env + job = _make_job(asset_id="") + _, _, asset_repo = _start(control, job, []) + try: + result = ingest_mod.ingest_asset("job-1") + finally: + _stop(control) + + assert result["status"] == "completed" + assert len(asset_repo.created) == 1 + created = asset_repo.created[0] + assert created.status == AssetStatus.READY + assert created.storage_key == "uploads/proj/IMG_2282_h264.MOV" + assert asset_repo.updated == [] + + def test_non_hevc_placeholder_updated_no_create(self, transcode_env): + """非 HEVC(h264)不转码:占位按原始 key 找到并更新 READY,不新建。""" + control = transcode_env + control["codec"] = "h264" + placeholder = _make_placeholder() + job = _make_job() + mocks, _, asset_repo = _start(control, job, [placeholder]) + try: + result = ingest_mod.ingest_asset("job-1") + finally: + _stop(control) + + assert result["status"] == "completed" + assert asset_repo.created == [] + assert len(asset_repo.updated) == 1 + updated = asset_repo.updated[0] + assert updated.status == AssetStatus.READY + assert updated.storage_key == "uploads/proj/IMG_2282.MOV" # 未转码,key 不变 + mocks["upload"].assert_not_called() + + def test_invalid_media_marks_placeholder_error_no_create(self, transcode_env): + """无效媒体:占位标记 ERROR 并 update,禁止再 create 一条 ERROR。""" + control = transcode_env + control["download_ok"] = False # 下载失败 → extract_success=False → 无效媒体路径 + placeholder = _make_placeholder() + job = _make_job() + _, job_repo, asset_repo = _start(control, job, [placeholder]) + try: + result = ingest_mod.ingest_asset("job-1") + finally: + _stop(control) + + assert result["status"] == "failed" + assert asset_repo.created == [], "无效媒体不得新建 ERROR 记录" + assert len(asset_repo.updated) == 1 + assert asset_repo.updated[0].id == placeholder.id + assert asset_repo.updated[0].status == AssetStatus.ERROR + assert job_repo.updated.result_asset_id == placeholder.id + + def test_exception_path_marks_placeholder_error(self, transcode_env): + """ingest 主流程抛异常(如元数据提取炸了)→ 占位按原始 key 找到并标 ERROR。""" + control = transcode_env + placeholder = _make_placeholder() + job = _make_job() + started, _, asset_repo = _start(control, job, [placeholder]) + started["metadata"].side_effect = RuntimeError("boom in flow") + try: + result = ingest_mod.ingest_asset("job-1") + finally: + _stop(control) + + assert result["status"] == "failed" + # 异常路径把占位标 ERROR(旧实现用被改写的 _h264 key 回查会落空) + error_marked = [a for a in asset_repo.assets if a.id == placeholder.id and a.status == AssetStatus.ERROR] + assert error_marked, "异常路径必须把占位 asset 标为 ERROR" diff --git a/tests/unit/test_ingest_orphan_cleanup_1714.py b/tests/unit/test_ingest_orphan_cleanup_1714.py new file mode 100644 index 000000000..cceccf8d4 --- /dev/null +++ b/tests/unit/test_ingest_orphan_cleanup_1714.py @@ -0,0 +1,262 @@ +"""#1714 上传/转码链路(IngestJob + Asset)孤儿清理测试。 + +场景:worker 容器重启/进程 OOM 时,已 prefetch 的 transcode celery 消息丢失, +ingest_job 永久卡 pending/processing、asset 永久卡 processing/uploading。 +""" + +from __future__ import annotations + +import os +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path + +import pytest + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret") +os.environ.setdefault("DATABASE_URL", "sqlite:///test_ingest_orphan.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker")) +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) +sys.path.insert(0, str(Path(__file__).resolve().parents[2])) + +from sqlalchemy import create_engine # noqa: E402 +from sqlalchemy.orm import sessionmaker # noqa: E402 + +from packages.adapters.sqlalchemy_impl.models import AssetModel, Base, IngestJobModel # noqa: E402 +from packages.application.ingest_orphan_cleanup import ( # noqa: E402 + cleanup_orphan_processing_assets, + cleanup_stale_ingest_jobs, +) + + +@pytest.fixture() +def session(): + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(bind=engine) + Session = sessionmaker(bind=engine) + db = Session() + yield db + db.close() + + +def _mk_job(session, *, status="processing", celery_task_id="cel-1", asset_id="a-1", minutes_ago=90): + now = datetime.now(timezone.utc) + job = IngestJobModel( + id=f"job-{minutes_ago}-{status}-{celery_task_id}", + project_id="p-1", + library_id="lib-1", + storage_key="uploads/x/IMG_2285.MOV", + status=status, + asset_id=asset_id, + celery_task_id=celery_task_id, + created_at=now - timedelta(minutes=minutes_ago), + updated_at=now - timedelta(minutes=minutes_ago), + ) + session.add(job) + session.commit() + return job + + +def _mk_asset(session, *, id="a-1", status="processing", minutes_ago=90, file_size=0): + now = datetime.now(timezone.utc) + asset = AssetModel( + id=id, + project_id="p-1", + asset_library_id="lib-1", + name="IMG_2285.MOV", + file_type="video", + file_size=file_size, + file_url="https://example.com/x.mov", + storage_key="uploads/x/IMG_2285.MOV", + status=status, + uploaded_by_user_id="u-1", + created_at=now - timedelta(minutes=minutes_ago), + updated_at=now - timedelta(minutes=minutes_ago), + ) + session.add(asset) + session.commit() + return asset + + +class TestCleanupStaleIngestJobs: + def test_stale_processing_job_marked_failed_and_asset_to_error(self, session): + """processing 超 60 分钟 → job failed,关联 processing asset → error。""" + _mk_asset(session, id="a-1", status="processing") + _mk_job(session, status="processing", celery_task_id="cel-dead", asset_id="a-1", minutes_ago=90) + + items, asset_ids = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90) + + assert len(items) == 1 + assert items[0] == ("job-90-processing-cel-dead", "cel-dead") + assert asset_ids == ["a-1"] + db_job = session.query(IngestJobModel).one() + assert db_job.status == "failed" + assert "中断" in db_job.error_message + db_asset = session.query(AssetModel).one() + assert db_asset.status == "error" + + def test_stale_pending_job_marked_failed(self, session): + """pending 超 90 分钟(从未被消费)→ job failed。""" + _mk_asset(session, id="a-2", status="uploading") + _mk_job(session, status="pending", celery_task_id="", asset_id="a-2", minutes_ago=120) + + items, asset_ids = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90) + + assert len(items) == 1 + assert items[0][1] == "" # 无 celery task id + assert session.query(IngestJobModel).one().status == "failed" + assert session.query(AssetModel).one().status == "error" + + def test_recent_processing_job_not_touched(self, session): + """processing 仅 10 分钟(正常转码中)→ 不误杀。""" + _mk_asset(session, id="a-3", status="processing", minutes_ago=10) + _mk_job(session, status="processing", celery_task_id="cel-live", asset_id="a-3", minutes_ago=10) + + items, asset_ids = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90) + + assert items == [] + assert asset_ids == [] + assert session.query(IngestJobModel).one().status == "processing" + assert session.query(AssetModel).one().status == "processing" + + def test_recent_pending_job_not_touched(self, session): + """pending 仅 30 分钟(队列积压排队中)→ 不误杀。""" + _mk_job(session, status="pending", asset_id="", minutes_ago=30) + + items, _ = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90) + + assert items == [] + assert session.query(IngestJobModel).one().status == "pending" + + def test_terminal_job_not_touched(self, session): + """已 completed/failed 的 job 不动。""" + _mk_job(session, status="completed", celery_task_id="", asset_id="", minutes_ago=999) + _mk_job(session, status="failed", celery_task_id="", asset_id="", minutes_ago=999) + + items, _ = cleanup_stale_ingest_jobs(session) + + assert items == [] + statuses = sorted(j.status for j in session.query(IngestJobModel).all()) + assert statuses == ["completed", "failed"] + + def test_ready_asset_not_demoted(self, session): + """关联 asset 已是 ready(转码其实成功了,仅 job 回写失败)→ 不降级为 error。""" + _mk_asset(session, id="a-4", status="ready") + _mk_job(session, status="processing", celery_task_id="cel-x", asset_id="a-4", minutes_ago=90) + + _, asset_ids = cleanup_stale_ingest_jobs(session) + + assert asset_ids == [] # ready 不动 + assert session.query(AssetModel).one().status == "ready" + + +class TestCleanupOrphanProcessingAssets: + def test_orphan_asset_without_job_marked_error(self, session): + """无 ingest_job 关联、created 超 120 分钟的 processing 占位 → error。""" + _mk_asset(session, id="orphan-1", status="processing", minutes_ago=150) + + ids = cleanup_orphan_processing_assets(session, timeout_minutes=120) + + assert ids == ["orphan-1"] + assert session.query(AssetModel).one().status == "error" + + def test_asset_with_active_job_not_touched(self, session): + """有 processing job 关联的 asset 不由本函数处理(归 cleanup_stale_ingest_jobs)。""" + _mk_asset(session, id="a-5", status="processing", minutes_ago=150) + _mk_job(session, status="processing", asset_id="a-5", minutes_ago=150) + + ids = cleanup_orphan_processing_assets(session, timeout_minutes=120) + + assert ids == [] + assert session.query(AssetModel).one().status == "processing" + + def test_recent_orphan_asset_not_touched(self, session): + """无 job 但才创建 30 分钟 → 可能 complete 刚建、job 派单中,不动。""" + _mk_asset(session, id="orphan-2", status="processing", minutes_ago=30) + + ids = cleanup_orphan_processing_assets(session, timeout_minutes=120) + + assert ids == [] + assert session.query(AssetModel).one().status == "processing" + + +class TestRecoverStuckIngestJobsOnStartup: + def test_stuck_processing_job_requeued(self, session): + """processing 超 10 分钟 → 重置 pending 并重新 send_task,回写新 celery id。""" + from types import SimpleNamespace + + from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup + + job = _mk_job(session, status="processing", celery_task_id="old-cel-1", asset_id="a-1", minutes_ago=30) + + sent = [] + + def fake_send_task(name, args=None, **kw): + sent.append((name, args)) + return SimpleNamespace(id="new-cel-9") + + updated_ids = [] + recovered = recover_stuck_ingest_jobs_on_startup( + session, + send_task=fake_send_task, + update_celery_task_id=lambda jid, cid: updated_ids.append((jid, cid)), + stuck_minutes=10, + ) + + assert recovered == 1 + assert sent == [("worker.ingest_asset", [job.id])] + refreshed = session.query(IngestJobModel).filter_by(id=job.id).one() + assert refreshed.status == "pending" + assert refreshed.celery_task_id == "new-cel-9" + assert updated_ids == [(job.id, "new-cel-9")] + + def test_recent_processing_job_not_touched(self, session): + """processing 仅 5 分钟(正常转码中/部署交接窗口)→ 不抢。""" + from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup + + _mk_job(session, status="processing", celery_task_id="live", asset_id="", minutes_ago=5) + + sent = [] + recovered = recover_stuck_ingest_jobs_on_startup( + session, + send_task=lambda *a, **k: sent.append(a), + stuck_minutes=10, + ) + + assert recovered == 0 + assert sent == [] + assert session.query(IngestJobModel).one().status == "processing" + + def test_lock_not_acquired_skips(self, session): + """未抢到分布式锁(另一 worker 正在恢复)→ 跳过。""" + from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup + + _mk_job(session, status="processing", celery_task_id="x", asset_id="", minutes_ago=30) + + recovered = recover_stuck_ingest_jobs_on_startup( + session, + send_task=lambda *a, **k: None, + lock_acquire=lambda: False, + stuck_minutes=10, + ) + + assert recovered == 0 + assert session.query(IngestJobModel).one().status == "processing" + + def test_pending_and_terminal_not_requeued(self, session): + """pending/已终态 job 不在恢复范围。""" + from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup + + _mk_job(session, status="pending", celery_task_id="", asset_id="", minutes_ago=60) + _mk_job(session, status="failed", celery_task_id="", asset_id="", minutes_ago=60) + + recovered = recover_stuck_ingest_jobs_on_startup( + session, + send_task=lambda *a, **k: None, + stuck_minutes=10, + ) + + assert recovered == 0 + statuses = sorted(j.status for j in session.query(IngestJobModel).all()) + assert statuses == ["failed", "pending"] diff --git a/tests/unit/test_orphan_guard_purge_mocked_1714.py b/tests/unit/test_orphan_guard_purge_mocked_1714.py new file mode 100644 index 000000000..29122959d --- /dev/null +++ b/tests/unit/test_orphan_guard_purge_mocked_1714.py @@ -0,0 +1,360 @@ +"""#1714:孤儿消息撤销/清理逻辑测试(mock redis,CI 无真实 redis 时也产生覆盖)。 + +覆盖 packages/shared/celery_orphan_guard.py: +- _extract_business_ids:三元组 body / 裸 args body / dict args / headers 提取 / + 无 body / 坏 JSON / 坏 base64 / 空 args +- _purge_one_queue:biz id 命中、celery id 命中、未命中保序(重写 rpush)、 + lrange 异常、重写异常、空队列 +- purge_stale_messages_from_queues:空 ids 早退、redis 未安装、连接失败、 + 正常清理并 close +- revoke_and_purge:revoke 逐消息调用、revoke 异常不阻断、空 id 跳过 +- ensure_task_claimable:任务不存在返回空串、终态抛错、pending 放行 +""" + +from __future__ import annotations + +import base64 +import json +import sys +import types +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from packages.shared import celery_orphan_guard as guard # noqa: E402 + + +def _envelope(celery_id: str | None, body_payload) -> bytes: + """构造 Redis transport 存储的 celery 消息(JSON 信封)。""" + if body_payload is None: + body = None + else: + body = base64.b64encode(json.dumps(body_payload).encode()).decode() + envelope = {"body": body, "headers": {"id": celery_id, "task": "worker.generate_video"}} + return json.dumps(envelope).encode() + + +# ── _extract_business_ids ─────────────────────────────────────────────── + + +def test_extract_ids_standard_tuple_body(): + raw = _envelope("celery-1", [["biz-task-1"], {}, {"callbacks": None}]) + assert guard._extract_business_ids(raw) == ("celery-1", "biz-task-1") + + +def test_extract_ids_bare_args_body(): + raw = _envelope("celery-2", ["biz-task-2"]) + assert guard._extract_business_ids(raw) == ("celery-2", "biz-task-2") + + +def test_extract_ids_dict_body_with_args(): + raw = _envelope("celery-3", {"args": ["biz-task-3"], "kwargs": {}}) + assert guard._extract_business_ids(raw) == ("celery-3", "biz-task-3") + + +def test_extract_ids_non_dict_headers_returns_celery_id_none(): + raw = json.dumps({"body": base64.b64encode(json.dumps([["biz-4"]]).encode()).decode(), "headers": "x"}).encode() + celery_id, biz_id = guard._extract_business_ids(raw) + assert celery_id is None + assert biz_id == "biz-4" + + +def test_extract_ids_no_body_returns_celery_id_only(): + raw = json.dumps({"headers": {"id": "celery-5"}}).encode() + assert guard._extract_business_ids(raw) == ("celery-5", None) + + +def test_extract_ids_empty_args_returns_no_biz_id(): + raw = _envelope("celery-6", [[], {}, {}]) + assert guard._extract_business_ids(raw) == ("celery-6", None) + + +def test_extract_ids_args_first_none_returns_no_biz_id(): + raw = _envelope("celery-7", [[None], {}, {}]) + assert guard._extract_business_ids(raw) == ("celery-7", None) + + +def test_extract_ids_bad_json_returns_none_none(): + assert guard._extract_business_ids(b"not-json{") == (None, None) + + +def test_extract_ids_bad_base64_returns_none_none(): + raw = json.dumps({"body": "!!!not-base64!!!", "headers": {"id": "c"}}).encode() + assert guard._extract_business_ids(raw) == (None, None) + + +def test_extract_ids_int_arg_coerced_to_str(): + raw = _envelope("celery-9", [[12345], {}, {}]) + celery_id, biz_id = guard._extract_business_ids(raw) + assert celery_id == "celery-9" + assert biz_id == "12345" + + +# ── _purge_one_queue ──────────────────────────────────────────────────── + + +def _queue_with_messages(*payloads: bytes): + """返回 list-backed mock redis client(记录当前队列内容)。""" + client = MagicMock() + store: dict[str, list[bytes]] = {"q": list(payloads)} + + def lrange(name, start, end): # noqa: ARG001 + return list(store.get(name, [])) + + client.lrange.side_effect = lrange + + pipe = MagicMock() + pipe.delete.side_effect = lambda name: store.pop(name, None) + pipe.rpush.side_effect = lambda name, *items: store.setdefault(name, []).extend(items) + client.pipeline.return_value = pipe + return client, store, pipe + + +def test_purge_one_queue_removes_by_biz_id_and_keeps_order(): + stale = _envelope("c-stale", [["biz-stale"], {}, {}]) + keep1 = _envelope("c-keep-1", [["biz-keep-1"], {}, {}]) + keep2 = _envelope("c-keep-2", [["biz-keep-2"], {}, {}]) + client, store, pipe = _queue_with_messages(keep1, stale, keep2) + + removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set()) + assert removed == 1 + # 队列被 delete + rpush 重写,未命中消息保持相对顺序 + pipe.delete.assert_called_once_with("q") + pipe.rpush.assert_called_once() + args, _ = pipe.rpush.call_args + assert args[0] == "q" + assert list(args[1:]) == [keep1, keep2] + pipe.execute.assert_called_once() + + +def test_purge_one_queue_removes_by_celery_message_id(): + stale = _envelope("celery-xyz", [["biz-whatever"], {}, {}]) + keep = _envelope("celery-aaa", [["biz-keep"], {}, {}]) + client, store, pipe = _queue_with_messages(stale, keep) + + removed = guard._purge_one_queue(client, "q", set(), {"celery-xyz"}) + assert removed == 1 + args, _ = pipe.rpush.call_args + assert list(args[1:]) == [keep] + + +def test_purge_one_queue_no_hit_no_rewrite(): + msg1 = _envelope("c1", [["b1"], {}, {}]) + msg2 = _envelope("c2", [["b2"], {}, {}]) + client, store, pipe = _queue_with_messages(msg1, msg2) + + removed = guard._purge_one_queue(client, "q", {"other"}, {"other-c"}) + assert removed == 0 + # 没有命中:不重写队列 + pipe.delete.assert_not_called() + pipe.rpush.assert_not_called() + + +def test_purge_one_queue_all_removed_deletes_without_rpush(): + stale = _envelope("c-stale", [["biz-stale"], {}, {}]) + client, store, pipe = _queue_with_messages(stale) + + removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set()) + assert removed == 1 + pipe.delete.assert_called_once_with("q") + pipe.rpush.assert_not_called() + + +def test_purge_one_queue_lrange_exception_returns_zero(): + client = MagicMock() + client.lrange.side_effect = RuntimeError("redis down") + assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0 + + +def test_purge_one_queue_empty_queue_returns_zero(): + client = MagicMock() + client.lrange.return_value = [] + assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0 + client.pipeline.assert_not_called() + + +def test_purge_one_queue_rewrite_exception_returns_zero(): + stale = _envelope("c-stale", [["biz-stale"], {}, {}]) + client, store, pipe = _queue_with_messages(stale) + pipe.execute.side_effect = RuntimeError("write fail") + + removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set()) + assert removed == 0 + + +def test_purge_one_queue_unparseable_message_conservatively_kept(): + stale = _envelope("c-stale", [["biz-stale"], {}, {}]) + garbage = b"garbage-not-a-message" + client, store, pipe = _queue_with_messages(garbage, stale) + + removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set()) + assert removed == 1 + args, _ = pipe.rpush.call_args + # 无法解析的消息保守保留,绝不误删 + assert list(args[1:]) == [garbage] + + +# ── purge_stale_messages_from_queues ──────────────────────────────────── + + +def test_purge_queues_no_ids_returns_zero_without_connecting(): + assert guard.purge_stale_messages_from_queues("redis://x", ("q",)) == 0 + + +def test_purge_queues_blank_ids_filtered_out(): + assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["", None]) == 0 + + +def test_purge_queues_redis_not_installed(monkeypatch): + """redis-py 不可用(ImportError)时安全返回 0。""" + import builtins + + real_import = builtins.__import__ + + def fake_import(name, *args, **kwargs): + if name == "redis": + raise ImportError("no redis") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", fake_import) + assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["b1"]) == 0 + + +def test_purge_queues_connection_failure_returns_zero(): + fake_redis = types.ModuleType("redis") + + class _FakeRedis: + @classmethod + def from_url(cls, url): # noqa: ARG003 + client = MagicMock() + client.ping.side_effect = ConnectionError("connect refused") + return client + + fake_redis.Redis = _FakeRedis + sys.modules["redis"] = fake_redis + try: + assert guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"]) == 0 + finally: + sys.modules.pop("redis", None) + + +def test_purge_queues_happy_path_closes_client(): + stale = _envelope("c-stale", [["biz-stale"], {}, {}]) + fake_redis = types.ModuleType("redis") + + client = MagicMock() + client.lrange.return_value = [stale] + pipe = MagicMock() + client.pipeline.return_value = pipe + + class _FakeRedis: + @classmethod + def from_url(cls, url): # noqa: ARG003 + return client + + fake_redis.Redis = _FakeRedis + sys.modules["redis"] = fake_redis + try: + removed = guard.purge_stale_messages_from_queues( + "redis://x", ("generation", "transcode"), business_task_ids=["biz-stale"] + ) + finally: + sys.modules.pop("redis", None) + + # mock client 对两个队列都返回同一条作废消息 → 各移除 1 条 + assert removed == 2 + client.ping.assert_called_once() + client.close.assert_called_once() + # 两个队列都扫描 + assert client.lrange.call_count == 2 + + +def test_purge_queues_close_exception_swallowed(): + fake_redis = types.ModuleType("redis") + + client = MagicMock() + client.lrange.return_value = [] + client.close.side_effect = RuntimeError("close fail") + + class _FakeRedis: + @classmethod + def from_url(cls, url): # noqa: ARG003 + return client + + fake_redis.Redis = _FakeRedis + sys.modules["redis"] = fake_redis + try: + removed = guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"]) + finally: + sys.modules.pop("redis", None) + assert removed == 0 + + +# ── revoke_and_purge ──────────────────────────────────────────────────── + + +def test_revoke_and_purge_revokes_each_message(monkeypatch): + fake_app = MagicMock() + purge_mock = MagicMock(return_value=2) + monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock) + + removed = guard.revoke_and_purge( + fake_app, + "redis://x", + business_task_ids=["b1"], + celery_task_ids=["c1", "c2"], + queue_names=("generation",), + ) + assert removed == 2 + assert fake_app.control.revoke.call_count == 2 + fake_app.control.revoke.assert_any_call("c1") + fake_app.control.revoke.assert_any_call("c2") + purge_mock.assert_called_once_with( + "redis://x", ("generation",), business_task_ids=["b1"], celery_task_ids=["c1", "c2"] + ) + + +def test_revoke_and_purge_revoke_exception_does_not_block(monkeypatch): + fake_app = MagicMock() + fake_app.control.revoke.side_effect = RuntimeError("broadcast fail") + purge_mock = MagicMock(return_value=0) + monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock) + + removed = guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["c1"]) + assert removed == 0 + purge_mock.assert_called_once() + + +def test_revoke_and_purge_skips_blank_ids(monkeypatch): + fake_app = MagicMock() + purge_mock = MagicMock(return_value=0) + monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock) + + guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["", None]) + fake_app.control.revoke.assert_not_called() + + +# ── ensure_task_claimable ─────────────────────────────────────────────── + + +def test_ensure_claimable_missing_task_returns_empty(): + assert guard.ensure_task_claimable("t1", lambda _tid: None) == "" + + +def test_ensure_claimable_terminal_raises(): + with pytest.raises(guard.StaleTaskDiscarded) as exc_info: + guard.ensure_task_claimable("t1", lambda _tid: "failed") + assert exc_info.value.task_id == "t1" + assert exc_info.value.status == "failed" + + +def test_ensure_claimable_cancelled_raises(): + with pytest.raises(guard.StaleTaskDiscarded): + guard.ensure_task_claimable("t1", lambda _tid: "cancelled") + + +def test_ensure_claimable_pending_passes(): + assert guard.ensure_task_claimable("t1", lambda _tid: "pending") == "pending" diff --git a/tests/unit/test_patch_me_profile_1718.py b/tests/unit/test_patch_me_profile_1718.py new file mode 100644 index 000000000..d2444f77b --- /dev/null +++ b/tests/unit/test_patch_me_profile_1718.py @@ -0,0 +1,222 @@ +"""#1718:PATCH /auth/me 资料更新接口测试。 + +覆盖: +- 正常更新昵称并落库 +- strip 生效(前后空白去除) +- 纯空白/超长 -> 422(pydantic 校验) +- 首次设置昵称 profile_completed False->True +- 已完成用户重复提交幂等(仍 True) +- 未登录由 get_current_user 依赖保证 401(框架行为,这里验证路由声明了该依赖) +- 响应结构 {user: {...}} 含 wechat_bound/profile_completed 全字段 +- 微信新建用户 profile_completed 默认 False(wechat_sync _create_wechat_user) +""" + +from __future__ import annotations + +import asyncio +import os +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from pydantic import ValidationError + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from app.api.routes import auth as auth_route # noqa: E402 + +from packages.adapters.in_memory.user_repository import InMemoryUserRepository # noqa: E402 +from packages.domain.entities import User # noqa: E402 + + +def _auth_user(user): + return SimpleNamespace(user=user, session_id="s-1", token_type="user_auth") + + +def _make_user(**kw): + defaults = dict( + id="u-1", + email="user@example.com", + username="user", + display_name="微信用户", + password_hash="x", + email_verified=True, + profile_completed=False, + ) + defaults.update(kw) + return User(**defaults) + + +# ---------- 请求体校验 ---------- + + +def test_display_name_strips_whitespace(): + req = auth_route.UpdateProfileRequest(display_name=" ying123 ") + assert req.display_name == "ying123" + + +def test_display_name_blank_rejected(): + with pytest.raises(ValidationError) as exc: + auth_route.UpdateProfileRequest(display_name=" ") + assert "空白" in str(exc.value) + + +def test_display_name_empty_rejected(): + with pytest.raises(ValidationError): + auth_route.UpdateProfileRequest(display_name="") + + +def test_display_name_too_long_rejected(): + with pytest.raises(ValidationError) as exc: + auth_route.UpdateProfileRequest(display_name="甲" * 21) + assert "1-20" in str(exc.value) + + +def test_display_name_max_length_accepted(): + req = auth_route.UpdateProfileRequest(display_name="甲" * 20) + assert req.display_name == "甲" * 20 + + +# ---------- 路由逻辑 ---------- + + +def test_patch_me_updates_display_name_and_persists(): + user = _make_user() + repo = InMemoryUserRepository() + repo.save(user) + + resp = asyncio.run( + auth_route.update_current_user_profile( + auth_route.UpdateProfileRequest(display_name=" ying123 "), + current_user=_auth_user(user), + user_repository=repo, + ) + ) + assert resp.user.display_name == "ying123" + assert resp.user.profile_completed is True + assert resp.user.wechat_bound is False + # 落库验证 + fresh = repo.find_by_id("u-1") + assert fresh.display_name == "ying123" + assert fresh.profile_completed is True + + +def test_patch_me_first_time_sets_profile_completed_true(): + user = _make_user(profile_completed=False) + repo = InMemoryUserRepository() + repo.save(user) + assert repo.find_by_id("u-1").profile_completed is False + + asyncio.run( + auth_route.update_current_user_profile( + auth_route.UpdateProfileRequest(display_name="小虾"), + current_user=_auth_user(user), + user_repository=repo, + ) + ) + assert repo.find_by_id("u-1").profile_completed is True + + +def test_patch_me_idempotent_for_completed_user(): + user = _make_user(display_name="老名字", profile_completed=True) + repo = InMemoryUserRepository() + repo.save(user) + + resp = asyncio.run( + auth_route.update_current_user_profile( + auth_route.UpdateProfileRequest(display_name="新名字"), + current_user=_auth_user(user), + user_repository=repo, + ) + ) + assert resp.user.profile_completed is True + assert resp.user.display_name == "新名字" + # 再提交一次同样内容,不报错、状态稳定 + resp2 = asyncio.run( + auth_route.update_current_user_profile( + auth_route.UpdateProfileRequest(display_name="新名字"), + current_user=_auth_user(repo.find_by_id("u-1")), + user_repository=repo, + ) + ) + assert resp2.user.profile_completed is True + + +def test_patch_me_response_contains_all_me_fields(): + user = _make_user(wechat_openid="wx-1", phone="13800000000", phone_verified=True) + repo = InMemoryUserRepository() + repo.save(user) + + resp = asyncio.run( + auth_route.update_current_user_profile( + auth_route.UpdateProfileRequest(display_name="昵称"), + current_user=_auth_user(user), + user_repository=repo, + ) + ) + payload = resp.user.model_dump() + for field in ( + "user_id", + "email", + "username", + "display_name", + "email_verified", + "phone", + "phone_verified", + "binding_complete", + "wechat_bound", + "profile_completed", + ): + assert field in payload, f"missing field {field}" + assert payload["wechat_bound"] is True + assert payload["phone"] == "13800000000" + + +def test_get_me_includes_profile_completed_flag(): + # 未完成 + u = _make_user(profile_completed=False) + resp = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(u))) + assert resp.profile_completed is False + assert resp.wechat_bound is False + + # 已完成 + 已绑微信 + u2 = _make_user(profile_completed=True, wechat_openid="wx-9") + resp2 = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(u2))) + assert resp2.profile_completed is True + assert resp2.wechat_bound is True + + +def test_patch_me_requires_auth_dependency(): + # 路由签名必须依赖 get_current_user,未携带 token 时框架返回 401 + params = ( + auth_route.update_current_user_profile.__wrapped__ + if hasattr(auth_route.update_current_user_profile, "__wrapped__") + else auth_route.update_current_user_profile + ) + import inspect + + sig = inspect.signature(params) + dep = sig.parameters.get("current_user") + assert dep is not None + assert dep.default is not None and getattr(dep.default, "dependency", None) is auth_route.get_current_user + + +def test_wechat_new_user_created_with_profile_completed_false(): + # 微信同步建号:新用户 profile_completed=False(需引导设置昵称) + from packages.application.auth.wechat_sync_use_case import ( + WechatSyncRequest, + WechatSyncUseCase, + ) + + repo = InMemoryUserRepository() + # session_store 用 mock,不依赖 redis + use_case = WechatSyncUseCase(user_repository=repo, session_store=MagicMock(), jwt_secret_key="test-secret") + resp, err = use_case.execute(WechatSyncRequest(openid="wx-new-openid", nickname="微信测试", source="web")) + assert err is None + user = repo.find_by_id(resp.user_id) + assert user.profile_completed is False diff --git a/tests/unit/test_persist_celery_id_routes_1714.py b/tests/unit/test_persist_celery_id_routes_1714.py new file mode 100644 index 000000000..ddfa89dda --- /dev/null +++ b/tests/unit/test_persist_celery_id_routes_1714.py @@ -0,0 +1,259 @@ +"""#1714:入队后 celery_task_id 持久化路径覆盖(routes / enqueue / celery_app / 仓储)。 + +CI 无 redis、不走完整 HTTP 流程,这些 try/except 与早退分支此前覆盖率为 0。 +用真实 SQLite 仓储 + monkeypatch celery_app.send_task 直接驱动路由函数: +- routes/ingest_jobs.submit_ingest_job:正常持久化 + 持久化异常吞掉不影响响应 +- routes/task_center.retry_project_task(ingest 分支):重试后持久化 + 异常吞掉 +- routes/upload._persist_celery_task_id:空 id 早退 + 异常吞掉 +- core/task_enqueue.safe_enqueue_generation_task:持久化失败仅 warning,入队仍 True +- core/celery_app:apply_queue_settings 抛异常时 API 启动不炸 +- adapters/ingest_job_repository.update:写 celery_task_id 分支落库 +""" + +from __future__ import annotations + +import importlib +import importlib.util +import os +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +API_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "api") +if API_PATH not in sys.path: + sys.path.insert(0, API_PATH) + +import pytest # noqa: E402 +from app.api.routes import ingest_jobs as ingest_jobs_route # noqa: E402 +from app.api.routes import task_center as task_center_route # noqa: E402 +from app.api.routes import upload as upload_route # noqa: E402 +from app.schemas.ingest_job import SubmitIngestJobRequest # noqa: E402 +from sqlalchemy import create_engine # noqa: E402 +from sqlalchemy.orm import sessionmaker # noqa: E402 + +from packages.adapters.sqlalchemy_impl.ingest_job_repository import ( # noqa: E402 + SQLAlchemyIngestJobRepository, +) +from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402 +from packages.domain import IngestJob, IngestJobStatus # noqa: E402 + + +def _ingest_repo(): + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(engine) + session = sessionmaker(bind=engine)() + return SQLAlchemyIngestJobRepository(session), session + + +def _fake_celery_result(task_id: str = "celery-route-msg-1"): + result = MagicMock() + result.id = task_id + return result + + +# ── routes/ingest_jobs.submit_ingest_job ──────────────────────────────── + + +def test_submit_ingest_job_persists_celery_task_id(monkeypatch): + repo, session = _ingest_repo() + monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result())) + + request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/x.mov") + response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=repo) + + assert response.status == "pending" + saved = repo.get(response.id) + assert saved.celery_task_id == "celery-route-msg-1" + + +def test_submit_ingest_job_persist_failure_swallowed(monkeypatch): + repo, _ = _ingest_repo() + + class _BoomRepo: + def __init__(self, inner): + self.inner = inner + + def create(self, job): + return self.inner.create(job) + + def get(self, job_id): + return self.inner.get(job_id) + + def update(self, job): # noqa: ARG002 + raise RuntimeError("db write fail") + + boom_repo = _BoomRepo(repo) + monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result())) + + request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/y.mov") + # 持久化异常被吞掉,主流程(响应)不受影响 + response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=boom_repo) + assert response.id + assert response.status == "pending" + + +# ── routes/task_center.retry_project_task(ingest 分支) ──────────────── + + +def _auth_user(): + user = SimpleNamespace(id="user-1") + return SimpleNamespace(user=user, session_id=None, token_type=None) + + +def test_retry_ingest_job_persists_celery_task_id(monkeypatch): + repo, session = _ingest_repo() + # 造一条 failed 的 ingest job + job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/z.mov") + job.status = IngestJobStatus.FAILED + repo.create(job) + + monkeypatch.setattr( + task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result("celery-retry-1")) + ) + + response = task_center_route.retry_project_task( + "ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo + ) + assert response.task_type == "ingest" + new_id = response.id.split("ingest:")[1] + retried = repo.get(new_id) + assert retried is not None + assert retried.celery_task_id == "celery-retry-1" + + +def test_retry_ingest_job_persist_failure_swallowed(monkeypatch): + repo, _ = _ingest_repo() + job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/w.mov") + job.status = IngestJobStatus.FAILED + repo.create(job) + + real_update = repo.update + + def _update_that_booms(entity): + # 仅在写 celery_task_id 的那次 update 抛错(新建 job 后路由内的持久化) + if getattr(entity, "celery_task_id", ""): + raise RuntimeError("db write fail") + return real_update(entity) + + repo.update = _update_that_booms # type: ignore[method-assign] + monkeypatch.setattr(task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result())) + + # 持久化异常吞掉,重试接口仍正常返回 + response = task_center_route.retry_project_task( + "ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo + ) + assert response.task_type == "ingest" + + +# ── routes/upload._persist_celery_task_id ─────────────────────────────── + + +def test_upload_persist_helper_empty_id_early_return(): + repo = MagicMock() + job = MagicMock() + upload_route._persist_celery_task_id(repo, job, "") + repo.update.assert_not_called() + upload_route._persist_celery_task_id(repo, job, None) # type: ignore[arg-type] + repo.update.assert_not_called() + + +def test_upload_persist_helper_exception_swallowed(): + repo = MagicMock() + repo.update.side_effect = RuntimeError("db fail") + job = MagicMock() + # 不抛异常 + upload_route._persist_celery_task_id(repo, job, "celery-upload-1") + repo.update.assert_called_once() + assert job.celery_task_id == "celery-upload-1" + + +# ── core/task_enqueue:持久化失败仅 warning ───────────────────────────── + + +def test_safe_enqueue_persist_failure_still_returns_true(monkeypatch): + from app.core import task_enqueue + + class _FakeTask: + def __init__(self): + self.id = "task-enqueue-persist-fail" + self.status = "pending" + self.celery_task_id = "" + + def mark_failed(self, msg): # noqa: ARG002 + self.status = "failed" + + class _FakeRepo: + def count_pending_total(self): + return 0 + + def count_pending_by_user(self, user_id): # noqa: ARG002 + return 0 + + def update(self, task): # noqa: ARG002 + raise RuntimeError("persist celery_task_id failed") + + fake_result = MagicMock() + fake_result.id = "celery-enqueue-fail-1" + mock_celery = MagicMock() + mock_celery.send_task.return_value = fake_result + monkeypatch.setattr(task_enqueue, "celery_app", mock_celery) + + task = _FakeTask() + repo = _FakeRepo() + ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1") + # 持久化失败不影响入队结果 + assert ok is True + mock_celery.send_task.assert_called_once() + + +# ── core/celery_app:队列配置失败不阻断 API 启动 ──────────────────────── + + +def test_api_celery_app_survives_queue_settings_failure(monkeypatch): + """apply_queue_settings 抛异常时 API 启动不炸(core/celery_app.py 的 try/except 分支)。 + + 通过让 `from packages.shared.celery_queues import apply_queue_settings` 本身 + 抛异常来触发 except 分支;用全新模块名 reload,不替换已被其他模块持有的 + app.core.celery_app 模块对象,避免污染 task_enqueue 等导入方。 + """ + import builtins + + real_import = builtins.__import__ + + def _failing_import(name, globals=None, locals=None, fromlist=(), level=0): # noqa: A002 + if name == "packages.shared.celery_queues" and "apply_queue_settings" in (fromlist or ()): + raise RuntimeError("config boom") + return real_import(name, globals, locals, fromlist, level) + + monkeypatch.setattr(builtins, "__import__", _failing_import) + + spec = importlib.util.find_spec("app.core.celery_app") + fresh_mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(fresh_mod) # 异常在模块内被 try/except 吞掉 + assert fresh_mod.celery_app is not None + assert fresh_mod.celery_app.main == "xiaoxia-saas-api" + + # 已加载的原模块对象不受影响(无 reload 污染) + import app.core.celery_app as api_celery_mod + + assert api_celery_mod.celery_app is not None + + +# ── 仓储:update 写 celery_task_id 落库 ───────────────────────────────── + + +def test_ingest_repo_update_persists_celery_task_id(): + repo, session = _ingest_repo() + job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/repo.mov") + repo.create(job) + + job.celery_task_id = "celery-repo-update-1" + repo.update(job) + + session.expire_all() + saved = repo.get(job.id) + assert saved.celery_task_id == "celery-repo-update-1" diff --git a/tests/unit/test_phash_threshold_calibration_1658.py b/tests/unit/test_phash_threshold_calibration_1658.py index 6e01b4a6a..cfefdd077 100644 --- a/tests/unit/test_phash_threshold_calibration_1658.py +++ b/tests/unit/test_phash_threshold_calibration_1658.py @@ -102,6 +102,7 @@ from video_processing.dedup import ( # noqa: E402 DUPLICATE_THRESHOLD, HISTOGRAM_WEIGHT, MATCH_RATIO_THRESHOLD, + PHASH_THRESHOLD, PHASH_WEIGHT, VideoDeduplicator, ) @@ -128,11 +129,17 @@ _ZERO_HIST = [0.0] * 96 # 全黑视频的全零直方图(有效数据) class TestThresholdCalibration: - """pHash 阈值由 10 收紧到 8(Issue #1658)。""" + """pHash 阈值校准(#1658 收紧到 8,#1702 两轮真实数据重校准 12→16)。 - def test_phash_threshold_is_8(self): - """PHASH_THRESHOLD 必须为 8(旧值 10 会放过 8~9 汉明距离的不同视频)。""" - assert VideoDeduplicator.PHASH_THRESHOLD == 8 + #1702 第一轮 staging 离线实验:同帧两次 2-5% 随机裁剪距离 4~10;同源成片 + (密集 1s 采样)<=12 命中 4/11、异源成片最小距离 24 → 初定 12。 + #1702 第二轮(证据视频 B->A 仍漏检)扩样本到该用户 15 个真实成片实测: + 同源降重对中位数距离 14、<=16 命中 8/11=0.73;异源 13 个候选 <=16 命中 + 全 0、每帧全局最近邻最小距离 18 → 校准为 16(与异源仍有 >=2bit 裕度)。 + """ + + def test_phash_threshold_is_calibrated(self): + assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD == 16 def test_match_ratio_threshold_constant(self): assert MATCH_RATIO_THRESHOLD == 0.7 @@ -144,22 +151,22 @@ class TestThresholdCalibration: assert PHASH_WEIGHT == 0.7 assert HISTOGRAM_WEIGHT == 0.3 - def test_threshold_tightening_excludes_distance_8_and_9(self): - """距离 8、9 的帧:旧阈值 10 下算匹配,新阈值 8 下不算匹配。 + def test_threshold_matching_semantics(self): + """阈值比较统一为 <=(帧匹配与片段匹配同一口径)。 - 场景:5 个关键帧距离为 [7, 7, 7, 9, 9]。 - - 旧阈值 10:5 帧全部 < 10 → match_ratio = 1.0(误放过) - - 新阈值 8:仅 3 帧 < 8 → match_ratio = 0.6 < 0.7(正确跳过) + 场景:5 个关键帧距离为 [10, 14, 16, 18, 26]。 + - <=16(#1702 二次校准阈值):3 帧匹配 → 0.6 < 0.7,被帧比例门槛 + 拦截(异源安全边界:真实数据异源最近邻最小距离 18,<=16 命中 0) + - 距离正好 16 的同源降重帧应算匹配(< 与 <= 口径统一) """ - distances = [7, 7, 7, 9, 9] + distances = [10, 14, 16, 18, 26] + matched = sum(1 for d in distances if d <= VideoDeduplicator.PHASH_THRESHOLD) + assert matched == 3 + assert matched / len(distances) == 0.6 + assert matched / len(distances) < MATCH_RATIO_THRESHOLD - matched_old = sum(1 for d in distances if d < 10) - assert matched_old == 5 # 旧行为:全匹配 → 误判风险 - - matched_new = sum(1 for d in distances if d < VideoDeduplicator.PHASH_THRESHOLD) - assert matched_new == 3 - assert matched_new / len(distances) == 0.6 - assert matched_new / len(distances) < MATCH_RATIO_THRESHOLD # 被帧比例门槛拦截 + # 异源安全边界(实测最小距离 18)及以上绝不匹配 + assert not any(d <= VideoDeduplicator.PHASH_THRESHOLD for d in (18, 24, 26, 30)) # ── TestComputeFusionScore:统一融合得分方法 ──────────────────── diff --git a/tests/unit/test_prepare_dedup_1714.py b/tests/unit/test_prepare_dedup_1714.py new file mode 100644 index 000000000..cb0ad9ff8 --- /dev/null +++ b/tests/unit/test_prepare_dedup_1714.py @@ -0,0 +1,471 @@ +"""#1714 prepare_direct_upload 去重 + 预建 asset 测试。 + +覆盖 4 类用例: +- 第一次上传:prepare 返回 duplicated=false + asset_id 非空 +- 第二次同 hash:prepare 返回 duplicated=true, skip_transfer=true +- 同 client_upload_id 重试:prepare 也直接跳过 +- file_hash 空:走老逻辑,duplicated=false,无 asset_id + +以及: +- pre-create 的 PROCESSING 占位不被"文件名兜底去重"误命中 +- _create_pending_asset find-or-create 复用现有记录 +""" + +from __future__ import annotations + +import os +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) +sys.path.insert(0, str(Path(__file__).resolve().parents[2])) + +from apps.api.app.api.routes import upload as upload_route # noqa: E402 +from packages.domain.entities import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, Project # noqa: E402 + +# --------------------------------------------------------------------------- +# Fake repository +# --------------------------------------------------------------------------- + + +class _FakeAssetRepo: + """内存 asset 仓储:实现 prepare/complete 去重需要的所有方法。""" + + def __init__(self): + self.assets = {} # id -> Asset + self.saved = 0 + self.updated = 0 + + def create(self, asset): + self.assets[asset.id] = asset + self.saved += 1 + return asset + + def update(self, asset): + self.assets[asset.id] = asset + self.updated += 1 + return asset + + def find_by_id(self, asset_id): + return self.assets.get(asset_id) + + def find_by_library_and_file_hash(self, library_id, file_hash): + if not file_hash: + return None + for a in self.assets.values(): + if a.library_id == library_id and a.file_hash == file_hash: + return a + return None + + def find_by_library_and_client_upload_id(self, library_id, client_upload_id): + if not client_upload_id: + return None + for a in self.assets.values(): + if a.library_id == library_id and a.client_upload_id == client_upload_id: + return a + return None + + def find_recent_active_by_library_and_name(self, library_id, name, within_minutes=30, file_size=0): + return None + + +def _make_asset(**kw): + defaults = dict( + project_id="p-1", + library_id="lib-1", + name="existing.mp4", + storage_key="uploads/old/existing.mp4", + mime_type="video/mp4", + status=AssetStatus.READY, + file_hash="existinghash", + ) + defaults.update(kw) + return Asset(id=defaults.pop("id", "existing-asset"), **defaults) + + +def _make_pending(**kw): + defaults = dict( + project_id="p-1", + library_id="lib-1", + name="test.mp4", + storage_key="uploads/abc/test.mp4", + mime_type="video/mp4", + status=AssetStatus.PROCESSING, + file_hash="abc123", + ) + defaults.update(kw) + return Asset(id=defaults.pop("id", "pending-asset"), **defaults) + + +def _user(): + return SimpleNamespace(user=SimpleNamespace(id="user-1"), session_id="s", token_type="t") + + +class _StubProjectRepo: + def __init__(self, project): + self._p = project + + def get(self, pid): + return self._p if self._p.id == pid else None + + def find_by_id(self, pid): + return self._p if self._p.id == pid else None + + +class _StubLibraryRepo: + def __init__(self, lib): + self._lib = lib + + def find_by_project(self, pid, kind=None): + if self._lib.project_id == pid: + return [self._lib] + return [] + + +_FIXTURE_PROJECT = Project(id="p-1", owner_user_id="user-1", name="proj", description="") +_FIXTURE_LIBRARY = AssetLibrary( + id="lib-1", project_id="p-1", name="videos", kind=AssetLibraryKind.VIDEO, asset_count=0, total_size=0 +) + + +def _storage(): + s = MagicMock() + s.create_direct_upload_post.return_value = { + "url": "https://bucket.oss.example.com", + "method": "POST", + "storage_key": "uploads/abc/test.mp4", + "expires_at": "2026-01-01T00:00:00Z", + "fields": {"key": "uploads/abc/test.mp4"}, + } + return s + + +# --------------------------------------------------------------------------- +# 场景 1:第一次上传(无 file_hash) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_prepare_first_upload_no_hash_returns_no_dedup(): + repo = _FakeAssetRepo() + req = SimpleNamespace( + project_id="p-1", + library_id="lib-1", + filename="test.mp4", + content_type="video/mp4", + file_size=1024, + file_hash="", + client_upload_id="", + ) + resp = await upload_route.prepare_direct_upload( + request=req, + authenticated_user=_user(), + project_repository=_StubProjectRepo(_FIXTURE_PROJECT), + asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY), + asset_repository=repo, + storage_service=_storage(), + ) + assert resp.duplicated is False + assert resp.skip_transfer is False + assert resp.asset_id == "" # file_hash 空,不预建 + assert repo.saved == 0 + + +# --------------------------------------------------------------------------- +# 场景 2:第一次上传带 file_hash → duplicated=false + asset_id 非空 +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_prepare_first_upload_with_hash_creates_pending(): + repo = _FakeAssetRepo() + req = SimpleNamespace( + project_id="p-1", + library_id="lib-1", + filename="test.mp4", + content_type="video/mp4", + file_size=1024, + file_hash="abc123", + client_upload_id="", + ) + resp = await upload_route.prepare_direct_upload( + request=req, + authenticated_user=_user(), + project_repository=_StubProjectRepo(_FIXTURE_PROJECT), + asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY), + asset_repository=repo, + storage_service=_storage(), + ) + assert resp.duplicated is False + assert resp.skip_transfer is False + assert resp.asset_id != "" + # 预建记录确实落库 + assert repo.saved == 1 + pending = repo.find_by_id(resp.asset_id) + assert pending is not None + assert pending.file_hash == "abc123" + assert pending.status == AssetStatus.PROCESSING + + +# --------------------------------------------------------------------------- +# 场景 3:第二次同 hash → duplicated=true, skip_transfer=true +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_prepare_second_upload_same_hash_returns_duplicated(): + repo = _FakeAssetRepo() + repo.create(_make_pending(file_hash="abc123", id="existing-asset")) + req = SimpleNamespace( + project_id="p-1", + library_id="lib-1", + filename="test.mp4", + content_type="video/mp4", + file_size=1024, + file_hash="abc123", + client_upload_id="", + ) + resp = await upload_route.prepare_direct_upload( + request=req, + authenticated_user=_user(), + project_repository=_StubProjectRepo(_FIXTURE_PROJECT), + asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY), + asset_repository=repo, + storage_service=_storage(), + ) + assert resp.duplicated is True + assert resp.skip_transfer is True + assert resp.asset_id == "existing-asset" + assert resp.upload_url == "" # 未签名 OSS + # 未新增记录 + assert repo.saved == 1 # 只有初始那条 + + +# --------------------------------------------------------------------------- +# 场景 4:同 client_upload_id 重试 → 直接跳过 +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_prepare_retry_same_client_upload_id_skips(): + repo = _FakeAssetRepo() + repo.create( + _make_pending( + file_hash="abc123", + client_upload_id="cuid-xyz", + id="existing-asset", + ) + ) + # 即使 file_hash 不同(理论上不会),client_upload_id 命中也直接跳过 + req = SimpleNamespace( + project_id="p-1", + library_id="lib-1", + filename="test.mp4", + content_type="video/mp4", + file_size=1024, + file_hash="different-hash", + client_upload_id="cuid-xyz", + ) + resp = await upload_route.prepare_direct_upload( + request=req, + authenticated_user=_user(), + project_repository=_StubProjectRepo(_FIXTURE_PROJECT), + asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY), + asset_repository=repo, + storage_service=_storage(), + ) + assert resp.duplicated is True + assert resp.skip_transfer is True + assert resp.asset_id == "existing-asset" + + +# --------------------------------------------------------------------------- +# 兜底:文件名兜底去重不误命中 PROCESSING 占位 +# --------------------------------------------------------------------------- + + +def test_filename_fallback_does_not_match_processing_pending(): + """_find_duplicate_asset 按文件名兜底时,不能命中 pre-create 的 PROCESSING 记录。""" + repo = _FakeAssetRepo() + repo.create(_make_pending(id="p1")) + result = upload_route._find_duplicate_asset( + repo, + library_id="lib-1", + file_hash="", # 无 hash + client_upload_id="", # 无 cuid + filename="test.mp4", # 同名 + file_size=1024, + ) + assert result is None # PROCESSING 占位不被兜底命中 + + +def test_filename_fallback_matches_stable_ready_record(): + """READY 状态的已存在记录能被文件名兜底命中。""" + repo = _FakeAssetRepo() + repo.create(_make_asset(status=AssetStatus.READY, id="ready-asset")) + # 伪造 find_recent_active_by_library_and_name 返回 READY 记录 + repo.find_recent_active_by_library_and_name = lambda **kw: repo.assets["ready-asset"] + result = upload_route._find_duplicate_asset( + repo, + library_id="lib-1", + file_hash="", + client_upload_id="", + filename="existing.mp4", + file_size=1024, + ) + assert result is not None + assert result.id == "ready-asset" + + +# --------------------------------------------------------------------------- +# _create_pending_asset find-or-create +# --------------------------------------------------------------------------- + + +def test_create_pending_asset_reuses_existing_by_hash(): + """_create_pending_asset:file_hash 命中现有 PROCESSING 记录则复用,不新建。""" + repo = _FakeAssetRepo() + repo.create(_make_pending(file_hash="abc123", client_upload_id="", id="p1")) + # 复用 + result = upload_route._create_pending_asset( + asset_repository=repo, + project_id="p-1", + library_id="lib-1", + storage_key="uploads/new/test.mp4", + filename="test.mp4", + mime_type="video/mp4", + user_id="user-1", + file_hash="abc123", + client_upload_id="cuid-new", + ) + assert result.id == "p1" + assert repo.saved == 1 # 没新增 + assert repo.updated >= 1 # 字段补齐触发 update + assert result.client_upload_id == "cuid-new" + + +def test_create_pending_asset_creates_when_no_match(): + """无匹配时正常新建。""" + repo = _FakeAssetRepo() + result = upload_route._create_pending_asset( + asset_repository=repo, + project_id="p-1", + library_id="lib-1", + storage_key="uploads/new/test.mp4", + filename="test.mp4", + mime_type="video/mp4", + user_id="user-1", + file_hash="newhash", + client_upload_id="newcuid", + ) + assert result.id != "" + assert result.file_hash == "newhash" + assert result.client_upload_id == "newcuid" + assert repo.saved == 1 + + +# --------------------------------------------------------------------------- +# 兜底去重:PROCESSING 占位 hash 不同时跳过 +# --------------------------------------------------------------------------- + + +def test_filename_fallback_skips_processing_with_different_hash(): + """PROCESSING/UPLOADING 占位记录仅当 hash 一致(或占位无 hash)才命中;hash 不同跳过。""" + repo = _FakeAssetRepo() + repo.create(_make_pending(id="p1", file_hash="oldhash")) + repo.find_recent_active_by_library_and_name = lambda **kw: repo.assets["p1"] + result = upload_route._find_duplicate_asset( + repo, + library_id="lib-1", + file_hash="differenthash", # 新上传内容不同 + client_upload_id="", + filename="test.mp4", + file_size=1024, + ) + assert result is None + + +def test_filename_fallback_matches_processing_with_same_hash(): + """PROCESSING 占位 hash 与请求一致时命中(重试场景)。""" + repo = _FakeAssetRepo() + repo.create(_make_pending(id="p1", file_hash="samehash")) + repo.find_recent_active_by_library_and_name = lambda **kw: repo.assets["p1"] + result = upload_route._find_duplicate_asset( + repo, + library_id="lib-1", + file_hash="samehash", + client_upload_id="", + filename="test.mp4", + file_size=1024, + ) + assert result is not None + assert result.id == "p1" + + +# --------------------------------------------------------------------------- +# prepare 预建失败降级:不阻塞签名 +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_prepare_pending_asset_create_failure_degrades_gracefully(): + """预建 asset 抛异常时,prepare 仍正常返回签名(duplicated=False, asset_id 空)。""" + + class _BrokenRepo(_FakeAssetRepo): + def create(self, asset): + raise RuntimeError("db down") + + repo = _BrokenRepo() + req = SimpleNamespace( + project_id="p-1", + library_id="lib-1", + filename="test.mp4", + content_type="video/mp4", + file_size=1024, + file_hash="abc123", + client_upload_id="cuid-1", + ) + resp = await upload_route.prepare_direct_upload( + request=req, + authenticated_user=_user(), + project_repository=_StubProjectRepo(_FIXTURE_PROJECT), + asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY), + asset_repository=repo, + storage_service=_storage(), + ) + assert resp.duplicated is False + assert resp.skip_transfer is False + assert resp.asset_id == "" # 预建失败,降级无 asset_id + assert resp.upload_url != "" # 签名仍正常返回 + + +def test_create_pending_asset_update_failure_swallowed(): + """复用占位记录时字段补齐 update 抛异常被吞掉,不阻塞返回。""" + + class _UpdateBrokenRepo(_FakeAssetRepo): + def update(self, asset): + raise RuntimeError("db down") + + repo = _UpdateBrokenRepo() + repo.create(_make_pending(file_hash="abc123", client_upload_id="", id="p1")) + result = upload_route._create_pending_asset( + asset_repository=repo, + project_id="p-1", + library_id="lib-1", + storage_key="uploads/new/test.mp4", + filename="test.mp4", + mime_type="video/mp4", + user_id="user-1", + file_hash="abc123", + client_upload_id="cuid-new", + file_size=1024, + ) + assert result.id == "p1" # 仍复用,不抛异常 + assert repo.saved == 1 diff --git a/tests/unit/test_stale_task_revoke_1714.py b/tests/unit/test_stale_task_revoke_1714.py new file mode 100644 index 000000000..4ac086402 --- /dev/null +++ b/tests/unit/test_stale_task_revoke_1714.py @@ -0,0 +1,204 @@ +"""Issue #1714:孤儿/超时清理标记 failed 时必须撤销并清除 Redis 队列消息。 + +覆盖: +- cleanup_stale_pending_with_session_ids:超时 pending 标记 failed 并返回 + (task_id, celery_task_id),worker 清理流程据此 revoke + purge 队列消息 +- 队列中对应业务任务的 celery 消息被物理移除(作废消息不会重投执行) +- 旧仓储(无 _with_ids 方法)降级为计数模式,不抛异常 +- cleanup_stale_running_with_ids 同样返回 id 列表 +""" + +from __future__ import annotations + +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path +from unittest.mock import MagicMock + +import pytest +from sqlalchemy import create_engine, text +from sqlalchemy.orm import sessionmaker + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from packages.adapters.sqlalchemy_impl.generation_task_repository import ( # noqa: E402 + SQLAlchemyGenerationTaskRepository, +) +from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402 +from packages.domain import GenerationTask, GenerationTaskStatus # noqa: E402 + +BROKER_URL = "redis://localhost:6379/15" +TEST_QUEUE = "_test_revoke_q" + + +def _repository(): + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(engine) + session = sessionmaker(bind=engine)() + return SQLAlchemyGenerationTaskRepository(session), session, engine + + +def _make_task(**kwargs) -> GenerationTask: + defaults = dict(project_id="proj-1", asset_library_id="lib-1", created_by_user_id="user-1") + defaults.update(kwargs) + return GenerationTask.create(**defaults) + + +def _redis_available() -> bool: + try: + import redis + + return bool(redis.Redis.from_url(BROKER_URL).ping()) + except Exception: + return False + + +# ── 仓储层:返回 ids ──────────────────────────────────────────────────── + + +def test_cleanup_stale_pending_returns_ids_with_celery_task_id(): + repo, _, engine = _repository() + task = _make_task() + task.celery_task_id = "celery-msg-id-001" + repo.create(task) + # created_at 改到 60 分钟前 + with engine.connect() as conn: + conn.execute( + text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"), + {"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id}, + ) + conn.commit() + + items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45) + assert len(items) == 1 + biz_id, celery_id = items[0] + assert biz_id == task.id + assert celery_id == "celery-msg-id-001" + + saved = repo.get(task.id) + assert saved.status == GenerationTaskStatus.FAILED + + +def test_cleanup_stale_running_returns_ids(): + repo, _, engine = _repository() + task = _make_task() + repo.create(task) + task.mark_processing() + task.celery_task_id = "celery-msg-id-002" + repo.update(task) + with engine.connect() as conn: + conn.execute( + text("UPDATE generation_tasks SET updated_at = :ts WHERE id = :id"), + {"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id}, + ) + conn.commit() + + items = repo.cleanup_stale_running_with_ids(timeout_minutes=20) + assert len(items) == 1 + assert items[0][0] == task.id + assert items[0][1] == "celery-msg-id-002" + assert repo.get(task.id).status == GenerationTaskStatus.FAILED + + +def test_legacy_repo_without_with_ids_falls_back_to_count(): + """旧仓储只有 cleanup_stale_pending(返回 int)时降级可用,不抛异常。""" + # worker 模块加载(标准 mock 模式) + saved = set(sys.modules.keys()) + mock_db = MagicMock() + mock_db.SessionLocal = MagicMock() + sys.modules["worker_app.db"] = mock_db + sys.modules["worker_app.core.config"] = MagicMock() + mock_celery = MagicMock() + mock_celery.celery_app.task = MagicMock( + side_effect=(lambda *a, **k: (a[0] if a and callable(a[0]) else (lambda f: f))) + ) + sys.modules["worker_app.celery_app"] = mock_celery + worker_path = str(Path(__file__).resolve().parents[2] / "apps" / "worker") + if worker_path not in sys.path: + sys.path.insert(0, worker_path) + + from worker_app.tasks import _startup # noqa: E402 + + class LegacyRepo: + def cleanup_stale_pending(self, timeout_minutes): # noqa: ARG002 + return 3 + + def cleanup_stale_running(self, timeout_minutes): # noqa: ARG002 + return 2 + + items_p = _startup.cleanup_stale_pending_with_session_ids(LegacyRepo(), 45) + items_r = _startup.cleanup_stale_running_with_session_ids(LegacyRepo(), 20) + assert len(items_p) == 3 + assert len(items_r) == 2 + + for key in list(sys.modules.keys()): + if key not in saved and not key.startswith("video_processing"): + del sys.modules[key] + + +# ── 端到端:清理 → 队列消息被移除(作废消息不重投) ──────────────────── + + +@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用") +def test_stale_pending_cleanup_purges_redis_message(): + """任务标 failed 后,其在 Redis 队列里的 celery 消息被清除,不会被重投。""" + import redis + from celery import Celery + from kombu import Queue + from kombu.pools import producers + + from packages.shared.celery_orphan_guard import purge_stale_messages_from_queues + + repo, _, engine = _repository() + task = _make_task() + task.celery_task_id = "celery-stale-xyz" + repo.create(task) + with engine.connect() as conn: + conn.execute( + text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"), + {"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id}, + ) + conn.commit() + + # 模拟该任务的 celery 消息仍在 generation 队列里(worker 下线期间未消费) + client = redis.Redis.from_url(BROKER_URL) + client.delete(TEST_QUEUE) + app = Celery("test-e2e-revoke") + app.conf.broker_url = BROKER_URL + with app.connection_for_write() as conn: + with producers[conn].acquire(block=True) as prod: + # 作废任务消息 + prod.publish( + (task.id,), + exchange="", + routing_key=TEST_QUEUE, + serializer="json", + headers={"id": "celery-stale-xyz", "task": "worker.generate_video"}, + retry=False, + delivery_mode=1, + declare=[Queue(TEST_QUEUE, routing_key=TEST_QUEUE, durable=False)], + ) + # 另一条正常任务消息(必须保留) + prod.publish( + ("other-task-id",), + exchange="", + routing_key=TEST_QUEUE, + serializer="json", + headers={"id": "celery-keep", "task": "worker.generate_video"}, + retry=False, + delivery_mode=1, + ) + + assert client.llen(TEST_QUEUE) == 2 + + # 执行清理(与 worker beat 相同流程:标 failed → 拿 ids → purge) + items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45) + biz_ids = [bid for bid, _ in items] + celery_ids = [cid for _, cid in items if cid] + removed = purge_stale_messages_from_queues( + BROKER_URL, (TEST_QUEUE,), business_task_ids=biz_ids, celery_task_ids=celery_ids + ) + + assert removed == 1 + assert client.llen(TEST_QUEUE) == 1 # 正常任务消息保留 + client.delete(TEST_QUEUE) diff --git a/tests/unit/test_task_discard_guard_1714.py b/tests/unit/test_task_discard_guard_1714.py new file mode 100644 index 000000000..faa1cfd96 --- /dev/null +++ b/tests/unit/test_task_discard_guard_1714.py @@ -0,0 +1,243 @@ +"""Issue #1714:任务执行前状态守卫 — 已作废消息必须丢弃,禁止非法转换后继续跑。 + +覆盖: +- ingest_asset:job 已 failed/completed 时直接返回 discarded,不下载、不转码、不回写 +- generate_video:GenerationTask 已 failed 时返回 discarded,不进入渲染 +- generate_video:pending → running 标记失败(非法转换)时安全中止 +""" + +from __future__ import annotations + +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock + +# ── worker 模块标准加载方式 ── +# 显式保存将要覆盖的注入键旧值:全量收集时更早的测试文件(如 +# test_ingest_validation.py)可能已向 sys.modules 注入 worker_app.* mock, +# 导入完成后必须精确恢复旧值,否则本文件的 bind 感知透传装饰器会残留, +# 污染后续懒加载短路径 worker_app.celery_app 的 worker 测试。 +_INJECTED_KEYS = ("worker_app.db", "worker_app.core.config", "worker_app.celery_app") +_SAVED_MODULE_VALUES = {k: sys.modules.get(k) for k in _INJECTED_KEYS} +_SAVED_MODULES_KEYS = set(sys.modules.keys()) + +_mock_db_module = MagicMock() +_mock_db_module.SessionLocal = MagicMock() +sys.modules["worker_app.db"] = _mock_db_module +sys.modules["worker_app.core.config"] = MagicMock() + +_mock_celery_module = MagicMock() + + +def _passthrough_decorator(*args, **kwargs): + if len(args) == 1 and callable(args[0]): + return args[0] + bind = kwargs.get("bind", False) + + def _wrap(f): + if bind: + # 模拟 celery bind=True:task(task_id) 调用时注入 self(MagicMock) + return lambda *a, **kw: f(MagicMock(), *a, **kw) + return f + + return _wrap + + +_mock_celery_module.celery_app.task = MagicMock(side_effect=_passthrough_decorator) +sys.modules["worker_app.celery_app"] = _mock_celery_module + +_WORKER_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "worker") +sys.path.insert(0, _WORKER_PATH) + +import pytest # noqa: E402 +from worker_app.tasks import ingest as ingest_mod # noqa: E402 + +# video_processing 相关 mock(generation 模块导入链) +for _mod_name in [ + "video_processing", + "video_processing.ffmpeg_utils", + "video_processing.oss_helpers", +]: + sys.modules.setdefault(_mod_name, MagicMock()) + +from worker_app.tasks import generation as gen_mod # noqa: E402 + +from packages.domain import IngestJobStatus # noqa: E402 + +# 模块导入完成后立即清理:删除本次 import 新引入的模块缓存(本模块已通过名字绑定 +# 持有 ingest_mod/gen_mod/IngestJobStatus,删除缓存不影响调用),再把三个注入键 +# 精确恢复为注入前的旧值(旧值不存在则移除),杜绝 mock 残留污染其他 worker 测试。 +for _key in list(sys.modules.keys()): + if _key not in _SAVED_MODULES_KEYS and not _key.startswith("video_processing"): + del sys.modules[_key] +for _k, _v in _SAVED_MODULE_VALUES.items(): + if _v is None: + sys.modules.pop(_k, None) + else: + sys.modules[_k] = _v +del _SAVED_MODULES_KEYS, _SAVED_MODULE_VALUES + + +# ── ingest 守卫 ──────────────────────────────────────────────────────── + + +class _FakeJobRepo: + def __init__(self, job): + self.job = job + + def get(self, job_id): + return self.job + + +def _make_ingest_job(status): + job = MagicMock() + job.id = "job-stale-1" + job.storage_key = "uploads/proj/stale.mov" + job.status = status + job.file_hash = "h" + job.asset_id = "" + return job + + +def test_ingest_discards_failed_job_message(): + """job 已 failed:消息丢弃,不进入下载/转码/回写。""" + job = _make_ingest_job(IngestJobStatus.FAILED) + fake_session = MagicMock() + _mock_db_module.SessionLocal = MagicMock(return_value=fake_session) + + # SQLAlchemy 仓储构造返回 fake + fake_job_repo = _FakeJobRepo(job) + fake_asset_repo = MagicMock() + + orig_job_repo = ingest_mod.SQLAlchemyIngestJobRepository + orig_asset_repo = ingest_mod.SQLAlchemyAssetRepository + ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo) + ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=fake_asset_repo) + try: + result = ingest_mod.ingest_asset("job-stale-1") + finally: + ingest_mod.SQLAlchemyIngestJobRepository = orig_job_repo + ingest_mod.SQLAlchemyAssetRepository = orig_asset_repo + + assert result["status"] == "discarded" + # 没有任何 update / commit / 下载动作 + fake_session.commit.assert_not_called() + fake_asset_repo.create.assert_not_called() + + +def test_ingest_discards_completed_job_message(): + job = _make_ingest_job(IngestJobStatus.COMPLETED) + fake_session = MagicMock() + _mock_db_module.SessionLocal = MagicMock(return_value=fake_session) + fake_job_repo = _FakeJobRepo(job) + + orig = ingest_mod.SQLAlchemyIngestJobRepository + ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo) + ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=MagicMock()) + try: + result = ingest_mod.ingest_asset("job-stale-1") + finally: + ingest_mod.SQLAlchemyIngestJobRepository = orig + + assert result["status"] == "discarded" + + +# ── generation 守卫 ──────────────────────────────────────────────────── + + +def _make_gen_task(status_value: str): + from packages.domain import GenerationTask + + task = GenerationTask.create(project_id="p", asset_library_id="l", created_by_user_id="u") + task.status = type(task.status)(status_value) + return task + + +def test_generate_video_discards_failed_task(monkeypatch): + """GenerationTask 已 failed:直接 discarded,不加载渲染数据。""" + failed_task = _make_gen_task("failed") + + fake_repo = MagicMock() + fake_repo.get.return_value = failed_task + + fake_session = MagicMock() + _mock_db_module.SessionLocal = MagicMock(return_value=fake_session) + + import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod + + orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository + gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo) + + update_status_mock = MagicMock(return_value=False) + monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock) + monkeypatch.setattr( + gen_mod, + "_load_task_info", + lambda task_id: { + "project_id": "p", + "template_id": "", + "task_asset_ids": [], + "batch_id": "", + "user_id": "u", + "mode": "one_take", + }, + ) + monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None) + + task_fn = gen_mod.generate_video + if hasattr(task_fn, "__wrapped__"): + task_fn = task_fn.__wrapped__ + try: + result = task_fn("task-stale-1") + finally: + gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig + + assert result["status"] == "discarded" + # 状态守卫命中终态,根本不应尝试 mark_processing + update_status_mock.assert_not_called() + + +def test_generate_video_aborts_when_claim_fails(monkeypatch): + """pending 但 mark_processing 返回 False(状态机非法转换)时安全中止。""" + pending_task = _make_gen_task("pending") + + fake_repo = MagicMock() + fake_repo.get.return_value = pending_task + fake_session = MagicMock() + _mock_db_module.SessionLocal = MagicMock(return_value=fake_session) + + import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod + + orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository + gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo) + + monkeypatch.setattr( + gen_mod, + "_load_task_info", + lambda task_id: { + "project_id": "p", + "template_id": "", + "task_asset_ids": [], + "batch_id": "", + "user_id": "u", + "mode": "one_take", + }, + ) + monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None) + # 模拟 mark_processing 失败(failed→running 非法转换被 _update_task_status 吞掉返回 False) + update_status_mock = MagicMock(return_value=False) + monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock) + render_mock = MagicMock(side_effect=AssertionError("must not render")) + monkeypatch.setattr(gen_mod, "_render_from_edit_plan", render_mock) + + task_fn = gen_mod.generate_video + if hasattr(task_fn, "__wrapped__"): + task_fn = task_fn.__wrapped__ + try: + result = task_fn("task-claim-fail") + finally: + gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig + + assert result["status"] == "discarded" + render_mock.assert_not_called() diff --git a/tests/unit/test_task_fault_tolerance_1709.py b/tests/unit/test_task_fault_tolerance_1709.py new file mode 100644 index 000000000..756543d2a --- /dev/null +++ b/tests/unit/test_task_fault_tolerance_1709.py @@ -0,0 +1,315 @@ +"""Issue #1709 任务容错:孤儿任务恢复 + 429 限流结构化提示。 + +覆盖: +1. 仓储层:count_running_by_user/count_running_total 计数正确(预览/正式任务都计入) +2. 仓储层:estimate_avg_duration_seconds 耗时估算(有历史/无历史) +3. 限流核心:build_rate_limit_detail 返回结构化 code/message/排队数/预计等待 +4. worker 侧:cleanup_stale_running/pending 核心函数——中断任务被重置为 failed + 且原因写明(容器重启/超时中断),正常任务不受影响 +""" + +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path +from unittest.mock import MagicMock + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker")) + +# 预注入 mock worker_app.db,防止真实数据库连接初始化(与其他 worker 测试同模式) +_mock_db = MagicMock() +_mock_db.SessionLocal = MagicMock() +sys.modules.setdefault("worker_app.db", _mock_db) + +from app.core import task_enqueue # noqa: E402 +from sqlalchemy import create_engine, text # noqa: E402 +from sqlalchemy.orm import sessionmaker # noqa: E402 +from worker_app.tasks import _startup # noqa: E402 + +from packages.adapters.sqlalchemy_impl.generation_task_repository import ( # noqa: E402 + SQLAlchemyGenerationTaskRepository, +) +from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402 +from packages.domain import GenerationTask, GenerationTaskStatus # noqa: E402 + + +def _repository(): + engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}) + Base.metadata.create_all(engine) + session = sessionmaker(bind=engine)() + return SQLAlchemyGenerationTaskRepository(session), session, engine + + +def _make_task(**kwargs) -> GenerationTask: + defaults = dict( + project_id="proj-1", + asset_library_id="lib-1", + created_by_user_id="user-1", + ) + defaults.update(kwargs) + return GenerationTask.create(**defaults) + + +def _age_task(engine, task_id, *, updated_minutes=None, created_minutes=None): + """用 SQL 直接把 updated_at/created_at 改到过去(模拟孤儿任务)。""" + sets, params = [], {"id": task_id} + if updated_minutes is not None: + sets.append("updated_at = :uts") + params["uts"] = datetime.now(timezone.utc) - timedelta(minutes=updated_minutes) + if created_minutes is not None: + sets.append("created_at = :cts") + params["cts"] = datetime.now(timezone.utc) - timedelta(minutes=created_minutes) + with engine.connect() as conn: + conn.execute(text(f"UPDATE generation_tasks SET {', '.join(sets)} WHERE id = :id"), params) + conn.commit() + + +# --------------------------------------------------------------------------- +# 1. running 计数(限流"渲染中"数量) +# --------------------------------------------------------------------------- + + +def test_count_running_by_user_mix_statuses(): + """count_running_by_user 只统计该用户 running,不含 pending/completed/failed。""" + repo, _, _ = _repository() + t1 = _make_task(project_id="p1") + repo.create(t1) # pending + t2 = _make_task(project_id="p2") + repo.create(t2) + t2.mark_processing() + repo.update(t2) + t3 = _make_task(project_id="p3") + repo.create(t3) + t3.mark_processing() + repo.update(t3) + t4 = _make_task(project_id="p4") + repo.create(t4) + t4.mark_processing() + repo.update(t4) + t4.mark_completed() + repo.update(t4) + t5 = _make_task(project_id="p5", created_by_user_id="user-2") + repo.create(t5) + t5.mark_processing() + repo.update(t5) + + assert repo.count_running_by_user("user-1") == 2 + assert repo.count_running_by_user("user-2") == 1 + assert repo.count_running_total() == 3 + + +def test_count_running_total_empty(): + repo, _, _ = _repository() + assert repo.count_running_total() == 0 + assert repo.count_running_by_user("nobody") == 0 + + +def test_preview_tasks_counted_in_running(): + """预览任务(is_preview=True,工单实测卡 80% 的那种)同样计入 running。""" + repo, _, _ = _repository() + t = _make_task(is_preview=True) + repo.create(t) + t.mark_processing() + repo.update(t) + assert repo.count_running_by_user("user-1") == 1 + assert repo.count_running_total() == 1 + + +# --------------------------------------------------------------------------- +# 2. 平均耗时估算(429 等待预估依据) +# --------------------------------------------------------------------------- + + +def _complete_task(repo, engine, task, duration_seconds: float): + repo.create(task) + task.mark_processing() + repo.update(task) + task.mark_completed() + repo.update(task) + now = datetime.now(timezone.utc) + with engine.connect() as conn: + conn.execute( + text("UPDATE generation_tasks SET started_at = :s, completed_at = :c WHERE id = :id"), + {"s": now - timedelta(seconds=duration_seconds), "c": now, "id": task.id}, + ) + conn.commit() + + +def test_estimate_avg_duration_with_history(): + """有历史完成任务时返回平均耗时(秒)。""" + repo, _, engine = _repository() + _complete_task(repo, engine, _make_task(project_id="p1"), 60.0) + _complete_task(repo, engine, _make_task(project_id="p2"), 180.0) + + avg = repo.estimate_avg_duration_seconds(default_seconds=120.0) + assert 119.0 < avg < 121.0 # (60+180)/2 = 120 + + +def test_estimate_avg_duration_no_history_returns_default(): + """无历史数据时返回默认值。""" + repo, _, _ = _repository() + assert repo.estimate_avg_duration_seconds(default_seconds=90.0) == 90.0 + + +# --------------------------------------------------------------------------- +# 3. build_rate_limit_detail 结构化提示(前端区分"排队"与"创建失败") +# --------------------------------------------------------------------------- + + +def test_user_rate_limit_detail_structure(): + """429 用户限流:返回 USER_QUEUE_FULL + 排队/渲染数 + 预计等待。""" + repo, _, _ = _repository() + for i in range(2): # 2 个渲染中 + t = _make_task(project_id=f"rp{i}") + repo.create(t) + t.mark_processing() + repo.update(t) + + exc = task_enqueue.UserPendingLimitExceeded(user_id="user-1", pending_count=3, limit=3) + detail = task_enqueue.build_rate_limit_detail(exc, repo, scope="user") + + assert detail["code"] == task_enqueue.ERROR_CODE_USER_QUEUE_FULL + assert detail["queued_count"] == 3 + assert detail["running_count"] == 2 + assert detail["limit"] == 3 + assert detail["estimated_wait_seconds"] > 0 + assert "排队" in detail["message"] + assert "user-1" not in detail["message"] # 不泄露内部 ID + + +def test_global_rate_limit_detail_structure(): + """503 全局繁忙:返回 SYSTEM_QUEUE_FULL。""" + repo, _, _ = _repository() + exc = task_enqueue.GlobalQueueFull(pending_count=20, limit=20) + detail = task_enqueue.build_rate_limit_detail(exc, repo, scope="global") + + assert detail["code"] == task_enqueue.ERROR_CODE_SYSTEM_QUEUE_FULL + assert detail["queued_count"] == 20 + assert detail["limit"] == 20 + assert detail["estimated_wait_seconds"] > 0 + assert "系统繁忙" in detail["message"] + + +def test_wait_estimate_uses_concurrency(): + """等待预估:排队 8 个 / 并发 4 = 2 批 × 平均耗时。""" + + class FakeRepo: + def estimate_avg_duration_seconds(self, limit=20, default_seconds=120.0): + return 100.0 + + wait = task_enqueue._estimate_wait_seconds(8, FakeRepo()) + assert wait == 200 # ceil(8/4)=2 批 × 100 秒 + + +def test_wait_estimate_repo_without_methods_uses_default(): + """仓储没有新方法(旧 mock/鸭子类型)时用默认 120 秒兜底,不抛错。""" + + class LegacyRepo: + """只实现旧接口的仓储(模拟未升级的调用方)。""" + + def count_pending_total(self): + return 0 + + wait = task_enqueue._estimate_wait_seconds(4, LegacyRepo()) + assert wait == 120 # ceil(4/4)=1 批 × 120 默认 + + +def test_rate_limit_detail_running_count_falls_back_to_zero(): + """仓储不支持 running 计数时,running_count 优雅降级为 0。""" + + class LegacyRepo: + def count_pending_total(self): + return 0 + + exc = task_enqueue.GlobalQueueFull(pending_count=20, limit=20) + detail = task_enqueue.build_rate_limit_detail(exc, LegacyRepo(), scope="global") + assert detail["running_count"] == 0 + assert detail["code"] == task_enqueue.ERROR_CODE_SYSTEM_QUEUE_FULL + + +# --------------------------------------------------------------------------- +# 4. worker 清理核心:中断任务被重置(worker 重启/超时恢复) +# --------------------------------------------------------------------------- + + +def test_worker_cleanup_resets_interrupted_running_task(): + """模拟 worker 重启:running 超 20 分钟无更新的任务被重置为 failed,原因写明。""" + repo, _, engine = _repository() + + t = _make_task(is_preview=True) # 预览任务 + repo.create(t) + t.mark_processing() # running + repo.update(t) + _age_task(engine, t.id, updated_minutes=25) # 25 分钟无进度更新 + + cleaned = _startup.cleanup_stale_running_with_session(repo, 20) + assert cleaned == 1 + + saved = repo.get(t.id) + assert saved.status == GenerationTaskStatus.FAILED + assert "中断" in saved.error_message + assert saved.error_info.get("error_type") == "WorkerInterrupted" + assert saved.completed_at is not None + + +def test_worker_cleanup_keeps_healthy_running_task(): + """正常运行中(5 分钟前有更新)的任务不被误杀。""" + repo, _, engine = _repository() + + t = _make_task() + repo.create(t) + t.mark_processing() + repo.update(t) + _age_task(engine, t.id, updated_minutes=5) + + assert _startup.cleanup_stale_running_with_session(repo, 20) == 0 + assert repo.get(t.id).status == GenerationTaskStatus.RUNNING + + +def test_worker_cleanup_resets_stale_pending_task(): + """卡 pending 超 15 分钟(worker 停止消费)的任务被重置,释放限流名额。""" + repo, _, engine = _repository() + + t = _make_task(is_preview=True) + repo.create(t) # 一直 pending + _age_task(engine, t.id, created_minutes=20) + + cleaned = _startup.cleanup_stale_pending_with_session(repo, 15) + assert cleaned == 1 + + saved = repo.get(t.id) + assert saved.status == GenerationTaskStatus.FAILED + assert saved.error_info.get("error_type") == "PendingTimeout" + # 释放名额后 pending 计数归零,新请求不再被 429 误伤 + assert repo.count_pending_total() == 0 + + +def test_worker_cleanup_pending_keeps_recent(): + """刚创建 3 分钟的 pending 任务不清理。""" + repo, _, engine = _repository() + + t = _make_task() + repo.create(t) + _age_task(engine, t.id, created_minutes=3) + + assert _startup.cleanup_stale_pending_with_session(repo, 15) == 0 + assert repo.get(t.id).status == GenerationTaskStatus.PENDING + + +def test_worker_cleanup_multiple_orphans_all_reset(): + """3 个卡死 running 任务(工单实测:3 个预览卡 80% 超 10 小时)全部恢复。""" + repo, _, engine = _repository() + + ids = [] + for i in range(3): + t = _make_task(project_id=f"p{i}", is_preview=True) + repo.create(t) + t.mark_processing() + repo.update(t) + _age_task(engine, t.id, updated_minutes=600) # 10 小时 + ids.append(t.id) + + cleaned = _startup.cleanup_stale_running_with_session(repo, 20) + assert cleaned == 3 + for tid in ids: + assert repo.get(tid).status == GenerationTaskStatus.FAILED diff --git a/tests/unit/test_task_queue_limit.py b/tests/unit/test_task_queue_limit.py index dc4e56b41..b78baaf96 100644 --- a/tests/unit/test_task_queue_limit.py +++ b/tests/unit/test_task_queue_limit.py @@ -164,7 +164,9 @@ class TestSafeEnqueueWithLimits: result = safe_enqueue_generation_task(task, repo, user_id="user-1") assert result is True mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"]) - assert len(repo.updated_tasks) == 0 # 成功不需要更新状态 + # 成功入队后持久化 celery 消息 ID(#1714:孤儿清理据此 revoke/清队列) + assert len(repo.updated_tasks) == 1 + assert task.celery_task_id def test_user_limit_rejected_with_failed_status(self, mock_celery): """用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。""" @@ -335,7 +337,9 @@ class TestPostEnqueueFinalCheck: assert result is True mock_celery.assert_called_once() assert task.status == "pending" # 状态没变 - assert len(repo.updated_tasks) == 0 # 没更新 DB + # 入队成功后持久化 celery_task_id(#1714),业务状态不变 + assert len(repo.updated_tasks) == 1 + assert task.celery_task_id def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery): """不传 user_id 时,入队后校验也跳过用户级,只查全局。""" diff --git a/tests/unit/test_upload_complete_idempotency_1714.py b/tests/unit/test_upload_complete_idempotency_1714.py new file mode 100644 index 000000000..c0436b065 --- /dev/null +++ b/tests/unit/test_upload_complete_idempotency_1714.py @@ -0,0 +1,422 @@ +"""Issue #1714:POST /upload/direct/complete 幂等 + multipart 幂等。 + +覆盖: +- 同 client_upload_id 重复 complete → 只建一条 asset、不重复派 ingest job +- 同 file_hash 重复 complete → 返回已存在记录 +- 旧客户端不传 hash/token:近期同库同名 processing 占位 → 兜底幂等返回 +- 旧客户端不传 hash/token:READY 历史同名 → 不兜底(正常新建) +- 兜底窗口外(>30 分钟)→ 不兜底 +- 旧仓储(无新方法)鸭子类型降级 → 不报错、正常新建 +- 重复 complete 时即使 OSS 已无文件(file_exists=False)也返回已存在记录 + (模拟 complete 超时后 OSS 侧对象已过期/清理,重试仍不重复建库) +- multipart 上传:同 client_upload_id 重复提交 → 第二次直接 duplicated,不再传 OSS +""" + +from __future__ import annotations + +import os +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path +from unittest.mock import MagicMock + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from fastapi import FastAPI # noqa: E402 +from fastapi.testclient import TestClient # noqa: E402 + +from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, IngestJob, Project # noqa: E402 + + +class StubProjectRepository: + def __init__(self, projects: dict | None = None): + self._projects = projects or {} + + def get(self, project_id: str): + return self._projects.get(project_id) + + def find_by_id(self, project_id: str): + return self._projects.get(project_id) + + +class StubAssetLibraryRepository: + def __init__(self, libraries: dict | None = None): + self._libraries = libraries or {} + + def find_by_project(self, project_id: str, kind=None) -> list: + return list(self._libraries.values()) + + +class StubAssetRepository: + """支持三种幂等查询的内存仓储,并统计 create 次数。""" + + def __init__(self, assets: list[Asset] | None = None): + self._assets = list(assets or []) + self.created: list[Asset] = [] + + def find_by_library_and_file_hash(self, library_id: str, file_hash: str) -> Asset | None: + if not file_hash: + return None + return next((a for a in self._assets if a.library_id == library_id and a.file_hash == file_hash), None) + + def find_by_library_and_client_upload_id(self, library_id: str, client_upload_id: str) -> Asset | None: + if not client_upload_id: + return None + return next( + (a for a in self._assets if a.library_id == library_id and a.client_upload_id == client_upload_id), + None, + ) + + def find_recent_active_by_library_and_name( + self, library_id: str, name: str, within_minutes: int = 30, file_size: int = 0 + ) -> Asset | None: + # 严格模式(#1714):大小未知(0)直接不命中,宁可漏判不可误杀 + if not file_size or file_size <= 0: + return None + cutoff = datetime.now(timezone.utc) - timedelta(minutes=within_minutes) + candidates = [ + a + for a in self._assets + if a.library_id == library_id + and a.name == name + and a.status in (AssetStatus.UPLOADING, AssetStatus.PROCESSING) + and a.created_at >= cutoff + and a.file_size == file_size + ] + return max(candidates, key=lambda a: a.created_at) if candidates else None + + def create(self, asset: Asset) -> Asset: + self._assets.append(asset) + self.created.append(asset) + return asset + + def update(self, asset: Asset) -> Asset: + return asset + + +class LegacyStubAssetRepository: + """旧仓储:只有 file_hash 去重,没有新方法(鸭子类型降级验证)。""" + + def __init__(self, assets: list[Asset] | None = None): + self._assets = list(assets or []) + self.created: list[Asset] = [] + + def find_by_library_and_file_hash(self, library_id: str, file_hash: str) -> Asset | None: + if not file_hash: + return None + return next((a for a in self._assets if a.library_id == library_id and a.file_hash == file_hash), None) + + def create(self, asset: Asset) -> Asset: + self._assets.append(asset) + self.created.append(asset) + return asset + + +class StubIngestJobRepository: + def __init__(self): + self._jobs: dict[str, IngestJob] = {} + self.created_count = 0 + + def create(self, job: IngestJob) -> IngestJob: + self._jobs[job.id] = job + self.created_count += 1 + return job + + def get(self, job_id: str) -> IngestJob | None: + return self._jobs.get(job_id) + + def update(self, job: IngestJob) -> IngestJob: + self._jobs[job.id] = job + return job + + +def _make_project() -> Project: + return Project(id="proj-1", name="Test Project", owner_user_id="user-1") + + +def _make_library() -> AssetLibrary: + return AssetLibrary(id="lib-1", name="Test Library", project_id="proj-1", kind=AssetLibraryKind.VIDEO) + + +def _build_app(asset_repo=None, ingest_repo=None, storage=None): + from app.api.routes.upload import router + from app.auth import AuthenticatedUser, get_current_user + from app.core.storage import get_storage_service + from app.dependencies import ( + get_asset_library_repository, + get_asset_repository, + get_ingest_job_repository, + get_project_repository, + ) + + app = FastAPI() + app.include_router(router, prefix="/api/v1") + + project_repo = StubProjectRepository({"proj-1": _make_project()}) + library_repo = StubAssetLibraryRepository({"lib-1": _make_library()}) + asset_repo = asset_repo or StubAssetRepository() + ingest_repo = ingest_repo or StubIngestJobRepository() + + storage = storage or MagicMock() + storage.is_configured = True + storage._normalize_storage_key = lambda key: key + storage.file_exists = MagicMock(return_value=True) + storage.upload_file = MagicMock(return_value="https://oss.example.com/file.mp4") + storage.get_url = MagicMock(return_value="https://oss.example.com/file.mp4") + + mock_user = MagicMock(spec=AuthenticatedUser) + mock_user.id = "user-1" + mock_user.user = MagicMock(id="user-1") + mock_user.email = "test@example.com" + + app.dependency_overrides[get_current_user] = lambda: mock_user + app.dependency_overrides[get_project_repository] = lambda: project_repo + app.dependency_overrides[get_asset_library_repository] = lambda: library_repo + app.dependency_overrides[get_asset_repository] = lambda: asset_repo + app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo + app.dependency_overrides[get_storage_service] = lambda: storage + return app, asset_repo, ingest_repo, storage + + +def _client(**kwargs): + app, asset_repo, ingest_repo, storage = _build_app(**kwargs) + return TestClient(app), asset_repo, ingest_repo, storage + + +COMPLETE_BODY = { + "project_id": "proj-1", + "library_id": "lib-1", + "storage_key": "uploads/abc/IMG_2282.MOV", +} + + +class TestDirectCompleteIdempotency: + def test_same_client_upload_id_creates_single_asset_and_job(self): + """同一 client_upload_id 连发两次 complete:只建 1 条 asset、1 个 job。""" + client, asset_repo, ingest_repo, _ = _client() + body = {**COMPLETE_BODY, "client_upload_id": "up-token-1", "file_size": 12345} + + r1 = client.post("/api/v1/direct/complete", json=body) + r2 = client.post("/api/v1/direct/complete", json={**body, "storage_key": "uploads/zzz/IMG_2282.MOV"}) + + assert r1.status_code == 200 and r2.status_code == 200 + b1, b2 = r1.json(), r2.json() + assert b1["duplicated"] is False + assert b2["duplicated"] is True + assert b1["asset_id"] == b2["asset_id"] + assert len(asset_repo.created) == 1 + assert ingest_repo.created_count == 1 + # 第二次返回的是已存在记录(其 storage_key 为第一次的 key) + assert b2["storage_key"] == "uploads/abc/IMG_2282.MOV" + + def test_same_file_hash_returns_existing(self): + """同 file_hash(不同 token)重复 complete → 返回已存在记录。""" + client, asset_repo, ingest_repo, _ = _client() + body1 = {**COMPLETE_BODY, "file_hash": "h" * 32, "client_upload_id": "tok-a"} + body2 = { + **COMPLETE_BODY, + "storage_key": "uploads/def/IMG_2282.MOV", + "file_hash": "h" * 32, + "client_upload_id": "tok-b", + } + + client.post("/api/v1/direct/complete", json=body1) + r2 = client.post("/api/v1/direct/complete", json=body2) + + assert r2.json()["duplicated"] is True + assert len(asset_repo.created) == 1 + assert ingest_repo.created_count == 1 + + def test_fallback_dedup_when_no_hash_no_token(self): + """旧客户端不传 hash/token:近期同库同名 processing 占位 → 兜底幂等。 + + 模拟 complete 超时重试:第一次已建好占位,第二次(OSS 重传拿到新 key) + 不应再建第二条。 + """ + client, asset_repo, ingest_repo, _ = _client() + # 第一次 complete(旧客户端无 token/hash,但 file_size 可知) + r1 = client.post( + "/api/v1/direct/complete", + json={**COMPLETE_BODY, "file_size": 5_000_000}, + ) + assert r1.json()["duplicated"] is False + # 重试:重新 prepare 产生新 storage_key(仅 uuid 目录不同,文件名一致—— + # 前端重试传的是同一个 File),且近期;同大小才允许兜底命中 + r2 = client.post( + "/api/v1/direct/complete", + json={ + **COMPLETE_BODY, + "storage_key": "uploads/retry/IMG_2282.MOV", + "file_size": 5_000_000, + }, + ) + assert r2.status_code == 200 + assert r2.json()["duplicated"] is True + assert r2.json()["asset_id"] == r1.json()["asset_id"] + assert len(asset_repo.created) == 1 + assert ingest_repo.created_count == 1 + + def test_fallback_dedup_skipped_when_file_size_unknown(self): + """file_size=0(未知)时不允许仅凭同名 + processing 判重,直接放行(#1714)。 + + 根因场景:complete 没传 file_size,30 分钟内同名占位(如 iPhone 的 + IMG_2285.MOV)会把内容/大小全新的视频误判为重复跳过。 + """ + client, asset_repo, _ingest_repo, _ = _client() + r1 = client.post( + "/api/v1/direct/complete", + json={**COMPLETE_BODY, "file_size": 0}, + ) + assert r1.json()["duplicated"] is False + # 第二个全新视频:同名(IMG_2285.MOV)、无 hash/token、file_size 仍未知 + r2 = client.post( + "/api/v1/direct/complete", + json={**COMPLETE_BODY, "storage_key": "uploads/retry2/IMG_2282.MOV", "file_size": 0}, + ) + assert r2.status_code == 200 + assert r2.json()["duplicated"] is False # 不能误杀 + assert len(asset_repo.created) == 2 # 两条记录,放行新上传 + + def test_fallback_dedup_skipped_when_same_name_but_different_size(self): + """同名但 file_size 不同 → 不判重,正常建记录(#1714)。""" + client, asset_repo, _ingest_repo, _ = _client() + r1 = client.post( + "/api/v1/direct/complete", + json={**COMPLETE_BODY, "file_size": 5_000_000}, + ) + assert r1.json()["duplicated"] is False + r2 = client.post( + "/api/v1/direct/complete", + json={ + **COMPLETE_BODY, + "storage_key": "uploads/retry3/IMG_2282.MOV", + "file_size": 9_999_999, # 同名但大小完全不同的新视频 + }, + ) + assert r2.status_code == 200 + assert r2.json()["duplicated"] is False + assert len(asset_repo.created) == 2 + + def test_fallback_dedup_skipped_when_hash_present_even_if_name_size_match(self): + """file_hash 非空且 hash 未命中时,不允许退回同名兜底(#1714)。 + + hash 已能代表内容:同名同大小但 hash 不同是真实的新内容,必须放行。 + """ + client, asset_repo, _ingest_repo, _ = _client() + # 第一次:某 hash 的视频 + r1 = client.post( + "/api/v1/direct/complete", + json={ + **COMPLETE_BODY, + "file_hash": "a" * 64, + "client_upload_id": "tok-1", + "file_size": 5_000_000, + }, + ) + assert r1.json()["duplicated"] is False + # 第二次:同名同大小但 hash 不同(新视频内容不同); + # 注意 client_upload_id 也必须不同,否则会先被 token 命中 + r2 = client.post( + "/api/v1/direct/complete", + json={ + **COMPLETE_BODY, + "storage_key": "uploads/retry4/IMG_2282.MOV", + "file_hash": "b" * 64, + "client_upload_id": "tok-2", + "file_size": 5_000_000, + }, + ) + assert r2.status_code == 200 + assert r2.json()["duplicated"] is False + assert len(asset_repo.created) == 2 + + def test_fallback_dedup_ignores_ready_history(self): + """READY 历史同名素材不触发兜底(允许用户再次上传同名文件)。""" + ready = Asset( + id="ready-1", + project_id="proj-1", + library_id="lib-1", + name="IMG_2282.MOV", + storage_key="uploads/old/IMG_2282.MOV", + mime_type="video/quicktime", + status=AssetStatus.READY, + ) + client, asset_repo, ingest_repo, _ = _client(asset_repo=StubAssetRepository([ready])) + r = client.post("/api/v1/direct/complete", json=COMPLETE_BODY) + assert r.status_code == 200 + assert r.json()["duplicated"] is False + assert len(asset_repo.created) == 1 + + def test_fallback_dedup_window_expired(self): + """占位记录超过 30 分钟 → 不再兜底(视为孤儿,正常新建)。""" + stale = Asset( + id="stale-1", + project_id="proj-1", + library_id="lib-1", + name="IMG_2282.MOV", + storage_key="uploads/stale/IMG_2282.MOV", + mime_type="video/quicktime", + status=AssetStatus.PROCESSING, + ) + stale.created_at = datetime.now(timezone.utc) - timedelta(minutes=45) + client, asset_repo, ingest_repo, _ = _client(asset_repo=StubAssetRepository([stale])) + r = client.post("/api/v1/direct/complete", json=COMPLETE_BODY) + assert r.status_code == 200 + assert r.json()["duplicated"] is False + assert len(asset_repo.created) == 1 + + def test_legacy_repo_without_new_methods_still_works(self): + """旧仓储没有新幂等方法 → 鸭子类型降级,不报错、正常创建。""" + client, asset_repo, ingest_repo, _ = _client(asset_repo=LegacyStubAssetRepository()) + r = client.post( + "/api/v1/direct/complete", + json={**COMPLETE_BODY, "client_upload_id": "tok-x", "file_hash": "f" * 32}, + ) + assert r.status_code == 200 + assert r.json()["duplicated"] is False + assert len(asset_repo.created) == 1 + + def test_duplicate_complete_returns_existing_even_if_oss_missing(self): + """重复 complete 幂等检查先于 OSS file_exists: + + 第一次成功建占位后,重试时即使 OSS 对象已不存在(file_exists=False), + 也必须返回已存在记录而不是 404/重复建库。""" + client, _, _, storage = _client() + body = {**COMPLETE_BODY, "client_upload_id": "tok-oss-gone"} + r1 = client.post("/api/v1/direct/complete", json=body) + assert r1.status_code == 200 + + storage.file_exists = MagicMock(return_value=False) + r2 = client.post( + "/api/v1/direct/complete", + json={**body, "storage_key": "uploads/retry2/IMG_2282.MOV"}, + ) + assert r2.status_code == 200 + assert r2.json()["duplicated"] is True + assert r2.json()["asset_id"] == r1.json()["asset_id"] + + +class TestMultipartUploadIdempotency: + def test_same_client_upload_id_second_submit_deduplicated(self): + """multipart 重复提交同 token:第二次直接 duplicated,不再上传 OSS。""" + client, asset_repo, ingest_repo, storage = _client() + + def _post(): + return client.post( + "/api/v1", + data={"project_id": "proj-1", "library_id": "lib-1", "client_upload_id": "mp-tok-1"}, + files={"file": ("IMG_2282.MOV", b"fake-mov-data", "video/quicktime")}, + ) + + r1 = _post() + r2 = _post() + assert r1.json()["duplicated"] is False + assert r2.json()["duplicated"] is True + assert r2.json()["asset_id"] == r1.json()["asset_id"] + assert len(asset_repo.created) == 1 + assert ingest_repo.created_count == 1 + # OSS 上传只发生一次(第二次在幂等检查处直接返回) + assert storage.upload_file.call_count == 1 diff --git a/tests/unit/test_wechat_bind_routes_1719.py b/tests/unit/test_wechat_bind_routes_1719.py new file mode 100644 index 000000000..9f07cc43d --- /dev/null +++ b/tests/unit/test_wechat_bind_routes_1719.py @@ -0,0 +1,235 @@ +"""#1719:微信绑定/解绑路由层测试(直接驱动路由函数)。 + +覆盖: +- GET /wechat/bind/url:调 oauth 生成链接、记日志 +- POST /wechat/bind:oauth 失败→400;绑定成功→success+user.wechat_bound=True; + use case 返回冲突→对应状态码透传 +- DELETE /wechat/bind:成功→success=True;use case 报错→状态码透传 +- /auth/me 返回 wechat_bound 字段 +""" + +from __future__ import annotations + +import asyncio +import os +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from app.api.routes import auth as auth_route # noqa: E402 +from fastapi import HTTPException # noqa: E402 + + +def _auth_user(user_id="u-1", openid=None): + user = SimpleNamespace( + id=user_id, + wechat_openid=openid, + email="user@example.com", + email_verified=True, + username="user", + display_name="用户", + phone="", + phone_verified=False, + profile_completed=True, + ) + return SimpleNamespace(user=user, session_id="s-1", token_type="user_auth") + + +def _patched_bind(result, error, status): + """构造打了补丁的 wechat_bind_use_case 模块""" + mod = SimpleNamespace( + WechatBindRequest=lambda **kw: SimpleNamespace(**kw), + WechatBindUseCase=MagicMock(), + WechatUnbindUseCase=MagicMock(), + ) + fake_bind_uc = MagicMock() + fake_bind_uc.bind.return_value = (result, error, status) + mod.WechatBindUseCase.return_value = fake_bind_uc + return mod + + +def test_get_bind_url_returns_url_and_state(): + fake_oauth = MagicMock() + fake_oauth.generate_auth_url.return_value = ("https://open.weixin.qq.com/qrconnect?xxx", "state-bind-1") + + import packages.application.auth.wechat_oauth_service as oauth_mod + + orig = oauth_mod.get_wechat_oauth_service + oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth) + try: + resp = asyncio.run(auth_route.get_wechat_bind_url(current_user=_auth_user())) + finally: + oauth_mod.get_wechat_oauth_service = orig + + assert resp.auth_url.startswith("https://open.weixin.qq.com") + assert resp.state == "state-bind-1" + + +def test_bind_oauth_error_returns_400(): + fake_oauth = MagicMock() + fake_oauth.handle_callback.return_value = (None, "无效的 state 参数") + + import packages.application.auth.wechat_oauth_service as oauth_mod + + orig = oauth_mod.get_wechat_oauth_service + oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth) + try: + with pytest.raises(HTTPException) as exc: + asyncio.run( + auth_route.wechat_bind( + SimpleNamespace(code="c-1", state="s-1"), + current_user=_auth_user(), + user_repository=MagicMock(), + ) + ) + finally: + oauth_mod.get_wechat_oauth_service = orig + + assert exc.value.status_code == 400 + assert "state" in exc.value.detail + + +def test_bind_success_returns_user_with_wechat_bound(): + fake_oauth = MagicMock() + fake_oauth.handle_callback.return_value = ( + SimpleNamespace(openid="wx-openid-1", unionid="wx-union-1"), + None, + ) + + bound_user = SimpleNamespace( + id="u-1", + wechat_openid="wx-openid-1", + email="user@example.com", + email_verified=True, + username="user", + display_name="用户", + phone="", + phone_verified=False, + profile_completed=True, + ) + + import packages.application.auth.wechat_oauth_service as oauth_mod + from packages.application.auth import wechat_bind_use_case as bind_mod + + orig_oauth = oauth_mod.get_wechat_oauth_service + fake_bind_uc = MagicMock() + fake_bind_uc.bind.return_value = (SimpleNamespace(user=bound_user), None, 200) + orig_bind = bind_mod.WechatBindUseCase + bind_mod.WechatBindUseCase = MagicMock(return_value=fake_bind_uc) + oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth) + try: + resp = asyncio.run( + auth_route.wechat_bind( + SimpleNamespace(code="c-1", state="s-1"), + current_user=_auth_user(), + user_repository=MagicMock(), + ) + ) + finally: + oauth_mod.get_wechat_oauth_service = orig_oauth + bind_mod.WechatBindUseCase = orig_bind + + assert resp.success is True + assert resp.user.wechat_bound is True + assert resp.user.user_id == "u-1" + # 绑定请求应带上当前用户 id 与微信 openid + call_kwargs = fake_bind_uc.bind.call_args[0][0] + assert call_kwargs.user_id == "u-1" + assert call_kwargs.openid == "wx-openid-1" + + +def test_bind_conflict_propagates_409(): + fake_oauth = MagicMock() + fake_oauth.handle_callback.return_value = ( + SimpleNamespace(openid="wx-openid-1", unionid=""), + None, + ) + + import packages.application.auth.wechat_oauth_service as oauth_mod + from packages.application.auth import wechat_bind_use_case as bind_mod + + orig_oauth = oauth_mod.get_wechat_oauth_service + fake_bind_uc = MagicMock() + fake_bind_uc.bind.return_value = (None, "该微信已绑定其他账号,请先在原账号解绑", 409) + orig_bind = bind_mod.WechatBindUseCase + bind_mod.WechatBindUseCase = MagicMock(return_value=fake_bind_uc) + oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth) + try: + with pytest.raises(HTTPException) as exc: + asyncio.run( + auth_route.wechat_bind( + SimpleNamespace(code="c-1", state="s-1"), + current_user=_auth_user(), + user_repository=MagicMock(), + ) + ) + finally: + oauth_mod.get_wechat_oauth_service = orig_oauth + bind_mod.WechatBindUseCase = orig_bind + + assert exc.value.status_code == 409 + assert "已绑定其他账号" in exc.value.detail + + +def test_unbind_success_returns_success_true(): + unbound_user = SimpleNamespace( + id="u-1", + wechat_openid=None, + email="user@example.com", + email_verified=True, + username="user", + display_name="用户", + phone="", + phone_verified=False, + ) + + from packages.application.auth import wechat_bind_use_case as bind_mod + + fake_uc = MagicMock() + fake_uc.unbind.return_value = (SimpleNamespace(user=unbound_user), None, 200) + orig = bind_mod.WechatUnbindUseCase + bind_mod.WechatUnbindUseCase = MagicMock(return_value=fake_uc) + try: + resp = asyncio.run( + auth_route.wechat_unbind(current_user=_auth_user(openid="wx-old"), user_repository=MagicMock()) + ) + finally: + bind_mod.WechatUnbindUseCase = orig + + assert resp.success is True + fake_uc.unbind.assert_called_once_with("u-1") + + +def test_unbind_rejected_no_other_login_propagates_400(): + from packages.application.auth import wechat_bind_use_case as bind_mod + + fake_uc = MagicMock() + fake_uc.unbind.return_value = (None, "账号需要至少一种其他登录方式(已验证手机或真实邮箱)后才能解绑微信", 400) + orig = bind_mod.WechatUnbindUseCase + bind_mod.WechatUnbindUseCase = MagicMock(return_value=fake_uc) + try: + with pytest.raises(HTTPException) as exc: + asyncio.run(auth_route.wechat_unbind(current_user=_auth_user(openid="wx-old"), user_repository=MagicMock())) + finally: + bind_mod.WechatUnbindUseCase = orig + + assert exc.value.status_code == 400 + assert "登录方式" in exc.value.detail + + +def test_me_includes_wechat_bound_flag(): + # 已绑定用户 + resp = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(openid="wx-openid-1"))) + assert resp.wechat_bound is True + + # 未绑定用户 + resp2 = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(openid=None))) + assert resp2.wechat_bound is False diff --git a/tests/unit/test_wechat_bind_use_case_1719.py b/tests/unit/test_wechat_bind_use_case_1719.py new file mode 100644 index 000000000..e933f8a25 --- /dev/null +++ b/tests/unit/test_wechat_bind_use_case_1719.py @@ -0,0 +1,251 @@ +"""#1719:已登录用户微信绑定/解绑 Use Case 测试。 + +覆盖: +- bind:幂等重复绑定、未绑定成功、当前账号已绑其他微信、openid/unionid 冲突 409、用户不存在 +- unbind:成功清 openid+unionid、未绑定拒绝、无其他登录方式拒绝、密码/手机/真实邮箱各兜底放行、用户不存在 +""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +from packages.application.auth.wechat_bind_use_case import ( + WechatBindRequest, + WechatBindUseCase, + WechatUnbindUseCase, +) + + +def _user( + user_id="u-1", + wechat_openid=None, + wechat_unionid=None, + password_hash="hashed-pw", + phone=None, + phone_verified=False, + email="user@example.com", + email_verified=True, +): + return SimpleNamespace( + id=user_id, + wechat_openid=wechat_openid, + wechat_unionid=wechat_unionid, + password_hash=password_hash, + phone=phone, + phone_verified=phone_verified, + email=email, + email_verified=email_verified, + ) + + +class _FakeRepo: + """内存仓储:按 id/openid/unionid 建索引,save 原地更新。""" + + def __init__(self, users): + self.users = {u.id: u for u in users} + self.saved = [] + + def find_by_id(self, user_id): + return self.users.get(user_id) + + def find_by_wechat_openid(self, openid): + for u in self.users.values(): + if u.wechat_openid == openid: + return u + return None + + def find_by_wechat_unionid(self, unionid): + if not unionid: + return None + for u in self.users.values(): + if u.wechat_unionid == unionid: + return u + return None + + def save(self, user): + self.saved.append(user) + + +# ==================== bind ==================== + + +def test_bind_success_when_not_bound(): + user = _user() + repo = _FakeRepo([user]) + result, err, status = WechatBindUseCase(repo).bind( + WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-1") + ) + assert err is None + assert status == 200 + assert result.user.wechat_openid == "wx-openid-1" + assert result.user.wechat_unionid == "wx-union-1" + assert repo.saved == [user] + + +def test_bind_idempotent_same_openid(): + user = _user(wechat_openid="wx-openid-1", wechat_unionid="wx-union-1") + repo = _FakeRepo([user]) + result, err, status = WechatBindUseCase(repo).bind( + WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-1") + ) + assert err is None + assert status == 200 + assert result.user is user + assert repo.saved == [] # 幂等不写库 + + +def test_bind_conflict_user_already_bound_other_wechat(): + user = _user(wechat_openid="wx-old") + repo = _FakeRepo([user]) + result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="u-1", openid="wx-new")) + assert result is None + assert status == 409 + assert "已绑定微信" in err + + +def test_bind_conflict_openid_used_by_other_user(): + user = _user(user_id="u-1") + other = _user(user_id="u-2", wechat_openid="wx-openid-1") + repo = _FakeRepo([user, other]) + result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="u-1", openid="wx-openid-1")) + assert result is None + assert status == 409 + assert "已绑定其他账号" in err + assert user.wechat_openid is None # 未写库 + + +def test_bind_conflict_unionid_used_by_other_user(): + user = _user(user_id="u-1") + # openid 不同,但 unionid 指向同一微信主体 + other = _user(user_id="u-2", wechat_openid="wx-other", wechat_unionid="wx-union-x") + repo = _FakeRepo([user, other]) + result, err, status = WechatBindUseCase(repo).bind( + WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-x") + ) + assert result is None + assert status == 409 + assert "微信主体" in err + + +def test_bind_missing_openid_returns_400(): + repo = _FakeRepo([_user()]) + result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="u-1", openid="")) + assert result is None + assert status == 400 + assert "openid" in err + + +def test_bind_user_not_found_returns_404(): + repo = _FakeRepo([]) + result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="ghost", openid="wx-openid-1")) + assert result is None + assert status == 404 + + +def test_bind_fills_unionid_when_existing_user_has_none(): + # 用户历史上只绑了 openid(unionid 为空),再次绑定时补齐 unionid 不冲突 + user = _user(wechat_openid="wx-openid-1", wechat_unionid=None) + repo = _FakeRepo([user]) + result, err, status = WechatBindUseCase(repo).bind( + WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-new") + ) + # openid 相同 → 幂等成功(不覆盖 unionid,保持数据稳定) + assert err is None + assert status == 200 + + +# ==================== unbind ==================== + + +def test_unbind_success_with_real_verified_email(): + # 默认 _user 即 real@example.com 且 email_verified=True + user = _user(wechat_openid="wx-openid-1", wechat_unionid="wx-union-1") + repo = _FakeRepo([user]) + result, err, status = WechatUnbindUseCase(repo).unbind("u-1") + assert err is None + assert status == 200 + assert result.user.wechat_openid is None + assert result.user.wechat_unionid is None + assert repo.saved == [user] + + +def test_unbind_rejected_when_only_random_password_hash(): + # 微信注册用户:随机密码 hash 存在、邮箱是 @wechat.local 占位、无手机 → 不允许解绑 + user = _user( + wechat_openid="wx-openid-1", + password_hash="random-secret-hash", + email="abc@wechat.local", + email_verified=True, + ) + repo = _FakeRepo([user]) + result, err, status = WechatUnbindUseCase(repo).unbind("u-1") + assert result is None + assert status == 400 + assert "登录方式" in err + assert user.wechat_openid == "wx-openid-1" # 未写库 + + +def test_unbind_allowed_with_verified_phone_even_without_password(): + user = _user( + wechat_openid="wx-openid-1", + password_hash="", + phone="13800000000", + phone_verified=True, + email="wx@wechat.local", + email_verified=True, + ) + repo = _FakeRepo([user]) + result, err, status = WechatUnbindUseCase(repo).unbind("u-1") + assert err is None + assert status == 200 + assert result.user.wechat_openid is None + + +def test_unbind_rejected_when_no_other_login_method(): + # 无手机、邮箱占位 → 唯一登录方式就是微信,禁止解绑 + user = _user( + wechat_openid="wx-openid-1", + password_hash="", + email="abc@wechat.local", + email_verified=True, + ) + repo = _FakeRepo([user]) + result, err, status = WechatUnbindUseCase(repo).unbind("u-1") + assert result is None + assert status == 400 + assert "登录方式" in err + assert user.wechat_openid == "wx-openid-1" # 未写库 + + +def test_unbind_not_bound_returns_400(): + user = _user() # 未绑定 + repo = _FakeRepo([user]) + result, err, status = WechatUnbindUseCase(repo).unbind("u-1") + assert result is None + assert status == 400 + assert "未绑定" in err + + +def test_unbind_user_not_found_returns_404(): + repo = _FakeRepo([]) + result, err, status = WechatUnbindUseCase(repo).unbind("ghost") + assert result is None + assert status == 404 + + +def test_unbind_unverified_phone_does_not_count(): + # 手机未验证不算有效登录方式 + user = _user( + wechat_openid="wx-openid-1", + password_hash="", + phone="13800000000", + phone_verified=False, + email="abc@wechat.local", + email_verified=True, + ) + repo = _FakeRepo([user]) + result, err, status = WechatUnbindUseCase(repo).unbind("u-1") + assert result is None + assert status == 400 diff --git a/tests/unit/test_wechat_callback_logging_1718.py b/tests/unit/test_wechat_callback_logging_1718.py new file mode 100644 index 000000000..4f45ea735 --- /dev/null +++ b/tests/unit/test_wechat_callback_logging_1718.py @@ -0,0 +1,115 @@ +"""#1718:微信回调路由可观测性日志分支覆盖(UA/state/错误透传)。 + +直接驱动 wechat_callback 路由函数,mock OAuth service 与用户仓储: +- 成功路径:日志记录 UA、state 校验通过(MicroMessenger 内置浏览器) +- 失败路径:OAuth 返回错误时记 warning 并抛 400 +""" + +from __future__ import annotations + +import asyncio +import os +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from app.api.routes import auth as auth_route # noqa: E402 +from fastapi import HTTPException # noqa: E402 + + +class _FakeRequest: + def __init__(self, ua: str): + self.headers = {"User-Agent": ua} + + +def _wechat_user(): + return SimpleNamespace( + openid="openid-callback-1", + unionid="union-callback-1", + nickname="微信用户", + avatar_url="http://x/a.png", + ) + + +def _fake_oauth_factory(success: bool): + service = MagicMock() + if success: + service.handle_callback.return_value = (_wechat_user(), None) + else: + service.handle_callback.return_value = (None, "无效的 state 参数,请求可能已过期或被篡改") + return service + + +def test_wechat_callback_success_logs_ua_and_state(caplog): + fake_repo = MagicMock() + sync_response = SimpleNamespace( + access_token="at", + refresh_token="rt", + user_id="u-1", + nickname="微信用户", + avatar_url="", + is_new_user=False, + expires_in=1800, + ) + fake_use_case = MagicMock() + fake_use_case.execute.return_value = (sync_response, None) + + user = SimpleNamespace( + id="u-1", + phone_verified=True, + email_verified=True, + email="u@example.com", + ) + fake_repo.find_by_id.return_value = user + + request_obj = SimpleNamespace(code="code-1", state="state-1") + fake_http = _FakeRequest("Mozilla/5.0 (Linux; Android 13) MicroMessenger/8.0.40 WeChat/8.0.40") + + import packages.application.auth.wechat_oauth_service as oauth_mod + import packages.application.auth.wechat_sync_use_case as sync_mod + + orig_oauth = oauth_mod.get_wechat_oauth_service + orig_sync = sync_mod.WechatSyncUseCase + oauth_mod.get_wechat_oauth_service = MagicMock(return_value=_fake_oauth_factory(success=True)) + sync_mod.WechatSyncUseCase = MagicMock(return_value=fake_use_case) + try: + with caplog.at_level("INFO", logger="app.api.routes.auth"): + resp = asyncio.run(auth_route.wechat_callback(request_obj, fake_http, user_repository=fake_repo)) + finally: + oauth_mod.get_wechat_oauth_service = orig_oauth + sync_mod.WechatSyncUseCase = orig_sync + + assert resp.user_id == "u-1" + assert resp.binding_complete is True + log_text = " ".join(rec.getMessage() for rec in caplog.records) + assert "微信回调" in log_text + assert "MicroMessenger" in log_text or "微信内置浏览器=True" in log_text + + +def test_wechat_callback_failure_raises_400_with_detail(caplog): + request_obj = SimpleNamespace(code="code-bad", state="state-bad") + fake_http = _FakeRequest("Mozilla/5.0 Chrome/127") + fake_service = _fake_oauth_factory(success=False) + + import packages.application.auth.wechat_oauth_service as oauth_mod + + orig = oauth_mod.get_wechat_oauth_service + oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_service) + try: + with caplog.at_level("WARNING", logger="app.api.routes.auth"): + with pytest.raises(HTTPException) as exc_info: + asyncio.run(auth_route.wechat_callback(request_obj, fake_http, user_repository=MagicMock())) + finally: + oauth_mod.get_wechat_oauth_service = orig + + assert exc_info.value.status_code == 400 + assert "state" in exc_info.value.detail + assert any("微信回调" in rec.getMessage() for rec in caplog.records) diff --git a/tests/unit/test_wechat_oauth_service.py b/tests/unit/test_wechat_oauth_service.py index 7a2ed3397..759732c62 100755 --- a/tests/unit/test_wechat_oauth_service.py +++ b/tests/unit/test_wechat_oauth_service.py @@ -388,3 +388,52 @@ class TestGetWechatOAuthService: """返回 WechatOAuthService 实例""" service = get_wechat_oauth_service() assert isinstance(service, WechatOAuthService) + + def test_singleton_same_instance_across_calls(self, monkeypatch): + """#1718 回归:工厂必须返回同一实例,否则 state store 不共享""" + import packages.application.auth.wechat_oauth_service as mod + + monkeypatch.setattr(mod, "_oauth_service_singleton", None) + s1 = get_wechat_oauth_service() + s2 = get_wechat_oauth_service() + assert s1 is s2 + + def test_state_survives_across_factory_calls(self, monkeypatch): + """#1718 回归:/wechat/url 与 /wechat/callback 经工厂拿到同一 state store + + 模拟两次请求各自调用工厂:第一个实例生成 state,第二个实例(同一单例) + 必须能校验通过。修复前工厂每次 new 一个实例,回调必现 400「无效的 state」。 + """ + import packages.application.auth.wechat_oauth_service as mod + + monkeypatch.setattr(mod, "_oauth_service_singleton", None) + monkeypatch.setenv("WECHAT_OPEN_APP_ID", "wx-test") + monkeypatch.setenv("WECHAT_OPEN_APP_SECRET", "secret-test") + monkeypatch.setenv("WECHAT_OPEN_REDIRECT_URI", "https://example.com/cb") + + # 请求1:生成授权链接(state 写入单例 store) + _, state = get_wechat_oauth_service().generate_auth_url() + + # 请求2:回调校验(应命中同一个 store;微信 API 用 mock 避免外网) + with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: + mock_get.return_value = MagicMock( + json=MagicMock( + return_value={ + "access_token": "at", + "openid": "oid", + "unionid": "uid", + "nickname": "n", + "headimgurl": "http://x/a.png", + } + ) + ) + user_info, error = get_wechat_oauth_service().handle_callback("code-x", state) + + assert error is None, f"state 应跨请求共享,实际报错: {error}" + assert user_info is not None + assert user_info.openid == "oid" + + # state 一次性消费,重放必须失败 + user_info2, error2 = get_wechat_oauth_service().handle_callback("code-y", state) + assert user_info2 is None + assert "state" in error2 diff --git a/tests/unit/test_wechat_state_redis_1718.py b/tests/unit/test_wechat_state_redis_1718.py new file mode 100644 index 000000000..3ab463a28 --- /dev/null +++ b/tests/unit/test_wechat_state_redis_1718.py @@ -0,0 +1,262 @@ +"""#1718:微信 OAuth state 存储 Redis 化 + 中文昵称 UTF-8 解码修复。 + +覆盖(全 mock/fake,CI 无真实 redis 也产生覆盖): +- RedisStateStore:put 用 SET NX EX、verify_and_consume 用 GETDEL 一次性消费、 + 重复消费返回 False、Redis 异常降级内存、client 注入 +- Redis 不可用(ping 失败)构造时降级内存,功能仍正常 +- GETDEL 不存在(老 Redis)走 GET+DELETE 兜底 +- handle_callback:微信 sns/userinfo 响应含中文 nickname,resp.encoding=utf-8 + 后解析不乱码;errcode 错误路径返回 errmsg +""" + +from __future__ import annotations + +import sys +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from packages.application.auth import wechat_oauth_service as oauth # noqa: E402 + + +class _FakeRedisClient: + """最小内存版 redis client,模拟 SET NX EX / GETDEL / GET / DELETE / ping。""" + + def __init__(self): + self.data: dict[str, str] = {} + self.ttl: dict[str, int] = {} + self.has_getdel = True + + def ping(self): + return True + + def set(self, key, value, nx=False, ex=None): # noqa: ARG002 + if nx and key in self.data: + return None + self.data[key] = value + if ex is not None: + self.ttl[key] = ex + return True + + def get(self, key): + return self.data.get(key) + + def getdel(self, key): + return self.data.pop(key, None) + + def delete(self, key): + return 1 if self.data.pop(key, None) is not None else 0 + + def eval(self, script, numkeys, key): # noqa: ARG002 + # 模拟 Lua:原子 GET + DEL + return self.data.pop(key, None) + + +# ── RedisStateStore ───────────────────────────────────────────────────── + + +def test_redis_state_store_put_and_consume_once(): + client = _FakeRedisClient() + store = oauth.RedisStateStore(client=client) + store.put("state-abc") + # key 带前缀、TTL 写入 + assert client.data.get("wechat:state:state-abc") is not None + assert client.ttl.get("wechat:state:state-abc") == oauth.STATE_TTL_SECONDS + # 一次性消费:第一次 True,第二次 False + assert store.verify_and_consume("state-abc") is True + assert store.verify_and_consume("state-abc") is False + + +def test_redis_state_store_unknown_state_returns_false(): + store = oauth.RedisStateStore(client=_FakeRedisClient()) + assert store.verify_and_consume("never-put") is False + + +def test_redis_state_store_eval_missing_falls_back_to_get_delete(): + """eval 不可用(如禁用脚本)时退化 GET+DELETE,仍一次性消费。""" + client = _FakeRedisClient() + + def _no_eval(script, numkeys, *keys): # noqa: ARG002 + raise RuntimeError("unknown command EVAL") + + client.eval = _no_eval # type: ignore[method-assign] + store = oauth.RedisStateStore(client=client) + store.put("state-old") + assert store.verify_and_consume("state-old") is True + # GET+DELETE 也消费掉了 + assert "wechat:state:state-old" not in client.data + assert store.verify_and_consume("state-old") is False + + +def test_redis_state_store_put_exception_falls_back_to_memory(): + client = MagicMock() + client.set.side_effect = RuntimeError("redis write fail") + # eval/get 也失败,确保降级到内存 + client.eval.side_effect = RuntimeError("redis read fail") + client.get.side_effect = RuntimeError("redis read fail") + store = oauth.RedisStateStore(client=client) + + store.put("state-fb") # 写 Redis 失败 → 内存 + assert store.verify_and_consume("state-fb") is True # 内存命中 + assert store.verify_and_consume("state-fb") is False + + +def test_redis_state_store_consume_exception_falls_back_to_memory(): + client = MagicMock() + client.set.return_value = True # put 走 Redis + client.eval.side_effect = RuntimeError("redis down") + client.get.side_effect = RuntimeError("redis down") + store = oauth.RedisStateStore(client=client) + + store.put("state-fb2") # 成功写 Redis + # 校验时 Redis 挂了 → 降级内存(内存里没有,返回 False,不报错) + assert store.verify_and_consume("state-fb2") is False + + +def test_redis_state_store_constructor_ping_failure_falls_back(): + """构造时 ping 失败(Redis 不可用)→ 内存降级,功能正常。""" + fake_redis_mod = MagicMock() + fake_client = MagicMock() + fake_client.ping.side_effect = ConnectionError("refused") + fake_redis_mod.Redis.from_url.return_value = fake_client + + with patch.dict(sys.modules, {"redis": fake_redis_mod}): + store = oauth.RedisStateStore(redis_url="redis://nonexistent:6379/0") + + # Redis 不可用 → 内存存储仍工作 + store.put("state-mem") + assert store.verify_and_consume("state-mem") is True + assert store.verify_and_consume("state-mem") is False + + +# ── handle_callback:state 校验 + UTF-8 中文昵称 ──────────────────────── + + +def _configured_service(state_store=None): + store = state_store or oauth.MemoryStateStore() + return oauth.WechatOAuthService( + app_id="wx-test", + app_secret="secret-test", + redirect_uri="https://staging.xiaoxiajianji.com/auth/wechat/callback", + state_store=store, + ) + + +class _FakeResponse: + def __init__(self, payload): + self._payload = payload + self.encoding = None # 模拟微信响应头不带 charset + + def json(self): + # 模拟 requests 行为:按 self.encoding 解码。这里直接返回 payload, + # 但记录 encoding 是否被设置为 utf-8(断言修复生效) + self._decoded_with = self.encoding + return self._payload + + +def test_handle_callback_chinese_nickname_decoded_utf8(monkeypatch): + """微信 userinfo 返回中文昵称,service 设置 encoding=utf-8 后不乱码。""" + service = _configured_service() + state = "state-cn-1" + service._state_store.put(state) + + token_resp = _FakeResponse({"access_token": "at-1", "openid": "openid-cn", "unionid": "union-cn"}) + user_resp = _FakeResponse( + {"openid": "openid-cn", "unionid": "union-cn", "nickname": "微信小应🎬", "headimgurl": ""} + ) + responses = iter([token_resp, user_resp]) + monkeypatch.setattr(oauth.requests, "get", lambda *a, **k: next(responses)) + + info, err = service.handle_callback("code-cn", state) + assert err is None + assert info is not None + assert info.openid == "openid-cn" + assert info.nickname == "微信小应🎬" + # 两个响应都被显式设为 utf-8 + assert token_resp.encoding == "utf-8" + assert user_resp.encoding == "utf-8" + + +def test_handle_callback_state_invalid_returns_error(): + service = _configured_service() + info, err = service.handle_callback("code-x", "state-not-exist") + assert info is None + assert "state" in err + + +def test_handle_callback_wechat_errcode_returns_errmsg(monkeypatch): + """微信返回 errcode(如 code 已被消费 40029)时返回 errmsg 原文。""" + service = _configured_service() + state = "state-err-1" + service._state_store.put(state) + + err_resp = _FakeResponse({"errcode": 40029, "errmsg": "invalid code"}) + monkeypatch.setattr(oauth.requests, "get", lambda *a, **k: err_resp) + + info, err = service.handle_callback("bad-code", state) + assert info is None + assert "invalid code" in err + assert err_resp.encoding == "utf-8" + + +def test_generate_auth_url_stores_state_in_redis(): + """generate_auth_url 生成的 state 写入 Redis(而非仅内存)。""" + client = _FakeRedisClient() + service = oauth.WechatOAuthService( + app_id="wx-test", + app_secret="secret-test", + redirect_uri="https://example.com/cb", + state_store=oauth.RedisStateStore(client=client), + ) + url, state = service.generate_auth_url() + assert f"wechat:state:{state}" in client.data + assert "open.weixin.qq.com" in url + + +# ── _build_default_state_store 工厂分支 ───────────────────────────────── + + +def test_build_default_state_store_uses_redis_when_broker_configured(): + """API settings 有 CELERY_BROKER_URL 时返回 RedisStateStore。""" + store = oauth._build_default_state_store() + # CI/本地通常配置了 redis://localhost:6379/...;无论 Redis 是否可达, + # 返回类型应为 RedisStateStore(内部降级内存) + assert isinstance(store, oauth.RedisStateStore) or isinstance(store, oauth.MemoryStateStore) + + +def test_build_default_state_store_env_fallback(monkeypatch): + """app.config 不可用(如纯 worker 环境)时从环境变量取 redis url。""" + import builtins + + real_import = builtins.__import__ + + def _failing_import(name, *args, **kwargs): + if name == "app.config": + raise ImportError("no app.config") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _failing_import) + monkeypatch.setenv("CELERY_BROKER_URL", "redis://localhost:6379/9") + store = oauth._build_default_state_store() + assert isinstance(store, oauth.RedisStateStore) + + +def test_build_default_state_store_no_config_returns_memory(monkeypatch): + """无任何 redis 配置时返回 MemoryStateStore。""" + import builtins + + real_import = builtins.__import__ + + def _failing_import(name, *args, **kwargs): + if name == "app.config": + raise ImportError("no app.config") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _failing_import) + monkeypatch.delenv("CELERY_BROKER_URL", raising=False) + monkeypatch.delenv("REDIS_URL", raising=False) + store = oauth._build_default_state_store() + assert isinstance(store, oauth.MemoryStateStore)