From c1763b995c2eeaa12295089bdc938eb3163196d2 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 00:44:24 +0800 Subject: [PATCH 01/33] =?UTF-8?q?feat(#1677):=20=E5=A4=9A=E8=A7=86?= =?UTF-8?q?=E9=A2=91=E6=89=B9=E9=87=8F=E7=94=9F=E6=88=90=E5=90=8E=E7=AB=AF?= =?UTF-8?q?=E8=A1=A5=E5=85=A8=20=E2=80=94=20=E6=89=B9=E9=87=8F=E9=A2=84?= =?UTF-8?q?=E8=A7=88=E5=8F=98=E4=BD=93=E6=95=B0=E7=BB=84=20+=20=E6=8C=89?= =?UTF-8?q?=E5=8F=98=E4=BD=93=E7=8B=AC=E7=AB=8B=E6=A0=87=E9=A2=98/?= =?UTF-8?q?=E9=85=8D=E9=9F=B3/=E5=B0=81=E9=9D=A2=20(#1701)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/api/app/api/routes/generation_preview.py | 363 ++++++++----- apps/api/app/api/routes/generation_tasks.py | 32 +- apps/api/app/schemas/generation_task.py | 71 ++- tests/unit/test_1677_batch_variants.py | 507 ++++++++++++++++++ tests/unit/test_generation_preview.py | 47 +- 5 files changed, 877 insertions(+), 143 deletions(-) create mode 100644 tests/unit/test_1677_batch_variants.py diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index c40fdf343..64e5b16ca 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -23,6 +23,7 @@ from app.dependencies import ( get_generation_task_repository, ) from app.schemas.generation_task import ( + BatchPreviewGenerationTaskResponse, CreatePreviewGenerationTaskRequest, PreviewGenerationTaskResponse, ) @@ -193,11 +194,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 +215,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,45 +224,95 @@ 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=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待后再提交", ) from e except GlobalQueueFull as e: raise HTTPException( @@ -273,14 +334,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 +347,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 +419,121 @@ 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, "系统队列已满") - raise HTTPException( - status_code=503, - detail="系统繁忙,请稍后再试", - ) from None + # 关联变体 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 - return _to_preview_response(task) + _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] = [] + 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: + _mark_task_failed(generation_task_repository, task, "待处理任务超限") + except GlobalQueueFull: + _mark_task_failed(generation_task_repository, task, "系统队列已满") + 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): + first_err = next((r.error_message for r in responses if r.error_message), "") + if "待处理任务" in first_err: + raise HTTPException(status_code=429, detail=first_err or "待处理任务超限") + if "队列" in first_err: + raise HTTPException(status_code=503, detail=first_err or "系统繁忙,请稍后再试") + + 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..32391a34d 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -47,6 +47,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, @@ -472,12 +481,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 +513,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 +553,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, ) 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/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_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 -- 2.54.0 From 28b30106682d76678803505071d40b375d9355fb Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 07:57:23 +0800 Subject: [PATCH 02/33] =?UTF-8?q?fix(dedup):=20=E4=BF=AE=E5=A4=8D=E6=9F=A5?= =?UTF-8?q?=E9=87=8D=E7=8E=87=E6=81=92=E4=B8=BA0%=E2=80=94=E2=80=94?= =?UTF-8?q?=E6=8C=87=E7=BA=B9=E7=BB=95=E5=BC=80=E9=99=8D=E9=87=8D=E8=A3=81?= =?UTF-8?q?=E5=89=AA+=E5=B1=80=E9=83=A8=E7=89=87=E6=AE=B5=E5=A4=8D?= =?UTF-8?q?=E7=94=A8+=E9=98=88=E5=80=BC=E6=A0=A1=E5=87=86+3=E4=B8=AA?= =?UTF-8?q?=E5=8D=95=E4=BD=8Dbug=20(#1702)=20(#1703)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/api/app/api/routes/videos.py | 10 +- apps/worker/video_processing/dedup.py | 672 +++++++++++------- apps/worker/video_processing/dedup_helpers.py | 3 +- tests/unit/test_bad_fingerprint_filter.py | 27 +- tests/unit/test_dedup_1702_zero_rate_fix.py | 361 ++++++++++ tests/unit/test_dedup_engine.py | 12 +- tests/unit/test_dedup_pure.py | 17 +- tests/unit/test_dedup_v2.py | 11 +- tests/unit/test_fingerprint_chunks.py | 42 +- .../test_phash_threshold_calibration_1658.py | 38 +- 10 files changed, 884 insertions(+), 309 deletions(-) create mode 100644 tests/unit/test_dedup_1702_zero_rate_fix.py 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/worker/video_processing/dedup.py b/apps/worker/video_processing/dedup.py index b4c6aaba3..a1f57b390 100755 --- a/apps/worker/video_processing/dedup.py +++ b/apps/worker/video_processing/dedup.py @@ -31,25 +31,60 @@ 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 真实数据回归校准(2026-09-05,worker 容器内离线实验): +# - 同源成片对(20s/11s,各自 2-5% 随机边缘裁剪降重,1s 密集采样): +# 全部帧对最小汉明距离 min=8,<=12 命中 10/31 帧(B->A 4/11) +# - 异源成片对(4 个不同项目真实视频):最小距离 24,<=16 命中 0 帧 +# 8(#1658 旧值)会漏掉同源裁剪(自对照实验:同帧两次 2-5% 随机裁剪距离 4~10), +# 12 能检出同源/局部复用且与异源分布(>=24)间隔 12bit,无误报空间。 +PHASH_THRESHOLD = 12 +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 +136,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 +251,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 +362,32 @@ 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 内优先选择与上一匹配帧 + 目标序号连贯(0 <= delta <= neighbor_window+1,允许 ±1 邻接窗口 / + 时序偏移对齐,缓解场景切割导致的切点、取帧错位)的候选;同距时 + 偏好大索引,避免重复 hash 塌缩到 target 首帧。 + 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 +395,93 @@ 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 内偏好与上一匹配帧目标序号连贯(0 <= 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 0 <= 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 +495,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 +521,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 +564,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 +589,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 +650,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 +707,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, @@ -649,6 +811,9 @@ class VideoDeduplicator: else: existing_videos = video_repo.list_by_project(project_id) + best_score = 0.0 + best_result: Optional[dict] = None + for existing in existing_videos: if not existing.video_fingerprint: continue @@ -677,61 +842,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 +937,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 +973,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 +1074,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 +1109,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 +1181,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) @@ -1045,7 +1229,9 @@ 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, ) 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..6d96eb9cf 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, diff --git a/tests/unit/test_bad_fingerprint_filter.py b/tests/unit/test_bad_fingerprint_filter.py index 05a18cb08..8e4ec78fe 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,7 +286,7 @@ class TestComputeDuplicateRateBadFingerprint: mock_session = MagicMock() videos = [ - _make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 5), + _make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 10), _make_existing_video("vid-b2", "md5_b2", ["bbbbbbbbbbbbbbbb"] * 5), ] mock_repo = MagicMock() 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..e8d2fdafd --- /dev/null +++ b/tests/unit/test_dedup_1702_zero_rate_fix.py @@ -0,0 +1,361 @@ +"""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 <= 12, "阈值应经校准保持在能检出同源裁剪的范围" + 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_duration_sec_not_divided_by_1000(self): + """fingerprint.duration 单位是秒,传给 check_duplicate 不应再 /1000。 + + 旧 bug:duration/1000 → duration_max≈0.0135s,所有真实视频被过滤。 + """ + + 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) + _, kwargs = repo.list_by_user.call_args + # ±15% 窗口:13.5s -> [11.475, 15.525] + assert 11.0 < kwargs["duration_min"] < 12.0 + assert 15.0 < kwargs["duration_max"] < 16.0 + + +# ── 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) 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..ce3a91213 100644 --- a/tests/unit/test_dedup_v2.py +++ b/tests/unit/test_dedup_v2.py @@ -484,8 +484,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 +497,9 @@ class TestConstants: """常量值验证 — 使用已在模块顶部导入的常量,避免重新 import.""" def test_segment_match_threshold(self): - # 从已导入的 find_duplicate_segments 默认参数间接验证 - assert SEGMENT_MATCH_THRESHOLD == 8 + # Issue #1702: pHash 阈值经 staging 真实同源/异源指纹回归校准 + # (同源密集采样 min=8、异源 min=24),统一为模块常量 PHASH_THRESHOLD=12。 + assert SEGMENT_MATCH_THRESHOLD == 12 def test_min_consecutive_matches(self): assert MIN_CONSECUTIVE_MATCHES == 5 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_phash_threshold_calibration_1658.py b/tests/unit/test_phash_threshold_calibration_1658.py index 6e01b4a6a..dc15dc120 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,15 @@ _ZERO_HIST = [0.0] * 96 # 全黑视频的全零直方图(有效数据) class TestThresholdCalibration: - """pHash 阈值由 10 收紧到 8(Issue #1658)。""" + """pHash 阈值校准(Issue #1658 收紧到 8,Issue #1702 经真实指纹分布重校准为 12)。 - 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 + 采样)最小距离 8、<=12 命中 10/31;异源成片最小距离 24。8 会漏检同源裁剪, + 12 检出同源且与异源分布(>=24)间隔充足。 + """ + + def test_phash_threshold_is_calibrated(self): + assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD == 12 def test_match_ratio_threshold_constant(self): assert MATCH_RATIO_THRESHOLD == 0.7 @@ -144,22 +149,21 @@ 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, 12, 12, 24, 26]。 + - <=12(#1702 校准阈值):3 帧匹配 → 0.6 < 0.7 被帧比例门槛拦截异源 + - 距离 12 的同源裁剪帧应算匹配(< 与 <= 口径统一) """ - distances = [7, 7, 7, 9, 9] + distances = [10, 12, 12, 24, 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 # 被帧比例门槛拦截 + # 异源典型距离(>=24)绝不匹配 + assert not any(d <= VideoDeduplicator.PHASH_THRESHOLD for d in (24, 26, 30)) # ── TestComputeFusionScore:统一融合得分方法 ──────────────────── -- 2.54.0 From 7a4aa27f71b22c529ae451afaf79a7b0c2576eb6 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 08:36:19 +0800 Subject: [PATCH 03/33] =?UTF-8?q?fix(dedup):=20recompute=E4=BB=BB=E5=8A=A1?= =?UTF-8?q?=E4=BB=8Efile=5Furl=E6=B4=BE=E7=94=9FOSS=E4=B8=8B=E8=BD=BDkey?= =?UTF-8?q?=EF=BC=8C=E4=BF=AE=E5=A4=8D=E9=87=8D=E7=AE=97404=20(#1702)=20(#?= =?UTF-8?q?1705)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/worker/video_processing/dedup.py | 14 ++++++++--- tests/unit/test_dedup_1702_zero_rate_fix.py | 28 +++++++++++++++++++++ 2 files changed, 39 insertions(+), 3 deletions(-) diff --git a/apps/worker/video_processing/dedup.py b/apps/worker/video_processing/dedup.py index a1f57b390..113908dfe 100755 --- a/apps/worker/video_processing/dedup.py +++ b/apps/worker/video_processing/dedup.py @@ -1216,9 +1216,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) diff --git a/tests/unit/test_dedup_1702_zero_rate_fix.py b/tests/unit/test_dedup_1702_zero_rate_fix.py index e8d2fdafd..6ddf8eab6 100644 --- a/tests/unit/test_dedup_1702_zero_rate_fix.py +++ b/tests/unit/test_dedup_1702_zero_rate_fix.py @@ -359,3 +359,31 @@ class TestZeroMatchLogging: 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"', + "", + ) -- 2.54.0 From af25045123e4d77f649de87e66f9489de56d25f9 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 09:03:39 +0800 Subject: [PATCH 04/33] =?UTF-8?q?feat(#1677):=20=E5=A4=9A=E8=A7=86?= =?UTF-8?q?=E9=A2=91=E6=89=B9=E9=87=8F=E7=94=9F=E6=88=90=E5=89=8D=E7=AB=AF?= =?UTF-8?q?=20=E2=80=94=20=E6=95=B0=E9=87=8F=E5=BC=B9=E7=AA=97/=E6=89=B9?= =?UTF-8?q?=E9=87=8F=E9=A2=84=E8=A7=88/=E7=8B=AC=E7=AB=8B=E6=A0=87?= =?UTF-8?q?=E9=A2=98=E9=85=8D=E9=9F=B3=E5=B0=81=E9=9D=A2/5=E6=AD=A5?= =?UTF-8?q?=E6=B5=81=E7=A8=8B=20(#1704)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/src/api/generation/types.ts | 29 +- apps/web/src/api/tasks/types.ts | 8 + apps/web/src/pages/generate/GeneratePage.tsx | 429 +++++++++++++++--- .../components/GenerateStepActions.tsx | 36 +- .../components/GenerateStepContent.tsx | 89 ++-- .../generate/components/PreviewCountModal.tsx | 127 ++++++ .../generate/components/ServerPreviewGrid.tsx | 127 ++++++ .../components/Step3VoiceWithMode.tsx | 102 +++++ .../components/Step4TitleSettings.tsx | 245 +++++----- .../generate/components/Step5VoiceSelect.tsx | 19 +- .../components/Step6CoverSettings.tsx | 188 +++++++- .../components/Step7ConfirmGenerate.tsx | 79 ---- .../components/step7-confirm/SummaryCard.tsx | 44 -- apps/web/src/pages/generate/constants.ts | 7 +- apps/web/src/pages/generate/generate.css | 378 +++++++++++++++ .../generate/hooks/generate-video/types.ts | 15 +- .../generate-video/useGenerationPolling.ts | 240 ++++++---- .../pages/generate/hooks/useBatchCovers.ts | 150 ++++++ .../pages/generate/hooks/useBatchPreview.ts | 285 ++++++++++++ .../hooks/useGenerateFormState/index.ts | 42 +- .../pages/generate/hooks/useGenerateVideo.ts | 47 +- .../pages/generate/hooks/useServerPreview.ts | 9 +- .../pages/generate/hooks/useStep7Generate.ts | 112 ----- .../pages/generate/hooks/useStepNavigation.ts | 49 +- .../src/test/pages/generate/smoke.test.tsx | 5 + 25 files changed, 2245 insertions(+), 616 deletions(-) create mode 100644 apps/web/src/pages/generate/components/PreviewCountModal.tsx create mode 100644 apps/web/src/pages/generate/components/ServerPreviewGrid.tsx create mode 100644 apps/web/src/pages/generate/components/Step3VoiceWithMode.tsx delete mode 100755 apps/web/src/pages/generate/components/Step7ConfirmGenerate.tsx delete mode 100755 apps/web/src/pages/generate/components/step7-confirm/SummaryCard.tsx create mode 100644 apps/web/src/pages/generate/hooks/useBatchCovers.ts create mode 100644 apps/web/src/pages/generate/hooks/useBatchPreview.ts delete mode 100644 apps/web/src/pages/generate/hooks/useStep7Generate.ts 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/pages/generate/GeneratePage.tsx b/apps/web/src/pages/generate/GeneratePage.tsx index db681da43..57fad236c 100644 --- a/apps/web/src/pages/generate/GeneratePage.tsx +++ b/apps/web/src/pages/generate/GeneratePage.tsx @@ -1,12 +1,11 @@ /** - * 智能剪辑页面 — 前端实时预览架构 - * 6 步向导:选择模板 → 素材 → 配音 → 标题(含预览) → 确认生成 → 选择封面 + * 智能剪辑页面(Issue #1677 多视频批量生成) + * 5 步向导:选择模板(弹数量) → 素材 → 配音 → 标题(预览+确认生成) → 封面 * * 架构: - * - 步骤 4 右侧显示 FrontendPreviewPlayer 实时预览 - * - 步骤 5 右侧内联播放生成中的/最终视频 - * - 步骤 6 封面从最终成片中智能选帧(MediaKit) - * - 点"确认生成"时调用 createGenerationTask 创建一次服务器渲染任务 + * - N=1:前端 Canvas 实时预览(FrontendPreviewPlayer),零回归 + * - N>1:服务器批量预览(POST /generation/preview?preview_count=N), + * N 个变体分别轮询,网格展示、独立可播放、CSS 标题浮层实时叠加、勾选批量生成 */ import React, { useMemo, useState, useEffect, useRef, useCallback } from "react" import { message } from "antd" @@ -17,16 +16,20 @@ import { useCloneProgress } from "@/hooks/useCloneProgress" import CloneModal from "@/components/voice/CloneModal" import GenerateHeader from "./components/GenerateHeader" import FrontendPreviewPlayer from "./components/FrontendPreviewPlayer" +import ServerPreviewGrid from "./components/ServerPreviewGrid" +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 { useBatchPreview } from "./hooks/useBatchPreview" import { usePreviewAssets } from "./hooks/usePreviewAssets" import { useTitleStyleUpdaters } from "./hooks/useStep4Title/useTitleStyleUpdaters" import { getAssetsByKind } from "@/api/assets" import { previewTts } from "@/api/tts" +import { calculateResolution } from "./utils/calculateResolution" import "./generate.css" const GeneratePage: React.FC = () => { @@ -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, @@ -127,7 +171,7 @@ const GeneratePage: React.FC = () => { }, [selectedVoice, selectedClonedVoice, titleSettings.title, voiceMaterials]) /* ── 克隆声音 ── */ - const { clones: clonedVoices, addClone, hasProcessing } = useCloneProgress() + const { addClone } = useCloneProgress() const handleCloneSuccess = (voice: VoiceClone) => { addClone(voice) @@ -156,19 +200,95 @@ const GeneratePage: React.FC = () => { [bgm, currentTemplate], ) - /* ── 加载素材详情(供前端预览播放器使用 + 配音时长校验) ── */ + /* ── 加载素材详情(供前端预览播放器使用) ── */ const previewAssetsEnabled = previewAssetIds.length > 0 const { assets: previewAssets, ready: previewAssetsReady } = usePreviewAssets( previewAssetIds, previewAssetsEnabled, ) - /* ── 预览就绪:素材已加载,且有模板 ── */ - const previewReady = useMemo( + /* ── 预览就绪 ── */ + const singlePreviewReady = useMemo( () => previewAssetsReady && !!currentTemplate, [previewAssetsReady, currentTemplate], ) + /* ── 批量服务器预览(N>1) ── */ + const buildPreviewRequest = useCallback(() => { + const { width, height } = calculateResolution(videoRatio || "9:16") + const voiceLibraryId = + voiceMode === "clone" ? selectedClonedVoice || selectedVoice || "" : selectedVoice || "" + return { + template_id: selectedTemplate, + asset_ids: previewAssetIds, + output_width: width, + output_height: height, + video_ratio: videoRatio, + voice_library_id: voiceLibraryId, + ...(voiceModePerVideo && voiceLibraryIds.some(Boolean) + ? { voice_library_ids: voiceLibraryIds.map((id) => id || voiceLibraryId) } + : {}), + preview_count: previewCount, + // 批量预览不传 titles/title_config:标题文字与样式由前端 CSS 浮层实时叠加 + // (用户改标题/样式即时可见,无需重渲染);正式生成时才把标题烧录进成片 + duration: duration || undefined, + bgm_config: { + enabled: bgm !== false, + ...(bgmConfig?.music_id ? { preset_id: bgmConfig.music_id } : {}), + }, + ...(storedSourceEditPlanId || sourceEditPlanId + ? { source_edit_plan_id: storedSourceEditPlanId || sourceEditPlanId || undefined } + : {}), + } + }, [ + videoRatio, + voiceMode, + selectedClonedVoice, + selectedVoice, + selectedTemplate, + previewAssetIds, + voiceModePerVideo, + voiceLibraryIds, + previewCount, + duration, + bgm, + bgmConfig, + storedSourceEditPlanId, + sourceEditPlanId, + ]) + + const { + variants, + status: batchPreviewStatus, + progress: batchPreviewProgress, + failedCount: batchFailedCount, + trigger: retryBatchPreview, + } = useBatchPreview({ + enabled: isBatch && currentStep >= 4 && previewAssetIds.length > 0 && !!selectedTemplate, + buildRequest: buildPreviewRequest, + onPreviewTasksCreated: (_taskIds, planId) => { + if (planId) setStoredSourceEditPlanId(planId) + }, + }) + + /** 批量预览就绪:全部变体渲染完成 */ + const batchPreviewReady = + isBatch && variants.length > 0 && variants.every((v) => v.status === "ready") + + /** 步骤4整体预览就绪状态 */ + const previewReady = isBatch ? batchPreviewReady : singlePreviewReady + + /* ── 勾选变体 ── */ + const toggleVariantSelect = useCallback( + (index: number) => { + setSelectedVariantIds((prev) => { + const list = prev || [] + return list.includes(index) ? list.filter((i) => i !== index) : [...list, index].sort() + }) + }, + [setSelectedVariantIds], + ) + /* ── 视频生成核心逻辑 ── */ const { generating, @@ -199,16 +319,61 @@ 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「确认生成视频」 ── */ const handleConfirmGenerate = useCallback(async () => { - if (!titleSettings.title.trim()) { - message.warning("请选择或输入标题") + // 标题校验 + if (previewTitles.some((t) => !t?.trim())) { + message.warning("请为每个视频输入标题") + return + } + if (isBatch && selectedVariantIds.length === 0) { + message.warning("请至少勾选一个视频") return } if (!previewReady) { @@ -216,10 +381,18 @@ const GeneratePage: React.FC = () => { return } const ok = await handleGenerate() - if (ok) { + if (ok && !isBatch) { setCurrentStep(5) } - }, [titleSettings.title, previewReady, handleGenerate, setCurrentStep]) + // 批量模式停留在步骤4,右侧网格显示生成进度,完成后点"下一步"进封面 + }, [ + isBatch, + selectedVariantIds.length, + previewReady, + previewTitles, + handleGenerate, + setCurrentStep, + ]) /* ── 步骤导航 ── */ const { goNext, goPrev } = useStepNavigation({ @@ -232,11 +405,24 @@ const GeneratePage: React.FC = () => { titleSettings, previewReady, generated, + previewTitles, + selectedCount: isBatch ? selectedVariantIds.length : 1, + onOpenCountModal: () => setCountModalOpen(true), }) - /* ── 最终成片(步骤5/6 右侧播放) ── */ + /* ── 最终成片(单视频右侧播放) ── */ const finalVideo = generatedVideos[0] + /** 批量生成进度文案 */ + const batchGeneratingText = useMemo(() => { + if (batchPreviewStatus === "loading") + return `AI 正在渲染 ${previewCount} 个预览视频… ${batchPreviewProgress}%` + if (batchPreviewStatus === "failed") return "预览渲染失败,请重试" + if (batchPreviewStatus === "partial_failed") + return `${batchFailedCount} 个预览失败,可重新生成或勾选成功的视频` + return "" + }, [batchPreviewStatus, batchPreviewProgress, batchFailedCount, previewCount]) + /* ================================================================ 渲染 ================================================================ */ @@ -247,8 +433,142 @@ const GeneratePage: React.FC = () => { -
- {/* ════ 左侧:表单区 ════ */} +
+ {/* ════ 步骤4:左侧预览大区域 ════ */} + {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} + /> + ) : ( + /* 批量:服务器预览网格 */ +
+
+

🎬 {previewCount} 个视频预览

+ {batchPreviewStatus === "loading" && ( + + {batchPreviewProgress}% + + )} + {(batchPreviewStatus === "failed" || batchPreviewStatus === "partial_failed") && ( + + )} +
+ + {batchGeneratingText && ( +
+ {batchGeneratingText} +
+ )} + + + + {/* 生成中进度(批量) */} + {generating && ( +
+
+
+
+ ⏳ 正在渲染 {selectedVariantIds.length} 个最终视频… {Math.round(progress)} + % +
+
+ 生成过程中可以切换到其他页面,完成后可在任务历史查看 +
+
+
+
+
+
+
+ )} + + {generateError && !generating && ( +
+
+
生成失败
+
{generateError}
+
+ +
+ )} + + {generated && !generating && ( +
+
+
✅ 视频生成完成!
+
+ 共生成 {generatedVideos.length} 条视频,点击「下一步」为每个视频选择封面 +
+
+
+ )} +
+ )} +
+ )} + + {/* ════ 右侧:步骤1~3 表单 / 步骤4 标题边栏 / 步骤5 封面 ════ */}
{ 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} @@ -296,7 +607,16 @@ const GeneratePage: React.FC = () => { generatedVideos={generatedVideos} onRetry={handleRetryGenerate} onDismissError={handleDismissError} - presetVoices={presetVoices} + 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(封面):成片播放器(单视频) ════ */} + {currentStep === 5 && !isBatch && generated && finalVideo && ( +
- )} -
+
+ )}
+ {/* 数量选择弹窗 */} + setCountModalOpen(false)} + /> + {/* 音色克隆弹窗 */} = ({ +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 ( ) } return ( ) } - /* 步骤 5:渲染中禁用,完成后下一步进入封面 */ - if (currentStep === 5) { - return ( - - ) - } - - /* 步骤 6(最后一步):无主按钮 */ + /* 步骤 5(封面,最后一步):无主按钮 */ return null } diff --git a/apps/web/src/pages/generate/components/GenerateStepContent.tsx b/apps/web/src/pages/generate/components/GenerateStepContent.tsx index 1b160b5ec..47a1a7ba8 100644 --- a/apps/web/src/pages/generate/components/GenerateStepContent.tsx +++ b/apps/web/src/pages/generate/components/GenerateStepContent.tsx @@ -1,19 +1,16 @@ /** * GeneratePage 步骤内容渲染 - * 步骤顺序(6步):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6) + * 步骤顺序(5步,Issue #1677):模板(1) → 素材(2) → 配音(3) → 标题+预览+确认生成(4) → 封面(5) */ 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 type { GeneratedVideo } from "@/api/template-editor" @@ -50,15 +47,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 @@ -67,12 +55,22 @@ export interface GenerateStepContentProps { generatedVideos: GeneratedVideo[] onRetry: () => void onDismissError: () => void - /* 其他 */ - presetVoices: PresetVoiceItem[] /** 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 +102,17 @@ export const GenerateStepContent: React.FC = (props) = selectedVoice, onSelectedVoiceChange, onServerClipsChange, - voiceMode, - selectedClonedVoice, - clonedVoices, - generating, - generated, - generateError, - progress, generatedVideos, - onRetry, - onDismissError, - presetVoices, + previewCount, + previewTitles, + onPreviewTitlesChange, + voiceModePerVideo, + onVoiceModePerVideoChange, + voiceLibraryIds, + onVoiceLibraryIdsChange, + previewCovers, + onPreviewCoversChange, + selectedVariantIds, } = props /* 当前模板的 segments,传给 Step2 构建 clips */ @@ -146,9 +144,14 @@ export const GenerateStepContent: React.FC = (props) = ) case 3: return ( - ) case 4: @@ -167,33 +170,12 @@ export const GenerateStepContent: React.FC = (props) = onApplyPreset={onApplyPreset} activePreset={activePreset} titlePresets={titlePresets} + previewCount={previewCount} + previewTitles={previewTitles} + onPreviewTitlesChange={onPreviewTitlesChange} /> ) case 5: - return ( - - ) - case 6: return ( = (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/ServerPreviewGrid.tsx b/apps/web/src/pages/generate/components/ServerPreviewGrid.tsx new file mode 100644 index 000000000..7e7b0d601 --- /dev/null +++ b/apps/web/src/pages/generate/components/ServerPreviewGrid.tsx @@ -0,0 +1,127 @@ +/** + * 批量预览网格(Issue #1677) + * N 个服务器渲染的预览视频,网格排列、各自独立播放、CSS 标题浮层实时叠加、勾选框批量选择 + */ +import React from "react" +import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons" +import type { VariantPreview } from "../hooks/useBatchPreview" + +interface ServerPreviewGridProps { + variants: VariantPreview[] + /** 每个变体的标题文字(实时叠加浮层) */ + titles: string[] + /** 标题样式(全局共用) */ + titleStyle: { + position: string + color: string + size: number + } + /** 勾选的变体索引 */ + selectedIds: number[] + onToggleSelect: (index: number) => void + /** 是否显示勾选框(确认生成前) */ + selectable?: boolean +} + +const ServerPreviewGrid: React.FC = ({ + variants, + titles, + titleStyle, + selectedIds, + onToggleSelect, + selectable = true, +}) => { + if (variants.length === 0) return null + + return ( +
+ {variants.map((v) => { + const selected = selectedIds.includes(v.index) + const titleText = titles[v.index] || "" + return ( +
{ + if (selectable && v.status === "ready") onToggleSelect(v.index) + }} + role="button" + tabIndex={0} + > + {/* 勾选框 */} + {selectable && v.status === "ready" && ( +
+ {selected && "✓"} +
+ )} + + {/* 变体序号 */} +
视频 {v.index + 1}
+ + {/* 视频区域 */} +
+ {v.status === "loading" && ( +
+ +
+
+
+ {v.progress}% +
+ )} + {v.status === "failed" && ( +
+ + {v.error || "预览失败"} +
+ )} + {v.status === "ready" && v.videoUrl && ( + <> +
+ + {/* 底部状态 */} +
+ {v.status === "ready" && selected && ( + + 已选择 + + )} + {v.status === "ready" && !selected && selectable && ( + 点击卡片取消/勾选 + )} +
+
+ ) + })} +
+ ) +} + +export default ServerPreviewGrid 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..7c92a93ba 100644 --- a/apps/web/src/pages/generate/components/Step4TitleSettings.tsx +++ b/apps/web/src/pages/generate/components/Step4TitleSettings.tsx @@ -1,12 +1,13 @@ /** - * Step 4 选择标题(合并原 Step4 标题输入 + Step5 标题样式面板) + * Step 4 选择标题(Issue #1677 批量生成改造) * - * 左侧:标题文字输入 + AI生成标题 + 样式设置(位置/字号/字体/颜色/样式/预设) - * 右侧:FrontendPreviewPlayer 实时预览(由 GeneratePage 统一渲染) + * 布局(由 GeneratePage 编排):左侧大区域预览,右侧边栏标题设置。 + * 本组件渲染在右侧边栏: + * - 标题文字:1 个视频 1 个输入框;N 个视频 N 个输入框各自独立 + * - 标题样式(字体/颜色/位置/大小/粗斜描边/预设):全局统一 */ import React from "react" -import { AutoComplete } from "antd" -import { PlayCircleOutlined } from "@ant-design/icons" +import { AutoComplete, Input } from "antd" import type { TitleSettings } from "../types" import { POSITION_OPTIONS, FONT_OPTIONS } from "../constants" import { useStep4Title } from "../hooks/useStep4Title" @@ -29,6 +30,12 @@ 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 } const Step4TitleSettings: React.FC = (props) => { @@ -44,122 +51,138 @@ const Step4TitleSettings: React.FC = (props) => { onApplyPreset, activePreset, titlePresets, + previewCount = 1, + previewTitles, + onPreviewTitlesChange, } = props + const isBatch = previewCount > 1 + + /** 更新单个变体标题;变体0同步写回 titleSettings.title(全局样式面板/草稿保存依赖) */ + const updateVariantTitle = (index: number, val: string) => { + if (!previewTitles || !onPreviewTitlesChange) return + const next = [...previewTitles] + next[index] = val + onPreviewTitlesChange(next) + if (index === 0) { + t.updateTitle(val) + } + } + 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={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()) + }} + /> +
+ + )} + ) : ( + /* ── 批量:N 个独立标题输入框(CSS 浮层实时叠加到对应预览) ── */ +
+
+ 为每个视频输入独立标题,修改会实时叠加到左侧对应视频上。标题样式(字体/颜色/位置)全局统一。 +
+ {Array.from({ length: previewCount }, (_, i) => ( +
+ + updateVariantTitle(i, e.target.value)} + /> +
+ ))} +
)} - {/* 标题样式面板(原 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) => { )}
- - )} + + 实时预览,勾选要生成的视频 +
- - {batchGeneratingText && ( -
- {batchGeneratingText} -
- )} - - - - {/* 生成中进度(批量) */} - {generating && ( -
-
-
-
- ⏳ 正在渲染 {selectedVariantIds.length} 个最终视频… {Math.round(progress)} - % -
-
- 生成过程中可以切换到其他页面,完成后可在任务历史查看 -
-
-
-
-
-
-
- )} - - {generateError && !generating && ( -
-
-
生成失败
-
{generateError}
-
- -
- )} - - {generated && !generating && ( -
-
-
✅ 视频生成完成!
-
- 共生成 {generatedVideos.length} 条视频,点击「下一步」为每个视频选择封面 -
-
-
- )}
)}
)} - {/* ════ 右侧:步骤1~3 表单 / 步骤4 标题边栏 / 步骤5 封面 ════ */} + {/* ════ 右侧:步骤1~3 表单 / 步骤4 标题边栏 / 步骤5 确认生成进度 / 步骤6 封面 ════ */}
{ progress={progress} generatedVideos={generatedVideos} onRetry={handleRetryGenerate} + onRetryBatchTask={handleRetryBatchTask} onDismissError={handleDismissError} + batchTasks={batchTasks} previewCount={previewCount} previewTitles={previewTitles} onPreviewTitlesChange={setPreviewTitles} @@ -631,16 +498,16 @@ const GeneratePage: React.FC = () => { />
- {/* ════ 步骤5(封面):成片播放器(单视频) ════ */} - {currentStep === 5 && !isBatch && generated && finalVideo && ( + {/* ════ 步骤5/6(单视频):右侧成片播放器 ════ */} + {currentStep >= 5 && !isBatch && generated && finalVideo && (
+
+ ) + })} +
+
+ ) +} + +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..1d0b17674 --- /dev/null +++ b/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx @@ -0,0 +1,95 @@ +/** + * 批量前端 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, +}) => { + 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 0f816e23d..c25e75abe 100644 --- a/apps/web/src/pages/generate/components/GenerateStepActions.tsx +++ b/apps/web/src/pages/generate/components/GenerateStepActions.tsx @@ -1,10 +1,10 @@ /** - * GeneratePage 步骤底部操作按钮(Issue #1677 改造后 5 步) + * GeneratePage 步骤底部操作按钮(Issue #1677 修正:固定 6 步) * * 步骤 1~3:上一步 / 下一步 - * 步骤 4(标题+预览+确认生成):确认生成按钮在右侧边栏底部(含勾选数量), - * 渲染中显示进度;生成完成后显示"下一步 → 选择封面" - * 步骤 5(选择封面):仅上一步 + * 步骤 4(选择标题):「✨ 确认生成视频 / 确认生成 N 个视频」→ 创建正式生成任务,成功后跳步骤5 + * 步骤 5(确认生成):渲染进度页,全部完成后「下一步:选择封面」;仅上一步 + * 步骤 6(选择封面):仅上一步 */ import React from "react" @@ -12,7 +12,7 @@ export interface GenerateStepActionsProps { currentStep: number onPrev: () => void onNext: () => void - /** 步骤4:确认生成视频(校验 + 创建渲染任务 + 成功后进入步骤5) */ + /** 步骤4:确认生成视频(校验 + 创建渲染任务) */ onConfirmGenerate: () => void | Promise generating: boolean generated: boolean @@ -41,26 +41,19 @@ const GenerateStepActions: React.FC = ({ ) } - /* 步骤 4:标题+预览+确认生成 */ + /* 步骤 4:选择标题 — 确认生成 */ if (currentStep === 4) { if (generating) { return ( ) } if (generateError) { return ( - ) - } - if (generated) { - return ( - ) } @@ -71,7 +64,23 @@ const GenerateStepActions: React.FC = ({ ) } - /* 步骤 5(封面,最后一步):无主按钮 */ + /* 步骤 5:确认生成进度页 — 全部完成后下一步进封面 */ + if (currentStep === 5) { + if (generated) { + return ( + + ) + } + return ( + + ) + } + + /* 步骤 6(封面,最后一步):无主按钮 */ return null } diff --git a/apps/web/src/pages/generate/components/GenerateStepContent.tsx b/apps/web/src/pages/generate/components/GenerateStepContent.tsx index 47a1a7ba8..639c99aa5 100644 --- a/apps/web/src/pages/generate/components/GenerateStepContent.tsx +++ b/apps/web/src/pages/generate/components/GenerateStepContent.tsx @@ -1,6 +1,7 @@ /** * GeneratePage 步骤内容渲染 - * 步骤顺序(5步,Issue #1677):模板(1) → 素材(2) → 配音(3) → 标题+预览+确认生成(4) → 封面(5) + * 步骤顺序(6步,Issue #1677 修正):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6) + * 步骤4预览(Canvas 网格)与步骤5进度(批量渲染网格)由 GeneratePage 直接渲染在左侧大区域。 */ import React from "react" import type { EditingTemplate } from "@/api/editing-planner" @@ -12,6 +13,8 @@ import Step2MaterialSelect from "../components/Step2MaterialSelect" import Step3VoiceWithMode from "./Step3VoiceWithMode" import Step4TitleSettings from "../components/Step4TitleSettings" 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 { @@ -54,7 +57,10 @@ export interface GenerateStepContentProps { progress: number generatedVideos: GeneratedVideo[] onRetry: () => void + onRetryBatchTask: (taskId: string) => void onDismissError: () => void + /** 批量:每个正式生成任务的独立状态(步骤5进度网格) */ + batchTasks: BatchTaskState[] /** BGM 开关 */ bgm: boolean /** BGM 配置(来自模板) */ @@ -102,7 +108,14 @@ export const GenerateStepContent: React.FC = (props) = selectedVoice, onSelectedVoiceChange, onServerClipsChange, + generating, + generated, + generateError, + progress, + onRetry, generatedVideos, + batchTasks, + onRetryBatchTask, previewCount, previewTitles, onPreviewTitlesChange, @@ -176,6 +189,62 @@ export const GenerateStepContent: React.FC = (props) = /> ) case 5: + /* 确认生成页:批量=逐任务进度网格;单视频=进度状态卡(成片播放器在左侧大区域) */ + if (previewCount > 1) { + return ( + + ) + } + /* 单视频:渲染进度 / 失败重试 / 完成提示(成片播放器在右侧栏) */ + return ( +
+

🎬 确认生成

+ {generating && ( +
+
+
+
+ ⏳ 视频渲染中… {Math.round(progress)}% +
+
+ 生成过程中可以切换到其他页面,完成后可在任务历史查看 +
+
+
+
+
+
+
+ )} + {generateError && !generating && ( +
+
+
生成失败
+
{generateError}
+
+ +
+ )} + {generated && !generating && ( +
+
+
✅ 视频生成完成!
+
右侧可预览成片,点击「下一步」选择封面
+
+
+ )} +
+ ) + case 6: return ( void - /** 是否显示勾选框(确认生成前) */ - selectable?: boolean -} - -const ServerPreviewGrid: React.FC = ({ - variants, - titles, - titleStyle, - selectedIds, - onToggleSelect, - selectable = true, -}) => { - if (variants.length === 0) return null - - return ( -
- {variants.map((v) => { - const selected = selectedIds.includes(v.index) - const titleText = titles[v.index] || "" - return ( -
{ - if (selectable && v.status === "ready") onToggleSelect(v.index) - }} - role="button" - tabIndex={0} - > - {/* 勾选框 */} - {selectable && v.status === "ready" && ( -
- {selected && "✓"} -
- )} - - {/* 变体序号 */} -
视频 {v.index + 1}
- - {/* 视频区域 */} -
- {v.status === "loading" && ( -
- -
-
-
- {v.progress}% -
- )} - {v.status === "failed" && ( -
- - {v.error || "预览失败"} -
- )} - {v.status === "ready" && v.videoUrl && ( - <> -
- - {/* 底部状态 */} -
- {v.status === "ready" && selected && ( - - 已选择 - - )} - {v.status === "ready" && !selected && selectable && ( - 点击卡片取消/勾选 - )} -
-
- ) - })} -
- ) -} - -export default ServerPreviewGrid diff --git a/apps/web/src/pages/generate/components/Step4TitleSettings.tsx b/apps/web/src/pages/generate/components/Step4TitleSettings.tsx index 7c92a93ba..064bc167a 100644 --- a/apps/web/src/pages/generate/components/Step4TitleSettings.tsx +++ b/apps/web/src/pages/generate/components/Step4TitleSettings.tsx @@ -1,18 +1,22 @@ /** - * Step 4 选择标题(Issue #1677 批量生成改造) + * Step 4 选择标题(Issue #1677 批量生成) * - * 布局(由 GeneratePage 编排):左侧大区域预览,右侧边栏标题设置。 - * 本组件渲染在右侧边栏: - * - 标题文字:1 个视频 1 个输入框;N 个视频 N 个输入框各自独立 + * 布局(由 GeneratePage 编排):左侧大区域实时预览(单=大播放器,批量=Canvas 网格), + * 右侧边栏标题设置。本组件渲染在右侧边栏: + * - 单视频:AI 标题生成器 + AutoComplete 标题库(与旧版完全一致,零回归) + * - 批量:N 个独立标题输入框(AutoComplete 支持标题库选择)+ 批量 AI 生成 + * (一次生成 N 个标题,分别填入各变体,可单独换一个) * - 标题样式(字体/颜色/位置/大小/粗斜描边/预设):全局统一 */ -import React from "react" -import { AutoComplete, Input } from "antd" +import React, { useMemo, useState } from "react" +import { AutoComplete, 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 TitleStylePanel from "./title/TitleStylePanel" +import { AI_TITLE_TEMPLATES } from "../constants" interface Step4TitleSettingsProps { titleSettings: TitleSettings @@ -38,6 +42,36 @@ interface Step4TitleSettingsProps { 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) => { const t = useStep4Title(props) const { @@ -57,8 +91,10 @@ const Step4TitleSettings: React.FC = (props) => { } = props const isBatch = previewCount > 1 + const [batchAiLoading, setBatchAiLoading] = useState(false) + const [batchAiTopic, setBatchAiTopic] = useState("") - /** 更新单个变体标题;变体0同步写回 titleSettings.title(全局样式面板/草稿保存依赖) */ + /** 更新单个变体标题;变体0同步写回 titleSettings.title(全局样式面板/草稿/TTS 链路依赖) */ const updateVariantTitle = (index: number, val: string) => { if (!previewTitles || !onPreviewTitlesChange) return const next = [...previewTitles] @@ -69,12 +105,43 @@ const Step4TitleSettings: React.FC = (props) => { } } + /** 批量 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 (

📝 选择标题

{!isBatch ? ( - /* ── 单视频:原有 AI 标题 + 输入框(保持不变) ── */ + /* ── 单视频:原有 AI 标题 + 输入框(保持不变,零回归) ── */ <> {t.titleSettings.aiAutoSelect ? ( <> @@ -144,7 +211,7 @@ const Step4TitleSettings: React.FC = (props) => { t.updateTitle(val || "") onPreviewTitlesChange?.([val || ""]) }} - options={t.userTitles.map((ut) => ({ label: ut.content, value: ut.content }))} + options={titleOptions} filterOption={(inputValue, option) => { const title = (option?.label || option?.value || "") as string return title.toLowerCase().includes((inputValue || "").toLowerCase()) @@ -155,7 +222,7 @@ const Step4TitleSettings: React.FC = (props) => { )} ) : ( - /* ── 批量:N 个独立标题输入框(CSS 浮层实时叠加到对应预览) ── */ + /* ── 批量:AI 批量生成 + N 个独立标题输入框(AutoComplete 支持标题库) ── */
= (props) => { > 为每个视频输入独立标题,修改会实时叠加到左侧对应视频上。标题样式(字体/颜色/位置)全局统一。
+ + {/* 批量 AI 标题 */} +
+ { + setBatchAiTopic(e.target.value) + t.setAiTitleInput(e.target.value) + }} + maxLength={30} + size="small" + style={{ flex: 1 }} + /> + + +
+ {Array.from({ length: previewCount }, (_, i) => (
- updateVariantTitle(i, e.target.value)} + style={{ width: "100%" }} + value={previewTitles?.[i] || undefined} + onChange={(val) => updateVariantTitle(i, val || "")} + options={titleOptions} + filterOption={(inputValue, option) => { + const title = (option?.label || option?.value || "") as string + return title.toLowerCase().includes((inputValue || "").toLowerCase()) + }} />
))} diff --git a/apps/web/src/pages/generate/constants.ts b/apps/web/src/pages/generate/constants.ts index a63fb827d..c0c37d3bc 100644 --- a/apps/web/src/pages/generate/constants.ts +++ b/apps/web/src/pages/generate/constants.ts @@ -33,7 +33,8 @@ export const STEPS = [ { key: 2, label: "选择素材" }, { key: 3, label: "选择配音" }, { key: 4, label: "选择标题" }, - { key: 5, label: "选择封面" }, + { key: 5, label: "确认生成" }, + { key: 6, label: "选择封面" }, ] /* ── 批量生成限制 ── */ diff --git a/apps/web/src/pages/generate/generate.css b/apps/web/src/pages/generate/generate.css index 192755553..08e17899b 100644 --- a/apps/web/src/pages/generate/generate.css +++ b/apps/web/src/pages/generate/generate.css @@ -3278,3 +3278,167 @@ max-height: none; } } + +/* ============================================================ + 批量前端 Canvas 预览网格(Issue #1677 修正:纯前端实时预览) + ============================================================ */ +.xx-canvas-grid { + display: grid; + grid-template-columns: repeat(2, 1fr); + gap: 16px; +} + +.xx-canvas-grid-card { + border: 2px solid var(--border-primary, #e2e8f0); + border-radius: 12px; + overflow: hidden; + background: #000; + transition: border-color 0.2s ease; + min-width: 0; +} + +.xx-canvas-grid-card.selected { + border-color: var(--primary-color, #1677ff); + box-shadow: 0 0 0 2px rgba(22, 119, 255, 0.15); +} + +.xx-canvas-grid-card-bar { + position: relative; + z-index: 2; + display: flex; + align-items: center; + padding: 6px 10px; + background: var(--bg-surface, #fff); + border-bottom: 1px solid var(--border-primary, #e2e8f0); +} + +.xx-canvas-grid-check { + display: inline-flex; + align-items: center; + gap: 6px; + font-size: 13px; + font-weight: 500; + color: var(--text-primary, #1a1a1a); + cursor: pointer; + user-select: none; +} + +.xx-canvas-grid-check input[type="checkbox"] { + width: 15px; + height: 15px; + cursor: pointer; + accent-color: var(--primary-color, #1677ff); +} + +/* ============================================================ + 批量标题:AI 一键生成行(Issue #1677) + ============================================================ */ +.xx-batch-ai-row { + display: flex; + flex-wrap: wrap; + align-items: center; + gap: 8px; + padding: 10px 12px; + margin-bottom: 12px; + background: var(--bg-secondary, #f7f8fa); + border: 1px dashed var(--border-primary, #d9d9d9); + border-radius: 10px; +} + +.xx-batch-ai-row .xx-form-field { + margin: 0; + flex: 1; + min-width: 140px; +} + +.xx-batch-titles { + display: flex; + flex-direction: column; + gap: 10px; +} + +/* ============================================================ + 第5步确认生成:批量渲染进度网格(Issue #1677) + ============================================================ */ +.xx-batch-gen-grid { + display: grid; + grid-template-columns: repeat(2, 1fr); + gap: 16px; +} + +.xx-batch-gen-card { + border: 1px solid var(--border-primary, #e2e8f0); + border-radius: 12px; + padding: 14px; + background: var(--bg-surface, #fff); + display: flex; + flex-direction: column; + gap: 10px; + min-width: 0; +} + +.xx-batch-gen-card.status-completed { + border-color: rgba(82, 196, 26, 0.4); + background: rgba(82, 196, 26, 0.04); +} + +.xx-batch-gen-card.status-failed { + border-color: rgba(239, 68, 68, 0.4); + background: rgba(239, 68, 68, 0.04); +} + +.xx-batch-gen-card-head { + display: flex; + align-items: center; + justify-content: space-between; + gap: 8px; +} + +.xx-batch-gen-card-title { + font-size: 14px; + font-weight: 600; + color: var(--text-primary, #1a1a1a); + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; +} + +.xx-batch-gen-card-body { + display: flex; + flex-direction: column; + gap: 8px; +} + +.xx-batch-gen-card-pct { + font-size: 13px; + color: var(--text-secondary, #666); + text-align: right; +} + +.xx-batch-gen-card-done { + font-size: 13px; + color: var(--success-color, #52c41a); + padding: 8px 0; +} + +.xx-batch-gen-card-failed { + display: flex; + flex-direction: column; + gap: 8px; + align-items: flex-start; +} + +.xx-batch-gen-card-err { + font-size: 13px; + color: var(--error-color, #ef4444); + line-height: 1.5; + word-break: break-word; +} + +/* ── 响应式:窄屏批量网格回退单列 ── */ +@media (max-width: 960px) { + .xx-canvas-grid, + .xx-batch-gen-grid { + grid-template-columns: 1fr; + } +} diff --git a/apps/web/src/pages/generate/hooks/generate-video/useGenerationPolling.ts b/apps/web/src/pages/generate/hooks/generate-video/useGenerationPolling.ts index 1661d9f81..cfa5be4fb 100644 --- a/apps/web/src/pages/generate/hooks/generate-video/useGenerationPolling.ts +++ b/apps/web/src/pages/generate/hooks/generate-video/useGenerationPolling.ts @@ -1,14 +1,28 @@ -import { useRef, useCallback } from "react" +import { useRef, useCallback, useState } from "react" import { message } from "antd" import axios from "axios" -import { getGenerationTask } from "@/api/tasks/tasks" +import { getGenerationTask, retryTask as retryGenerationTaskApi } from "@/api/tasks/tasks" import { getGenerationTaskResults } from "@/api/template-editor" import { safeExtractError } from "./errorUtils" +/** 批量生成时单个任务的实时状态(Issue #1677 第5步确认生成页) */ +export interface BatchTaskState { + taskId: string + /** 变体序号(0-based,与标题/封面数组对齐) */ + variantIndex: number + status: "running" | "completed" | "failed" + progress: number + error: string | null + /** 完成后的成片视频 */ + videos: unknown[] +} + interface UseGenerationPollingOptions { onProgress: (progress: number) => void onComplete: (videos: unknown[]) => void onFailed: (errorMsg: string) => void + /** 批量:单任务状态变化(第5步逐卡片展示) */ + onBatchTaskUpdate?: (taskId: string, patch: Partial) => void } /** 最大连续错误次数(仅对可重试错误),超过后终止轮询 */ @@ -17,20 +31,25 @@ const MAX_RETRYABLE_ERRORS = 10 const MAX_RESULTS_RETRIES = 3 /** - * 生成状态轮询 Hook(v3 — 支持批量多任务) + * 生成状态轮询 Hook(v4 — 批量任务独立状态 + 单任务重试) * * startPolling(taskId) 轮询单个任务; - * startPollingBatch(taskIds) 并行轮询 N 个任务,全部完成后聚合结果, - * 任一任务失败即整体失败(其余任务仍在后端继续,不影响)。 - * 进度为所有任务平均值。 + * startPollingBatch(tasks) 并行轮询 N 个任务: + * - 每个任务独立进度/状态/失败,通过 onBatchTaskUpdate 实时回传 + * - 全部成功才 onComplete(聚合视频按变体顺序);任一失败不影响其他任务继续 + * - retryTask(taskId) 单独重试失败任务(重新轮询,后端任务仍在跑则直接接续) */ export function useGenerationPolling({ onProgress, onComplete, onFailed, + onBatchTaskUpdate, }: UseGenerationPollingOptions) { const progressTimer = useRef[]>([]) const cancelledRef = useRef(false) + /** 批量任务上下文:taskId → 变体序号 */ + const batchContextRef = useRef>(new Map()) + const [, forceTick] = useState(0) const clearTimer = useCallback(() => { cancelledRef.current = true @@ -66,9 +85,21 @@ export function useGenerationPolling({ return safeExtractError(msg) } - /** 轮询单个任务,resolve 该任务的结果视频数组;失败时 reject(new Error(msg)) */ + /** + * 轮询单个任务。 + * - isBatch=true:状态变化通过 onBatchTaskUpdate 回传,不触发整体 onProgress/onComplete + * - resolve(videos) 成功;reject(Error) 失败 + */ const pollSingleTask = useCallback( - (taskId: string, runId: number, onTaskProgress?: (pct: number) => void): Promise => { + ( + taskId: string, + runId: number, + callbacks?: { + onTaskProgress?: (pct: number) => void + onTaskCompleted?: (videos: unknown[]) => void + onTaskFailed?: (msg: string) => void + }, + ): Promise => { return new Promise((resolve, reject) => { let consecutiveErrors = 0 let done = false @@ -85,9 +116,12 @@ export function useGenerationPolling({ const videos = await fetchResultsWithRetry(taskId) if (cancelledRef.current) return if (videos === null) { - reject(new Error("视频已生成,但获取结果列表失败,请稍后在任务列表查看")) + const msg = "视频已生成,但获取结果列表失败,请稍后在任务列表查看" + callbacks?.onTaskFailed?.(msg) + reject(new Error(msg)) return } + callbacks?.onTaskCompleted?.(videos) resolve(videos) return } @@ -98,14 +132,15 @@ export function useGenerationPolling({ task.error_info?.error_message || task.error_message || (task.status === "cancelled" ? "任务已取消" : "视频生成失败,请联系管理员或重试") - reject(new Error(safeExtractError(rawMsg))) + const msg = safeExtractError(rawMsg) + callbacks?.onTaskFailed?.(msg) + reject(new Error(msg)) return } const pct = Math.max(0, Math.min(99, Math.round(Number(task.progress) || 0))) - if (onTaskProgress) { - onTaskProgress(pct) - } else if (runId === 0) { + callbacks?.onTaskProgress?.(pct) + if (!callbacks && runId === 0) { onProgress(pct) } const timer = setTimeout(poll, 2000) @@ -116,13 +151,17 @@ export function useGenerationPolling({ const status = axios.isAxiosError(pollErr) ? pollErr.response?.status : undefined if (status && status >= 400 && status < 500) { done = true - reject(new Error(extractErrorMessage(pollErr, status))) + const msg = extractErrorMessage(pollErr, status) + callbacks?.onTaskFailed?.(msg) + reject(new Error(msg)) return } consecutiveErrors += 1 if (consecutiveErrors >= MAX_RETRYABLE_ERRORS) { done = true - reject(new Error("任务状态查询连续失败,请稍后在任务列表查看结果")) + const msg = "任务状态查询连续失败,请稍后在任务列表查看结果" + callbacks?.onTaskFailed?.(msg) + reject(new Error(msg)) return } const timer = setTimeout(poll, 3000) @@ -137,12 +176,12 @@ export function useGenerationPolling({ [onProgress, fetchResultsWithRetry], ) - /** 单任务轮询(兼容旧调用) */ + /** 单任务轮询(单视频,兼容旧调用) */ const startPolling = useCallback( (taskId: string) => { cancelledRef.current = false - const runId = 0 - pollSingleTask(taskId, runId) + batchContextRef.current.clear() + pollSingleTask(taskId, 0) .then((videos) => { if (cancelledRef.current) return onProgress(100) @@ -159,48 +198,114 @@ export function useGenerationPolling({ [pollSingleTask, onProgress, onComplete, onFailed], ) - /** 批量多任务轮询:全部完成后聚合结果;任一失败即整体失败 */ + /** + * 批量多任务轮询: + * - 每个任务独立进度/状态回传 onBatchTaskUpdate + * * 全部完成后按变体顺序聚合视频 onComplete + * - 部分失败:整体不 onFailed(第5步逐卡片展示失败+重试按钮);全部失败才 onFailed + */ const startPollingBatch = useCallback( - (taskIds: string[]) => { + (tasks: { taskId: string; variantIndex: number }[]) => { cancelledRef.current = false const runId = Date.now() const progressMap = new Map() + const resultMap = new Map() + const failureMap = new Map() + batchContextRef.current = new Map(tasks.map((t) => [t.taskId, t.variantIndex])) const reportAggregateProgress = () => { if (cancelledRef.current) return - const values = taskIds.map((id) => progressMap.get(id) ?? 0) + const values = tasks.map((t) => progressMap.get(t.taskId) ?? 0) const avg = Math.round(values.reduce((a, b) => a + b, 0) / Math.max(values.length, 1)) onProgress(Math.min(avg, 99)) } - const tasks = taskIds.map((taskId) => - pollSingleTask(taskId, runId, (pct) => { - progressMap.set(taskId, pct) - reportAggregateProgress() - }).then((videos) => { - progressMap.set(taskId, 100) - reportAggregateProgress() - return videos - }), - ) - - Promise.all(tasks) - .then((results) => { - if (cancelledRef.current) return + const checkAllSettled = () => { + if (resultMap.size + failureMap.size < tasks.length) return + if (resultMap.size === tasks.length) { onProgress(100) - const allVideos = results.flat() - onComplete(allVideos) - message.success(`全部 ${taskIds.length} 个视频生成完成!`) + const ordered = tasks.map((t) => resultMap.get(t.taskId) || []).flat() + onComplete(ordered) + message.success(`全部 ${tasks.length} 个视频生成完成!`) + } else if (resultMap.size > 0) { + // 部分失败:成功的视频聚合进成片列表(可进封面),失败卡片带重试按钮 + onProgress(100) + const ordered = tasks + .filter((t) => resultMap.has(t.taskId)) + .map((t) => resultMap.get(t.taskId) || []) + .flat() + onComplete(ordered) + message.warning( + `${failureMap.size} 个视频生成失败,可点击卡片上的「重试此视频」,成功的视频可先进入下一步`, + ) + } else { + const firstMsg = failureMap.get(tasks[0].taskId) || "全部视频生成失败" + onFailed(firstMsg) + } + } + + tasks.forEach(({ taskId, variantIndex }) => { + onBatchTaskUpdate?.(taskId, { + taskId, + variantIndex, + status: "running", + progress: 0, + error: null, + videos: [], }) - .catch((err: Error) => { - if (cancelledRef.current) return - console.error("[批量生成失败]", err.message) - onFailed(err.message) - message.error(err.message) + pollSingleTask(taskId, runId, { + onTaskProgress: (pct) => { + progressMap.set(taskId, pct) + onBatchTaskUpdate?.(taskId, { status: "running", progress: pct }) + reportAggregateProgress() + }, + onTaskCompleted: (videos) => { + progressMap.set(taskId, 100) + resultMap.set(taskId, videos) + onBatchTaskUpdate?.(taskId, { status: "completed", progress: 100, videos }) + reportAggregateProgress() + checkAllSettled() + }, + onTaskFailed: (msg) => { + failureMap.set(taskId, msg) + onBatchTaskUpdate?.(taskId, { status: "failed", error: msg }) + checkAllSettled() + }, + }).catch(() => { + // 失败已在 onTaskFailed 处理,这里吞掉 Promise rejection }) + }) }, - [pollSingleTask, onProgress, onComplete, onFailed], + [pollSingleTask, onProgress, onComplete, onFailed, onBatchTaskUpdate], ) - return { startPolling, startPollingBatch, clearTimer } + /** 单独重试失败任务(第5步卡片「重试此视频」):先调后端重试接口,再轮询 */ + const retryTask = useCallback( + async (taskId: string) => { + if (cancelledRef.current) cancelledRef.current = false + const variantIndex = batchContextRef.current.get(taskId) ?? 0 + onBatchTaskUpdate?.(taskId, { status: "running", progress: 0, error: null, videos: [] }) + try { + await retryGenerationTaskApi(taskId) + } catch (err) { + // 后端不支持重试或任务不可重试:直接重新轮询(任务可能已被自动恢复) + console.warn("[重试任务接口调用失败,改为直接轮询]", err) + } + pollSingleTask(taskId, Date.now(), { + onTaskProgress: (pct) => onBatchTaskUpdate?.(taskId, { status: "running", progress: pct }), + onTaskCompleted: (videos) => { + onBatchTaskUpdate?.(taskId, { status: "completed", progress: 100, videos }) + message.success(`视频 ${variantIndex + 1} 重试成功`) + }, + onTaskFailed: (msg) => onBatchTaskUpdate?.(taskId, { status: "failed", error: msg }), + }).catch(() => { + /* 失败已在回调处理 */ + }) + forceTick((n) => n + 1) + return variantIndex + }, + [pollSingleTask, onBatchTaskUpdate], + ) + + return { startPolling, startPollingBatch, retryTask, clearTimer } } diff --git a/apps/web/src/pages/generate/hooks/useBatchPreview.ts b/apps/web/src/pages/generate/hooks/useBatchPreview.ts deleted file mode 100644 index e8aacae52..000000000 --- a/apps/web/src/pages/generate/hooks/useBatchPreview.ts +++ /dev/null @@ -1,285 +0,0 @@ -/** - * 批量服务器预览 Hook(Issue #1677 多视频批量生成) - * - * 核心职责: - * 1. 调用 POST /generation/preview(preview_count=N)一次创建 N 个独立变体任务 - * 2. 对每个变体 task_id 分别轮询 GET /generation/preview/{task_id} - * 3. 返回每个变体的状态/进度/视频URL,供网格播放器展示 - * - * N=1 时不启用(走前端 Canvas 实时预览,零回归); - * N>1 时进入标题页自动触发;素材/配音等配置变化后重新触发。 - */ -import { useState, useCallback, useRef, useEffect } from "react" -import { createPreview, getPreviewStatus } from "@/api/generation/preview" -import type { CreatePreviewRequest } from "@/api/generation/types" - -export type VariantPreviewStatus = "loading" | "ready" | "failed" - -export interface VariantPreview { - /** 变体序号(0-based) */ - index: number - taskId: string - status: VariantPreviewStatus - progress: number - videoUrl: string | null - error: string | null -} - -interface UseBatchPreviewOptions { - /** 是否启用(仅 previewCount>1 且在标题页时启用) */ - enabled: boolean - /** 构建预览请求参数(每次触发时调用,获取最新配置) */ - buildRequest: () => CreatePreviewRequest - /** 批量预览任务创建成功回调(回传变体 taskId 列表与 source_edit_plan_id) */ - onPreviewTasksCreated?: (taskIds: string[], sourceEditPlanId?: string) => void -} - -interface UseBatchPreviewReturn { - variants: VariantPreview[] - /** 整体状态:loading=任一进行中,ready=全部完成,failed=有失败 */ - status: "idle" | "loading" | "ready" | "partial_failed" | "failed" - /** 总进度 0-100(各变体平均值) */ - progress: number - /** 失败的变体数量 */ - failedCount: number - /** 手动重新触发 */ - trigger: () => void -} - -const POLL_INTERVAL = 2000 -const POLL_TIMEOUT = 180_000 -const MAX_NETWORK_RETRIES = 2 - -/** - * 对配置参数做指纹,用于检测配置是否变化(标题文字/样式变化不触发重渲染,仅CSS浮层叠加) - */ -function buildFingerprint(req: CreatePreviewRequest): string { - // 不含 titles/title_config:标题文字与样式由 CSS 浮层实时叠加,变化不触发重渲染 - return JSON.stringify({ - t: req.template_id, - a: [...(req.asset_ids || [])].sort(), - r: req.video_ratio, - v: req.voice_library_id, - vs: req.voice_library_ids, - pc: req.preview_count, - b: req.bgm_config, - }) -} - -export function useBatchPreview({ - enabled, - buildRequest, - onPreviewTasksCreated, -}: UseBatchPreviewOptions): UseBatchPreviewReturn { - const [variants, setVariants] = useState([]) - const [status, setStatus] = useState<"idle" | "loading" | "ready" | "partial_failed" | "failed">( - "idle", - ) - const requestSeqRef = useRef(0) - const pollTimersRef = useRef[]>([]) - const timeoutTimerRef = useRef | null>(null) - const mountedRef = useRef(true) - - const buildRequestRef = useRef(buildRequest) - buildRequestRef.current = buildRequest - const onCreatedRef = useRef(onPreviewTasksCreated) - onCreatedRef.current = onPreviewTasksCreated - - const clearTimers = useCallback(() => { - pollTimersRef.current.forEach((t) => clearTimeout(t)) - pollTimersRef.current = [] - if (timeoutTimerRef.current) { - clearTimeout(timeoutTimerRef.current) - timeoutTimerRef.current = null - } - }, []) - - useEffect(() => { - mountedRef.current = true - return () => { - mountedRef.current = false - clearTimers() - } - }, [clearTimers]) - - /** 更新单个变体状态 */ - const patchVariant = useCallback((taskId: string, patch: Partial) => { - setVariants((prev) => prev.map((v) => (v.taskId === taskId ? { ...v, ...patch } : v))) - }, []) - - /** 轮询单个变体任务 */ - const pollVariant = useCallback( - async (taskId: string, seq: number, retries = 0) => { - if (seq !== requestSeqRef.current || !mountedRef.current) return - try { - const st = await getPreviewStatus(taskId) - if (seq !== requestSeqRef.current || !mountedRef.current) return - - if (st.status === "completed" && st.video_url) { - patchVariant(taskId, { - status: "ready", - videoUrl: st.video_url, - progress: 100, - error: null, - }) - return - } - if (st.status === "failed" || st.status === "cancelled") { - patchVariant(taskId, { - status: "failed", - error: - st.status === "cancelled" ? "预览任务已取消" : st.error_message || "预览渲染失败", - }) - return - } - if (typeof st.progress === "number") { - patchVariant(taskId, { progress: Math.round(st.progress) }) - } - const timer = setTimeout(() => pollVariant(taskId, seq), POLL_INTERVAL) - pollTimersRef.current.push(timer) - } catch (err) { - if (seq !== requestSeqRef.current || !mountedRef.current) return - if (retries < MAX_NETWORK_RETRIES) { - console.warn(`[BatchPreview] 变体 ${taskId} 轮询网络错误,第 ${retries + 1} 次重试`, err) - const timer = setTimeout(() => pollVariant(taskId, seq, retries + 1), POLL_INTERVAL * 2) - pollTimersRef.current.push(timer) - } else { - patchVariant(taskId, { status: "failed", error: "网络错误,无法获取预览状态" }) - } - } - }, - [patchVariant], - ) - - /** 创建批量预览任务并开始轮询 */ - const trigger = useCallback(() => { - if (!enabled) return - const request = buildRequestRef.current() - if (!request.template_id || !request.asset_ids?.length) return - const count = request.preview_count && request.preview_count > 1 ? request.preview_count : 0 - if (!count) return - - clearTimers() - const seq = ++requestSeqRef.current - setStatus("loading") - setVariants( - Array.from({ length: count }, (_, i) => ({ - index: i, - taskId: "", - status: "loading" as const, - progress: 0, - videoUrl: null, - error: null, - })), - ) - - createPreview(request) - .then((resp) => { - if (seq !== requestSeqRef.current || !mountedRef.current) return - const items = resp.items || [] - const taskIds = items.map((it) => it.task_id).filter(Boolean) - if (taskIds.length === 0) { - setStatus("failed") - setVariants((prev) => - prev.map((v) => ({ ...v, status: "failed", error: "未创建预览任务" })), - ) - return - } - onCreatedRef.current?.(taskIds, resp.source_edit_plan_id) - - // 用返回的 task_id 填充变体(按 variant_index 对齐) - setVariants((prev) => - prev.map((v) => { - const item = items.find((it) => it.variant_index === v.index) || items[v.index] - return item ? { ...v, taskId: item.task_id } : v - }), - ) - - // 超时保护 - timeoutTimerRef.current = setTimeout(() => { - if (seq !== requestSeqRef.current || !mountedRef.current) return - setVariants((prev) => - prev.map((v) => - v.status === "loading" - ? { ...v, status: "failed", error: "预览渲染超时,请重试" } - : v, - ), - ) - }, POLL_TIMEOUT) - - // 分别轮询每个变体 - items.forEach((item) => { - if (item.task_id) pollVariant(item.task_id, seq) - }) - }) - .catch((err: unknown) => { - if (seq !== requestSeqRef.current || !mountedRef.current) return - console.error("[BatchPreview] 创建批量预览失败:", err) - const errData = (err as { response?: { data?: { detail?: string; message?: string } } }) - ?.response?.data - setStatus("failed") - setVariants((prev) => - prev.map((v) => ({ - ...v, - status: "failed", - error: errData?.detail || errData?.message || "预览任务创建失败,请重试", - })), - ) - }) - }, [enabled, clearTimers, pollVariant]) - - /* ── 自动触发 + 配置变更检测 ── */ - const request = enabled ? buildRequest() : null - const currentFingerprint = request - ? request.template_id && request.asset_ids?.length && (request.preview_count || 1) > 1 - ? buildFingerprint(request) - : "" - : "" - - const didInitRef = useRef(false) - useEffect(() => { - if (!enabled || !currentFingerprint) { - didInitRef.current = false - requestSeqRef.current += 1 - clearTimers() - setStatus("idle") - setVariants([]) - return - } - if (!didInitRef.current) { - didInitRef.current = true - trigger() - } - }, [enabled, currentFingerprint, trigger, clearTimers]) - - // 配置变更(素材/配音/数量)→ 重新渲染;标题文字变化不触发(CSS浮层实时叠加) - const prevFingerprintRef = useRef(currentFingerprint) - useEffect(() => { - if (!enabled || !currentFingerprint) return - const prev = prevFingerprintRef.current - prevFingerprintRef.current = currentFingerprint - if (!prev || prev === currentFingerprint) return - trigger() - }, [enabled, currentFingerprint, trigger]) - - /* ── 派生状态 ── */ - const progress = - variants.length > 0 - ? Math.round(variants.reduce((sum, v) => sum + v.progress, 0) / variants.length) - : 0 - const failedCount = variants.filter((v) => v.status === "failed").length - const readyCount = variants.filter((v) => v.status === "ready").length - - useEffect(() => { - if (status !== "loading" || variants.length === 0) return - if (readyCount === variants.length) { - setStatus("ready") - } else if (readyCount + failedCount === variants.length && failedCount > 0) { - setStatus(failedCount === variants.length ? "failed" : "partial_failed") - } - }, [variants, status, readyCount, failedCount]) - - return { variants, status, progress, failedCount, trigger } -} - -export default useBatchPreview diff --git a/apps/web/src/pages/generate/hooks/useGenerateVideo.ts b/apps/web/src/pages/generate/hooks/useGenerateVideo.ts index 6f77188ee..2d7a6d0d9 100755 --- a/apps/web/src/pages/generate/hooks/useGenerateVideo.ts +++ b/apps/web/src/pages/generate/hooks/useGenerateVideo.ts @@ -2,13 +2,13 @@ * 视频生成 Hook * 封装视频生成的核心逻辑、状态管理、轮询等 */ -import { useState, useCallback } from "react" +import { useState, useCallback, useEffect } from "react" import { message } from "antd" import { type GeneratedVideo, getEditPlanClips, createClipsFromAssets } from "@/api/template-editor" import { createGenerationTask } from "@/api/tasks/tasks" import type { UseGenerateVideoProps } from "./generate-video/types" import { getGenerationPhase } from "./generate-video/phase" -import { useGenerationPolling } from "./generate-video/useGenerationPolling" +import { useGenerationPolling, type BatchTaskState } from "./generate-video/useGenerationPolling" import { validateGenerateInputs } from "./generate-video/buildPayload" import { calculateResolution } from "../utils/calculateResolution" import { extractBackendError, translateError } from "./generate-video/errorUtils" @@ -22,6 +22,32 @@ export function useGenerateVideo(props: UseGenerateVideoProps) { const [generated, setGenerated] = useState(false) const [generateError, setGenerateError] = useState(null) const [generatedVideos, setGeneratedVideos] = useState([]) + /** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */ + const [batchTasks, setBatchTasks] = useState([]) + + const handleBatchTaskUpdate = useCallback((taskId: string, patch: Partial) => { + setBatchTasks((prev) => { + const list = prev || [] + const idx = list.findIndex((t) => t.taskId === taskId) + if (idx === -1) { + return [ + ...list, + { + taskId, + variantIndex: patch.variantIndex ?? 0, + status: "running", + progress: 0, + error: null, + videos: [], + ...patch, + }, + ] + } + const next = [...list] + next[idx] = { ...next[idx], ...patch } + return next + }) + }, []) const handleProgress = useCallback((p: number) => setProgress(p), []) const handleComplete = useCallback( @@ -29,6 +55,19 @@ export function useGenerateVideo(props: UseGenerateVideoProps) { setGenerating(false) setGenerated(true) setGeneratedVideos(videos as GeneratedVideo[]) + // 批量:成功任务的 videos 已通过 onBatchTaskUpdate 写入,这里同步兜底 + setBatchTasks((prev) => + (prev || []).map((t) => + t.status === "completed" && t.videos.length === 0 + ? { + ...t, + videos: (videos as GeneratedVideo[]).filter( + (v) => v.generation_task_id === t.taskId, + ), + } + : t, + ), + ) onGenerationSuccess?.() }, [onGenerationSuccess], @@ -38,10 +77,30 @@ export function useGenerateVideo(props: UseGenerateVideoProps) { setGenerateError(errorMsg) }, []) - const { startPolling, startPollingBatch, clearTimer } = useGenerationPolling({ + /* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入), + 按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */ + useEffect(() => { + if (batchTasks.length === 0) return + const byVariant = new Map() + batchTasks.forEach((t) => { + if (t.status === "completed" && t.videos && t.videos.length > 0) { + byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo) + } + }) + const ordered = [...byVariant.entries()].sort((a, b) => a[0] - b[0]).map(([, v]) => v) + setGeneratedVideos((prev) => { + if (prev.length === ordered.length && prev.every((v, i) => v.id === ordered[i].id)) { + return prev + } + return ordered + }) + }, [batchTasks]) + + const { startPolling, startPollingBatch, retryTask, clearTimer } = useGenerationPolling({ onProgress: handleProgress, onComplete: handleComplete, onFailed: handleFailed, + onBatchTaskUpdate: handleBatchTaskUpdate, }) /* ── 生成视频 ── @@ -57,6 +116,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) { setProgress(0) setGenerated(false) setGenerateError(null) + setBatchTasks([]) clearTimer() try { @@ -170,7 +230,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) { throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看") } if (taskIds.length > 1) { - startPollingBatch(taskIds) + // 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建) + startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i }))) } else { startPolling(taskIds[0]) } @@ -196,6 +257,14 @@ export function useGenerateVideo(props: UseGenerateVideoProps) { generate() }, [generate]) + /** 第5步:单独重试某个失败任务 */ + const retryBatchTask = useCallback( + (taskId: string) => { + retryTask(taskId) + }, + [retryTask], + ) + const dismissError = useCallback(() => { setGenerateError(null) }, []) @@ -240,6 +309,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) { generatedVideos, generate, retry, + retryBatchTask, + batchTasks, dismissError, download, share, diff --git a/apps/web/src/pages/generate/hooks/useStepNavigation.ts b/apps/web/src/pages/generate/hooks/useStepNavigation.ts index d7e59bd58..12e8c51fd 100644 --- a/apps/web/src/pages/generate/hooks/useStepNavigation.ts +++ b/apps/web/src/pages/generate/hooks/useStepNavigation.ts @@ -1,6 +1,11 @@ /** - * GeneratePage 步骤导航(Issue #1677 改造后 5 步) - * 步骤:模板(1) → 素材(2) → 配音(3) → 标题+预览+确认生成(4) → 封面(5) + * GeneratePage 步骤导航(Issue #1677 修正:固定 6 步,单视频与批量一致) + * 步骤:模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6) + * + * - 步骤4底部按钮是「确认生成视频/确认生成 N 个视频」(由 GenerateStepActions 调 + * onConfirmGenerate),创建成功后跳转步骤5;本 hook 的 goNext 只负责 1→2→3→4 + * 和 5→6 的「下一步」。 + * - 步骤5(确认生成进度页):渲染全部完成(generated)后「下一步」解锁进封面。 */ import { message } from "antd" import type { TitleSettings } from "../types" @@ -13,14 +18,8 @@ export interface UseStepNavigationOptions { selectedMaterials: string[] smartSelectedIds: string[] titleSettings: TitleSettings - /** 预览是否已就绪(单视频=前端预览素材已加载;批量=服务器预览全部完成) */ - previewReady: boolean - /** 是否已完成视频生成(步骤4确认生成后才能进入封面) */ + /** 是否已完成视频生成(步骤5全部渲染完成后才能进入封面) */ generated: boolean - /** 批量模式下每个变体的标题 */ - previewTitles: string[] - /** 批量模式勾选的变体数 */ - selectedCount: number /** Step1 点下一步时弹出数量选择弹窗 */ onOpenCountModal: () => void } @@ -38,10 +37,7 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav materialMode, selectedMaterials, smartSelectedIds, - previewReady, generated, - previewTitles, - selectedCount, onOpenCountModal, } = options @@ -63,27 +59,14 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav message.warning("请先进行智能匹配并选择素材") return } - // Step4(标题+预览+确认生成):标题必填 + 预览必须已加载 - if (currentStep === 4) { - const allTitlesFilled = previewTitles.every((t) => t && t.trim()) - if (!allTitlesFilled) { - message.warning("请为每个视频输入标题") - return - } - if (selectedCount === 0) { - message.warning("请至少勾选一个视频") - return - } - if (!previewReady) { - message.warning("预览视频正在加载,请稍候") - return - } + // 步骤5(确认生成):全部渲染完成后才能下一步进封面 + if (currentStep === 5) { if (!generated) { - message.warning("请先点击「确认生成视频」完成渲染") + message.warning("视频还在渲染中,请等待生成完成") return } } - if (currentStep < 5) { + if (currentStep < 6) { setCurrentStep((s) => s + 1) } } diff --git a/apps/web/src/test/pages/generate/smoke.test.tsx b/apps/web/src/test/pages/generate/smoke.test.tsx index 90da8ddc3..9e58e0c6b 100755 --- a/apps/web/src/test/pages/generate/smoke.test.tsx +++ b/apps/web/src/test/pages/generate/smoke.test.tsx @@ -20,7 +20,8 @@ import "@/pages/generate/components/Step2MaterialSelect" import "@/pages/generate/components/Step4TitleSettings" import "@/pages/generate/components/Step5VoiceSelect" import "@/pages/generate/components/Step3VoiceWithMode" -import "@/pages/generate/components/ServerPreviewGrid" +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" @@ -48,7 +49,6 @@ describe("GeneratePage module smoke test", () => { }) }) import "@/pages/generate/hooks/useGenerateVideo" -import "@/pages/generate/hooks/useBatchPreview" import "@/pages/generate/hooks/useBatchCovers" import "@/pages/generate/hooks/usePreviewAssets" import "@/pages/generate/hooks/useSegmentScheduler" -- 2.54.0 From fbfd19fbb963933e11b2b2b5f4d5bbe85062686a Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 12:56:10 +0800 Subject: [PATCH 11/33] =?UTF-8?q?fix(#1677):=20AI=20Review=20=E8=B7=9F?= =?UTF-8?q?=E8=BF=9B=E4=BF=AE=E5=A4=8D=EF=BC=88=E6=A0=87=E9=A2=98=E6=A0=A1?= =?UTF-8?q?=E9=AA=8C/TTS=E4=BE=9D=E8=B5=96/=E6=8E=92=E5=BA=8F=E5=85=9C?= =?UTF-8?q?=E5=BA=95/=E9=A2=84=E8=A7=88=E6=95=B0=E7=A1=AC=E4=B8=8A?= =?UTF-8?q?=E9=99=90=EF=BC=89=20(#1712)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/src/pages/generate/GeneratePage.tsx | 16 +++++++++++----- .../generate/components/BatchGenerationGrid.tsx | 2 +- .../generate/components/CanvasPreviewGrid.tsx | 5 ++++- 3 files changed, 16 insertions(+), 7 deletions(-) diff --git a/apps/web/src/pages/generate/GeneratePage.tsx b/apps/web/src/pages/generate/GeneratePage.tsx index ef48442d7..32fa3d68a 100644 --- a/apps/web/src/pages/generate/GeneratePage.tsx +++ b/apps/web/src/pages/generate/GeneratePage.tsx @@ -132,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) @@ -141,7 +143,7 @@ const GeneratePage: React.FC = () => { } // 批量模式下 TTS 文案跟随变体0标题;单视频跟随主标题 - const ttsTitle = isBatch ? previewTitles[0] || "" : titleSettings.title + const ttsTitle = isBatch ? variant0Title || "" : titleSettings.title const voiceId = selectedClonedVoice || selectedVoice if (!voiceId || !ttsTitle) { setPreviewVoiceAudioUrl(null) @@ -174,7 +176,7 @@ const GeneratePage: React.FC = () => { selectedVoice, selectedClonedVoice, titleSettings.title, - previewTitles, + variant0Title, isBatch, voiceMaterials, ]) @@ -321,9 +323,13 @@ const GeneratePage: React.FC = () => { message.warning("请为每个勾选的视频输入标题") return } - } else if (!titleSettings.aiAutoSelect && !titleSettings.title?.trim()) { - message.warning("请先选择或输入标题") - return + } else if (!titleSettings.title?.trim()) { + // 与 buildPayload.validateGenerateInputs 一致:AI 自动选标题模式(aiAutoSelect) + // 允许空标题由后端生成;手动模式必须填写,避免提交空标题 + if (!titleSettings.aiAutoSelect) { + message.warning("请先选择或输入标题") + return + } } if (!previewReady) { message.warning("预览素材正在加载,请稍候") diff --git a/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx b/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx index 3761a4a47..0c52e355d 100644 --- a/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx +++ b/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx @@ -22,7 +22,7 @@ const BatchGenerationGrid: React.FC = ({ titles, onRetryTask, }) => { - const sorted = [...tasks].sort((a, b) => a.variantIndex - b.variantIndex) + const sorted = [...tasks].sort((a, b) => (a.variantIndex || 0) - (b.variantIndex || 0)) return (
diff --git a/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx b/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx index 1d0b17674..a034e29b3 100644 --- a/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx +++ b/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx @@ -12,6 +12,7 @@ import type { AssetItem } from "@/api/assets" import type { EditingTemplate } from "@/api/editing-planner" import type { TitleSettings } from "../types" import FrontendPreviewPlayer from "./FrontendPreviewPlayer" +import { MAX_PREVIEW_COUNT } from "../constants" interface CanvasPreviewGridProps { count: number @@ -41,9 +42,11 @@ const CanvasPreviewGrid: React.FC = ({ onToggleSelect, selectable = true, }) => { + // 硬上限保护:同时播放的媒体元素数量不超过 MAX_PREVIEW_COUNT(10),避免浏览器卡顿 + const safeCount = Math.max(1, Math.min(count, MAX_PREVIEW_COUNT)) return (
- {Array.from({ length: count }, (_, i) => { + {Array.from({ length: safeCount }, (_, i) => { const checked = selectedIds.includes(i) return (
Date: Sat, 5 Sep 2026 13:27:46 +0800 Subject: [PATCH 12/33] =?UTF-8?q?fix(#1677):=20=E7=A7=BB=E9=99=A4=E9=A2=84?= =?UTF-8?q?=E8=A7=88=E7=BD=91=E6=A0=BC=E6=88=AA=E6=96=AD+=E7=A9=BA?= =?UTF-8?q?=E6=8C=87=E9=92=88=E9=98=B2=E5=BE=A1+=E6=8E=92=E5=BA=8F?= =?UTF-8?q?=E6=98=BE=E5=BC=8FNumber=EF=BC=88AI=20Review=20=E7=AC=AC?= =?UTF-8?q?=E4=BA=8C=E8=BD=AE=EF=BC=89=20(#1713)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/src/pages/generate/GeneratePage.tsx | 2 +- .../src/pages/generate/components/BatchGenerationGrid.tsx | 4 +++- .../src/pages/generate/components/CanvasPreviewGrid.tsx | 7 +++---- 3 files changed, 7 insertions(+), 6 deletions(-) diff --git a/apps/web/src/pages/generate/GeneratePage.tsx b/apps/web/src/pages/generate/GeneratePage.tsx index 32fa3d68a..a31a87f58 100644 --- a/apps/web/src/pages/generate/GeneratePage.tsx +++ b/apps/web/src/pages/generate/GeneratePage.tsx @@ -133,7 +133,7 @@ const GeneratePage: React.FC = () => { const [previewVoiceAudioUrl, setPreviewVoiceAudioUrl] = useState(null) const ttsAbortRef = useRef(null) // TTS 试听文案:批量跟随变体0标题(仅取首项,避免编辑其他变体标题触发多余 TTS 请求) - const variant0Title = isBatch ? previewTitles[0] || "" : "" + const variant0Title = isBatch ? previewTitles?.[0] || "" : "" useEffect(() => { const voiceAsset = voiceMaterials.find((m) => m.id === selectedVoice) diff --git a/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx b/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx index 0c52e355d..40cf7938a 100644 --- a/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx +++ b/apps/web/src/pages/generate/components/BatchGenerationGrid.tsx @@ -22,7 +22,9 @@ const BatchGenerationGrid: React.FC = ({ titles, onRetryTask, }) => { - const sorted = [...tasks].sort((a, b) => (a.variantIndex || 0) - (b.variantIndex || 0)) + const sorted = [...tasks].sort( + (a, b) => Number(a.variantIndex || 0) - Number(b.variantIndex || 0), + ) return (
diff --git a/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx b/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx index a034e29b3..4dd74d0f3 100644 --- a/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx +++ b/apps/web/src/pages/generate/components/CanvasPreviewGrid.tsx @@ -12,7 +12,6 @@ import type { AssetItem } from "@/api/assets" import type { EditingTemplate } from "@/api/editing-planner" import type { TitleSettings } from "../types" import FrontendPreviewPlayer from "./FrontendPreviewPlayer" -import { MAX_PREVIEW_COUNT } from "../constants" interface CanvasPreviewGridProps { count: number @@ -42,11 +41,11 @@ const CanvasPreviewGrid: React.FC = ({ onToggleSelect, selectable = true, }) => { - // 硬上限保护:同时播放的媒体元素数量不超过 MAX_PREVIEW_COUNT(10),避免浏览器卡顿 - const safeCount = Math.max(1, Math.min(count, MAX_PREVIEW_COUNT)) + // count 上限已在源头 PreviewCountModal 的数量选择(1~MAX_PREVIEW_COUNT=10)clamp, + // 这里完整渲染所有变体,保证每个变体都有勾选/预览入口,UI 与数据不脱节 return (
- {Array.from({ length: safeCount }, (_, i) => { + {Array.from({ length: count }, (_, i) => { const checked = selectedIds.includes(i) return (
Date: Sat, 5 Sep 2026 16:07:45 +0800 Subject: [PATCH 13/33] =?UTF-8?q?fix(upload):=20complete=20=E6=8E=A5?= =?UTF-8?q?=E5=8F=A3=E5=B9=82=E7=AD=89=E5=8E=BB=E9=87=8D=20+=20HEVC=20?= =?UTF-8?q?=E8=BD=AC=E7=A0=81=E5=9B=9E=E5=86=99=E5=8D=A0=E4=BD=8D=20asset?= =?UTF-8?q?=20=E7=A6=81=E6=AD=A2=E5=85=9C=E5=BA=95=E6=96=B0=E5=BB=BA=20(#1?= =?UTF-8?q?714)=20(#1715)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- ...et_client_upload_id_ingest_job_asset_id.py | 34 ++ apps/api/app/api/routes/upload.py | 203 ++++++++--- apps/api/app/schemas/upload.py | 12 +- apps/worker/worker_app/tasks/ingest.py | 152 ++++++-- .../adapters/in_memory/asset_repository.py | 41 +++ .../sqlalchemy_impl/asset_repository.py | 54 +++ .../sqlalchemy_impl/ingest_job_repository.py | 5 + packages/adapters/sqlalchemy_impl/models.py | 2 + packages/application/ingest_jobs.py | 2 + packages/domain/entities.py | 6 + packages/ports/asset_repository.py | 20 ++ tests/unit/test_ingest_hevc_orphan_1714.py | 317 ++++++++++++++++ .../test_upload_complete_idempotency_1714.py | 338 ++++++++++++++++++ 13 files changed, 1100 insertions(+), 86 deletions(-) create mode 100644 alembic/versions/066_asset_client_upload_id_ingest_job_asset_id.py create mode 100644 tests/unit/test_ingest_hevc_orphan_1714.py create mode 100644 tests/unit/test_upload_complete_idempotency_1714.py 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/apps/api/app/api/routes/upload.py b/apps/api/app/api/routes/upload.py index 58269c35c..f5b046f30 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,8 +108,85 @@ 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 +ACTIVE_ASSET_STATUSES = (AssetStatus.UPLOADING, AssetStatus.PROCESSING) + + +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. 兜底:同库 + 同文件名(+同大小)且 30 分钟内仍处 uploading/processing + 的记录——旧客户端不传 hash/token 时,防止 complete 超时重试反复建占位。 + + 全部为鸭子类型调用:旧仓储无对应方法时静默跳过,不破坏既有实现。 + """ + 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 + if filename: + 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 or 0, + ) + if existing is not None and getattr(existing, "status", None) in ACTIVE_ASSET_STATUSES: + logger.info( + "素材幂等兜底命中(近期活动同名记录): library=%s name=%s asset=%s status=%s", + library_id, + filename, + getattr(existing, "id", "?"), + getattr(existing, "status", "?"), + ) + return existing + 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="", ): """立即创建一条 PROCESSING 状态的 Asset 记录,使前端能马上看到新素材。""" asset = Asset.create( @@ -111,6 +198,7 @@ def _create_pending_asset( status=AssetStatus.PROCESSING, uploaded_by_user_id=user_id, file_hash=file_hash, + client_upload_id=client_upload_id, ) return asset_repository.create(asset) @@ -121,6 +209,7 @@ def _submit_ingest_job( 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,6 +218,7 @@ 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]) @@ -202,7 +292,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 +302,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 +336,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 +347,7 @@ 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, ) job = _submit_ingest_job( @@ -264,6 +356,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 +376,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 +388,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 +441,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 +450,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/schemas/upload.py b/apps/api/app/schemas/upload.py index c6d798288..bc606649c 100644 --- a/apps/api/app/schemas/upload.py +++ b/apps/api/app/schemas/upload.py @@ -31,14 +31,16 @@ 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 +48,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/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index 7b48f0f45..c42169b92 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,10 @@ def ingest_asset(job_id: str) -> dict: if job is None: return {"status": "failed", "error": "job not found"} + # 记录原始 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 +676,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 +727,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 +759,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 +816,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/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..ba334d498 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,53 @@ 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 历史素材不拦。 + """ + from datetime import datetime, timedelta, timezone + + if not name: + 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, + ) + if file_size and file_size > 0: + query = query.filter(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/ingest_job_repository.py b/packages/adapters/sqlalchemy_impl/ingest_job_repository.py index f16dc735d..c2e11c24b 100644 --- a/packages/adapters/sqlalchemy_impl/ingest_job_repository.py +++ b/packages/adapters/sqlalchemy_impl/ingest_job_repository.py @@ -18,6 +18,7 @@ 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 "", created_at=job.created_at, updated_at=job.updated_at, ) @@ -38,6 +39,7 @@ 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 "", created_at=model.created_at, updated_at=model.updated_at, ) @@ -54,6 +56,9 @@ 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 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..3bc1c8def 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -94,6 +94,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 +245,7 @@ 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) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) diff --git a/packages/application/ingest_jobs.py b/packages/application/ingest_jobs.py index 75a708de7..576a2b25c 100644 --- a/packages/application/ingest_jobs.py +++ b/packages/application/ingest_jobs.py @@ -12,6 +12,7 @@ class SubmitIngestJobCommand: library_id: str storage_key: str file_hash: str = "" + asset_id: str = "" class SubmitIngestJobUseCase: @@ -24,5 +25,6 @@ class SubmitIngestJobUseCase: library_id=command.library_id, storage_key=command.storage_key, file_hash=command.file_hash, + asset_id=command.asset_id, ) return self.ingest_job_repository.create(job) diff --git a/packages/domain/entities.py b/packages/domain/entities.py index 4df342e4d..9ed9e5b8a 100755 --- a/packages/domain/entities.py +++ b/packages/domain/entities.py @@ -174,6 +174,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 +209,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 +237,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 +269,7 @@ class IngestJob: error_message: str = "" result_asset_id: str = "" file_hash: str = "" + asset_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 +280,7 @@ class IngestJob: library_id: str, storage_key: str, file_hash: str = "", + asset_id: str = "", ) -> "IngestJob": if not project_id.strip(): raise ValueError("project_id 不能为空") @@ -289,4 +294,5 @@ class IngestJob: library_id=library_id.strip(), storage_key=storage_key.strip(), file_hash=file_hash.strip(), + asset_id=asset_id.strip(), ) 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/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_upload_complete_idempotency_1714.py b/tests/unit/test_upload_complete_idempotency_1714.py new file mode 100644 index 000000000..1effb6e78 --- /dev/null +++ b/tests/unit/test_upload_complete_idempotency_1714.py @@ -0,0 +1,338 @@ +"""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: + 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 (not file_size or 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) + r1 = client.post("/api/v1/direct/complete", json=COMPLETE_BODY) + 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": 0}, + ) + 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_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 -- 2.54.0 From 6521be5426e938cc6d1941d1b959bd9f23ec6e41 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 16:40:05 +0800 Subject: [PATCH 14/33] =?UTF-8?q?fix(#1714):=20=E7=B4=A0=E6=9D=90=E4=B8=8A?= =?UTF-8?q?=E4=BC=A0=E5=8E=BB=E9=87=8D=E9=98=B2=E9=87=8D=20+=20complete?= =?UTF-8?q?=E5=B9=82=E7=AD=89=E4=BC=A0=E9=80=92=20+=20=E8=BD=AE=E8=AF=A2?= =?UTF-8?q?=E6=94=B6=E6=95=9B=EF=BC=88=E5=89=8D=E7=AB=AFP0=EF=BC=89=20(#17?= =?UTF-8?q?17)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/src/api/assets/upload.ts | 38 ++- apps/web/src/api/assets/uploadDedup.ts | 143 ++++++++++++ apps/web/src/pages/assets/AssetLibrary.tsx | 4 + apps/web/src/pages/assets/assets.css | 22 ++ .../src/pages/assets/components/AssetCard.tsx | 17 +- .../assets/components/AssetGridSection.tsx | 4 + .../assets/components/AssetUploadZone.tsx | 21 +- .../assets/components/UploadQueuePanel.tsx | 5 +- .../src/pages/assets/hooks/useAssetUpload.ts | 219 +++++++++++++++--- .../src/pages/assets/hooks/useAssetsData.ts | 49 +++- apps/web/src/test/api/uploadDedup.test.ts | 101 ++++++++ .../test/pages/assets/useAssetUpload.test.tsx | 163 +++++++++++-- 12 files changed, 720 insertions(+), 66 deletions(-) create mode 100644 apps/web/src/api/assets/uploadDedup.ts create mode 100644 apps/web/src/test/api/uploadDedup.test.ts diff --git a/apps/web/src/api/assets/upload.ts b/apps/web/src/api/assets/upload.ts index e396b5d77..ab450cd69 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,14 @@ 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 }): 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,6 +121,10 @@ 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() @@ -118,6 +134,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 +146,8 @@ 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, }), } } @@ -137,8 +157,20 @@ 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, + }) 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..67f31f7de --- /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) + */ + +/** 大文件抽样阈值:超过此大小只哈希头尾片段,避免上传前长时间卡 UI */ +export const HASH_FULL_READ_LIMIT = 256 * 1024 * 1024 // 256MB +/** 抽样读取的头尾片段大小(各 8MB) */ +export const HASH_SAMPLE_CHUNK = 8 * 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 字段长度一致)。 + * - ≤256MB:全量哈希,内容一致必然一致 + * - >256MB:哈希「头部 8MB + 尾部 8MB + 文件大小」,视频素材体积大、 + * 头部含 moov 元数据、尾部含 mdat 结尾,抽样碰撞概率可忽略, + * 且避免上传前对 2GB 文件全量读取造成长时间卡顿 + * + * 运行环境不支持 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/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..1c143fb79 100644 --- a/apps/web/src/pages/assets/assets.css +++ b/apps/web/src/pages/assets/assets.css @@ -1041,3 +1041,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..a88f4f9eb 100644 --- a/apps/web/src/pages/assets/components/UploadQueuePanel.tsx +++ b/apps/web/src/pages/assets/components/UploadQueuePanel.tsx @@ -80,6 +80,7 @@ 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}` : ""}
@@ -89,7 +90,9 @@ const UploadQueuePanel: React.FC = ({ +
+
+ ) + } + + return ( +
+
+ +

正在绑定微信...

+
+
+ ) +} + +export default WechatBindCallback diff --git a/apps/web/src/pages/auth/WechatCallback.tsx b/apps/web/src/pages/auth/WechatCallback.tsx index 3dc75a6e9..4186c3033 100644 --- a/apps/web/src/pages/auth/WechatCallback.tsx +++ b/apps/web/src/pages/auth/WechatCallback.tsx @@ -1,19 +1,20 @@ /** * 微信登录回调页 + * 扫码授权后由微信重定向回来:用 code 换登录态, + * 新用户/资料未完善 → 跳昵称引导页;老用户 → 回来源页/首页 */ 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 { useAuthStore } from "@/store/authStore" -import BindContactModal from "@/components/auth/BindContactModal" +import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh" 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) useEffect(() => { @@ -49,20 +50,21 @@ 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 (needOnboarding) { + navigate("/welcome/wechat", { replace: true }) + return } - } catch (err) { - setError("登录失败,请重试") + + // 老用户:回登录前页面或首页 + const redirect = localStorage.getItem("login_redirect") || "/" + localStorage.removeItem("login_redirect") + navigate(redirect, { replace: true }) + } catch { + setError("微信登录失败,请重试") setLoading(false) } } @@ -70,21 +72,6 @@ const WechatCallback: React.FC = () => { 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") - } - if (loading) { return (
{ >
-

正在登录...

-
-
- ) - } - - 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..a4893b846 --- /dev/null +++ b/apps/web/src/pages/auth/WechatOnboarding.tsx @@ -0,0 +1,102 @@ +/** + * 微信新用户昵称引导页 + * 新微信用户首次登录后强制填写昵称,完成后才进入主界面 + */ +import React 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 { 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() + + 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) => { + 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 { + message.error("保存失败,请重试") + } + } + + return ( +
+
+
+
+ 🦐 + 小虾智剪 +
+

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

+
+ +
+ + + + + + + +
+
+
+ ) +} + +export default WechatOnboarding diff --git a/apps/web/src/pages/profile/ProfileSettings.css b/apps/web/src/pages/profile/ProfileSettings.css index de21e497b..03a35583e 100644 --- a/apps/web/src/pages/profile/ProfileSettings.css +++ b/apps/web/src/pages/profile/ProfileSettings.css @@ -178,3 +178,35 @@ border-color: var(--border-color); margin: var(--space-lg) 0; } + +/* 微信账号绑定卡片 */ +.xx-settings-wechat { + display: flex; + align-items: center; + justify-content: space-between; + gap: var(--space-lg); + flex-wrap: wrap; +} + +.xx-settings-wechat-info { + display: flex; + align-items: center; + gap: var(--space-md); +} + +.xx-settings-wechat-info .xx-wechat-icon { + font-size: 28px; + line-height: 1; +} + +.xx-settings-wechat-info strong { + display: block; + color: var(--text-primary); + font-size: var(--font-size-md); +} + +.xx-settings-wechat-info p { + margin: 2px 0 0; + color: var(--text-secondary); + font-size: var(--font-size-sm); +} diff --git a/apps/web/src/pages/profile/Settings.tsx b/apps/web/src/pages/profile/Settings.tsx index d66eea05a..7513c3bcf 100644 --- a/apps/web/src/pages/profile/Settings.tsx +++ b/apps/web/src/pages/profile/Settings.tsx @@ -1,39 +1,111 @@ /** * 个人设置页面 - * P1-2: 添加 PageHead - * P1-3: antd Form/Input/Button/Alert → 自定义 UI 组件 + * - 个人资料(昵称)保存 + * - 微信账号绑定状态 / 绑定 / 解绑 */ -import React, { useState } from "react" +import React, { useEffect, useRef, useState } from "react" +import { useSearchParams } from "react-router-dom" +import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query" +import { message } from "antd" import { Button, Input, Modal } from "@/components/ui" +import { getCurrentUser, updateProfile, getWechatBindUrl, unbindWechat } from "@/api/auth" import { useAuthStore } from "@/store/authStore" import PageHead from "@/components/layout/PageHead" import "./ProfileSettings.css" const Settings: React.FC = () => { const user = useAuthStore((state) => state.user) + const setUser = useAuthStore((state) => state.setUser) + const queryClient = useQueryClient() + const [searchParams, setSearchParams] = useSearchParams() const [displayName, setDisplayName] = useState(user?.display_name || "") + const bindTipShownRef = useRef(false) - const handleSave = () => { - Modal.info({ - title: "提示", - content: "个人资料修改接口暂未开放,保存功能即将上线。", + // 拉取最新用户信息(微信绑定状态以后端为准) + const { data: freshUser } = useQuery({ + queryKey: ["currentUser"], + queryFn: getCurrentUser, + }) + + useEffect(() => { + if (freshUser) { + setUser(freshUser) + setDisplayName((prev) => prev || freshUser.display_name || "") + } + }, [freshUser, setUser]) + + // 绑定回调结果提示(?wechat_bind=success|failed) + useEffect(() => { + if (bindTipShownRef.current) return + const result = searchParams.get("wechat_bind") + if (!result) return + bindTipShownRef.current = true + if (result === "success") { + message.success("微信绑定成功") + } else if (result === "failed") { + message.error("微信绑定失败,请重试") + } + searchParams.delete("wechat_bind") + setSearchParams(searchParams, { replace: true }) + }, [searchParams, setSearchParams]) + + const wechatBound = user?.wechat_bound === true + + const saveProfileMutation = useMutation({ + mutationFn: () => updateProfile({ display_name: displayName.trim() }), + onSuccess: (updated) => { + setUser(updated) + message.success("资料已保存") + }, + onError: () => { + message.error("保存失败,请重试") + }, + }) + + const handleBindWechat = async () => { + try { + const result = await getWechatBindUrl() + localStorage.setItem("wechat_bind_state", result.state) + window.location.href = result.auth_url + } catch { + message.error("微信绑定暂不可用,请稍后重试") + } + } + + const unbindMutation = useMutation({ + mutationFn: unbindWechat, + onSuccess: () => { + message.success("已解绑微信") + queryClient.invalidateQueries({ queryKey: ["currentUser"] }) + // 本地立即更新,避免等待刷新 + if (user) { + setUser({ ...user, wechat_bound: false, wechat_nickname: "" }) + } + }, + onError: () => { + message.error("解绑失败,请重试") + }, + }) + + const handleUnbind = () => { + Modal.confirm({ + title: "解绑微信", + content: "解绑后将无法使用微信登录该账号,确定要解绑吗?", + okText: "确定解绑", + cancelText: "取消", + okButtonProps: { danger: true }, + onOk: () => unbindMutation.mutateAsync(), }) } + const displayNameDirty = displayName.trim() !== (user?.display_name || "") + return (

个人信息

-
- ℹ️ -
- 个人资料编辑暂未开放 -

当前仅展示登录用户信息,资料修改接口接入后再开放保存。

-
-
-
@@ -42,25 +114,76 @@ const Settings: React.FC = () => {
- -
- -
- setDisplayName(e.target.value)} - placeholder="请输入显示名称" + value={user?.email && !user.email.endsWith("@wechat.local") ? user.email : ""} + disabled + placeholder={user?.email?.endsWith("@wechat.local") ? "微信账号暂未绑定邮箱" : "邮箱"} />
-
+ +
+
+ +
+

微信账号

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

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

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

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

+ + )} +
+
+
+ {wechatBound ? ( + + ) : ( + + )} +
+
+
) } 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/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/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/auth/WechatCallback.test.tsx b/apps/web/src/test/pages/auth/WechatCallback.test.tsx index fe558f930..01675db9b 100644 --- a/apps/web/src/test/pages/auth/WechatCallback.test.tsx +++ b/apps/web/src/test/pages/auth/WechatCallback.test.tsx @@ -1,79 +1,127 @@ -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() +const mockSearchParams = [new URLSearchParams({ code: "test_code", state: "test_state" })] 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 callbackShouldFail = false + 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 (callbackShouldFail) throw new Error("fail") + 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 -
- ), -})) - -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() + callbackShouldFail = false + 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("state 不匹配显示安全错误", async () => { + localStorageStore.wechat_state = "other_state" + renderPage() + await waitFor(() => { + expect(screen.getByText("安全校验失败,请重新登录")).toBeTruthy() + }) + expect(mockNavigate).not.toHaveBeenCalled() + }) + + it("接口失败显示错误提示", async () => { + callbackShouldFail = true + renderPage() + await waitFor(() => { + expect(screen.getByText("微信登录失败,请重试")).toBeTruthy() + }) + }) + + it("处理中显示 loading", () => { + renderPage() + expect(screen.getByText("微信登录中...")).toBeTruthy() }) }) 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..b93218448 --- /dev/null +++ b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx @@ -0,0 +1,132 @@ +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("新用户可见昵称表单并能提交", 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 () => { + 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() + }) + }) +}) -- 2.54.0 From 9840d5d77842a8d9493af53c950233b7bd3870ed Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 17:39:46 +0800 Subject: [PATCH 17/33] =?UTF-8?q?deploy(#1718):=20=E5=BE=AE=E4=BF=A1OAuth?= =?UTF-8?q?=E5=87=AD=E8=AF=81=E7=BA=B3=E5=85=A5env=E6=A8=A1=E6=9D=BF+CI?= =?UTF-8?q?=E6=B8=B2=E6=9F=93=EF=BC=8C=E4=BF=AE=E5=A4=8Dstaging=E9=83=A8?= =?UTF-8?q?=E7=BD=B2=E8=A6=86=E7=9B=96=E4=B8=A2=E5=A4=B1=20(#1721)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- .gitea/workflows/ci-pipeline.yml | 2 ++ deploy/configs/.env.production | 7 +++++++ deploy/configs/.env.staging | 7 +++++++ scripts/render_env.sh | 2 +- 4 files changed, 17 insertions(+), 1 deletion(-) 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/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/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:-}" # 已经在环境中了,无需额外操作 -- 2.54.0 From df99305dd68c644c3e48ae2f18cc0b9420c8c39b Mon Sep 17 00:00:00 2001 From: saas-backend-bot Date: Sat, 5 Sep 2026 17:31:39 +0800 Subject: [PATCH 18/33] =?UTF-8?q?feat(worker):=20celery=20=E9=98=9F?= =?UTF-8?q?=E5=88=97=E9=9A=94=E7=A6=BB=20+=20=E5=AD=A4=E5=84=BF=E4=BB=BB?= =?UTF-8?q?=E5=8A=A1=E6=B6=88=E6=81=AF=E4=BD=9C=E5=BA=9F=20(#1714)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 问题:素材转码与视频生成共用 celery 默认队列、worker 单进程消费, 20+ 转码积压会把用户生成任务堵 40 分钟以上;孤儿清理把任务标 failed 后 Redis 队列消息未作废,消息被重投导致 failed→running 非法转换, worker 打印 ERROR 后继续产出半成品。 队列隔离: - 新增 packages/shared/celery_queues.py:generation/transcode/celery 三队列与 task_routes(generate_video→generation;ingest_asset/ classify_asset/duplication→transcode),apply_queue_settings() - worker 入口改双进程:generation worker 独占队列并内嵌 beat (prefetch=1, GENERATION_CONCURRENCY 默认 2),transcode worker 消费 transcode,celery(并发=总-2,最小 1),任一退出则整体终止 - compose/部署脚本/ps1 同步新增 GENERATION_CONCURRENCY 与健康检查 消息作废: - 新增 packages/shared/celery_orphan_guard.py:终态守卫 ensure_task_claimable、Redis 队列消息物理清理(JSON 信封解析, 按业务 id + celery headers.id 双匹配,未命中 rpush 保序)、 revoke_and_purge(control.revoke + 物理清队列双保险) - 入队点(生成/上传/分片/重试)send_task 后持久化 celery_task_id 到 generation_tasks/ingest_jobs(新列,067 迁移,失败仅 warning) - generate_video/ingest_asset 执行前校验 DB 状态:终态直接 discarded 不进业务逻辑;mark_processing 返回 False(非法转换)安全中止 - 孤儿/超时清理标 failed 时同时 revoke + 清队列消息 - pending 超时阈值 15→45 分钟,与 running 孤儿(20min)区分 测试:新增 22 个单测(路由表/真实 Redis 消息清理/终态守卫/ 非法转换中止/标 failed 后消息不重投/入队持久化),全量 14301 passed;067 迁移隔离 DDL 验证 upgrade/downgrade 通过。 --- alembic/versions/067_celery_task_id_revoke.py | 35 +++ apps/api/app/api/routes/chunked_upload.py | 8 +- apps/api/app/api/routes/ingest_jobs.py | 8 +- apps/api/app/api/routes/task_center.py | 8 +- apps/api/app/api/routes/upload.py | 14 +- apps/api/app/core/celery_app.py | 8 + apps/api/app/core/task_enqueue.py | 12 +- apps/worker/worker_app/celery_app.py | 12 + apps/worker/worker_app/tasks/_startup.py | 100 ++++++- apps/worker/worker_app/tasks/cleanup.py | 6 +- apps/worker/worker_app/tasks/generation.py | 32 ++- apps/worker/worker_app/tasks/ingest.py | 13 + infra/docker/compose.yml | 9 +- infra/docker/deploy-production-registry.sh | 3 +- infra/docker/deploy-staging-registry.sh | 3 +- infra/docker/entrypoint-worker.sh | 63 ++++- .../generation_task_repository.py | 65 ++--- .../sqlalchemy_impl/ingest_job_repository.py | 5 + packages/adapters/sqlalchemy_impl/models.py | 2 + packages/application/ingest_jobs.py | 2 + packages/domain/entities.py | 3 + packages/domain/generation_task.py | 1 + packages/shared/celery_orphan_guard.py | 222 ++++++++++++++++ packages/shared/celery_queues.py | 58 +++++ start-worker.ps1 | 2 +- .../unit/test_celery_queue_isolation_1714.py | 176 +++++++++++++ .../test_enqueue_persists_celery_id_1714.py | 57 ++++ tests/unit/test_stale_task_revoke_1714.py | 204 +++++++++++++++ tests/unit/test_task_discard_guard_1714.py | 243 ++++++++++++++++++ tests/unit/test_task_queue_limit.py | 8 +- 30 files changed, 1319 insertions(+), 63 deletions(-) create mode 100644 alembic/versions/067_celery_task_id_revoke.py create mode 100644 packages/shared/celery_orphan_guard.py create mode 100644 packages/shared/celery_queues.py create mode 100644 tests/unit/test_celery_queue_isolation_1714.py create mode 100644 tests/unit/test_enqueue_persists_celery_id_1714.py create mode 100644 tests/unit/test_stale_task_revoke_1714.py create mode 100644 tests/unit/test_task_discard_guard_1714.py 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/apps/api/app/api/routes/chunked_upload.py b/apps/api/app/api/routes/chunked_upload.py index 1db5a1647..cb97ff06b 100644 --- a/apps/api/app/api/routes/chunked_upload.py +++ b/apps/api/app/api/routes/chunked_upload.py @@ -381,7 +381,13 @@ 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]) + if getattr(celery_result, "id", ""): + try: + job.celery_task_id = celery_result.id + ingest_job_repository.update(job) + except Exception: # noqa: BLE001 + pass # Update metadata status meta["status"] = "completed" 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 f5b046f30..48b367a0e 100644 --- a/apps/api/app/api/routes/upload.py +++ b/apps/api/app/api/routes/upload.py @@ -203,6 +203,17 @@ def _create_pending_asset( 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, @@ -221,7 +232,8 @@ def _submit_ingest_job( 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 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 95e0d48cc..c3b11329c 100755 --- a/apps/api/app/core/task_enqueue.py +++ b/apps/api/app/core/task_enqueue.py @@ -279,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/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py index a5051e5a4..75477f583 100755 --- a/apps/worker/worker_app/celery_app.py +++ b/apps/worker/worker_app/celery_app.py @@ -6,6 +6,18 @@ 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 崩溃时未完成任务重回队列,由执行前守卫丢弃作废消息 + celery_app.conf.imports = ( "worker_app.tasks.health", "worker_app.tasks.ingest", diff --git a/apps/worker/worker_app/tasks/_startup.py b/apps/worker/worker_app/tasks/_startup.py index c3e6e8dfb..2d3f80b7c 100644 --- a/apps/worker/worker_app/tasks/_startup.py +++ b/apps/worker/worker_app/tasks/_startup.py @@ -14,7 +14,17 @@ def cleanup_stale_running_with_session(repo, timeout_minutes: int) -> int: Returns: 清理的任务数量 """ - return repo.cleanup_stale_running(timeout_minutes) + 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: @@ -23,7 +33,44 @@ def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int: Returns: 清理的任务数量 """ - return repo.cleanup_stale_pending(timeout_minutes) + 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 任务超过此时间无进度更新则视为卡死。 @@ -31,10 +78,12 @@ def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int: # 20 分钟阈值覆盖硬超时 + 重试 + 余量,绝不误杀正常任务。 ORPHAN_TASK_TIMEOUT_MINUTES = 20 -# Pending 任务超时阈值:任务创建后超过此时间仍未被 worker 拉取, -# 说明 worker 已停止消费(容器异常/卡死),清掉释放限流名额。 -# 依据:满队列(20 pending)× 平均 2 分钟 / 并发 4 ≈ 10 分钟,15 分钟留余量。 -PENDING_TASK_TIMEOUT_MINUTES = 15 +# 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 @@ -57,11 +106,14 @@ def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> session = SessionLocal() try: repo = SQLAlchemyGenerationTaskRepository(session) - count = cleanup_stale_running_with_session(repo, timeout_minutes) + 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 @@ -95,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 @@ -105,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) @@ -130,9 +186,12 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU session = SessionLocal() try: repo = SQLAlchemyGenerationTaskRepository(session) - count = cleanup_stale_pending_with_session(repo, 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 @@ -143,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 """统一清理所有超时的孤儿任务。 diff --git a/apps/worker/worker_app/tasks/cleanup.py b/apps/worker/worker_app/tasks/cleanup.py index 5948e8a0d..07a7de8dc 100644 --- a/apps/worker/worker_app/tasks/cleanup.py +++ b/apps/worker/worker_app/tasks/cleanup.py @@ -25,10 +25,12 @@ def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_ 每 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: 超时时间(分钟),默认 15 分钟 + timeout_minutes: 超时时间(分钟),默认 45 分钟(pending 排队阈值放宽, + 与 running 孤儿 20 分钟区分,避免正常排队任务被误杀) Returns: {"cleaned": int} 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 c42169b92..4d7094952 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -428,6 +428,19 @@ 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 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..2f20644b2 100755 --- a/infra/docker/entrypoint-worker.sh +++ b/infra/docker/entrypoint-worker.sh @@ -1,18 +1,63 @@ #!/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}" + -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/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py index 9ef851ad5..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 "", @@ -315,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 "" @@ -326,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) @@ -344,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 = { @@ -355,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 c2e11c24b..7c42fdeee 100644 --- a/packages/adapters/sqlalchemy_impl/ingest_job_repository.py +++ b/packages/adapters/sqlalchemy_impl/ingest_job_repository.py @@ -19,6 +19,7 @@ class SQLAlchemyIngestJobRepository: 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, ) @@ -40,6 +41,7 @@ class SQLAlchemyIngestJobRepository: 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, ) @@ -59,6 +61,9 @@ class SQLAlchemyIngestJobRepository: 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 3bc1c8def..0ee3e8565 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -246,6 +246,7 @@ class IngestJobModel(Base): 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)) @@ -298,6 +299,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/application/ingest_jobs.py b/packages/application/ingest_jobs.py index 576a2b25c..92f062ac2 100644 --- a/packages/application/ingest_jobs.py +++ b/packages/application/ingest_jobs.py @@ -13,6 +13,7 @@ class SubmitIngestJobCommand: storage_key: str file_hash: str = "" asset_id: str = "" + celery_task_id: str = "" class SubmitIngestJobUseCase: @@ -26,5 +27,6 @@ class SubmitIngestJobUseCase: 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/domain/entities.py b/packages/domain/entities.py index 9ed9e5b8a..8d1060c87 100755 --- a/packages/domain/entities.py +++ b/packages/domain/entities.py @@ -270,6 +270,7 @@ class IngestJob: 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)) @@ -281,6 +282,7 @@ class IngestJob: storage_key: str, file_hash: str = "", asset_id: str = "", + celery_task_id: str = "", ) -> "IngestJob": if not project_id.strip(): raise ValueError("project_id 不能为空") @@ -295,4 +297,5 @@ class IngestJob: 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/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/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_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_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_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_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 时,入队后校验也跳过用户级,只查全局。""" -- 2.54.0 From 6a1ec20e686e8c0b3590c0b6efe3f0f412118907 Mon Sep 17 00:00:00 2001 From: saas-backend-agent Date: Sat, 5 Sep 2026 19:52:50 +0800 Subject: [PATCH 19/33] =?UTF-8?q?test(#1714):=20=E8=A1=A5=20mock=20redis?= =?UTF-8?q?=20=E8=A6=86=E7=9B=96=E7=8E=87=E6=B5=8B=E8=AF=95=EF=BC=8Cdiff?= =?UTF-8?q?=20coverage=2098%?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit CI runner 无本地 redis,purge/revoke 与路由持久化的真实 redis 测试 全部 skip 导致这些改动零覆盖、diff coverage 45% 未达 60% 门禁: - 新增 test_orphan_guard_purge_mocked_1714.py(31 测试,全 mock): _extract_business_ids 各形态(三元组/裸 args/dict args/坏 JSON/坏 base64/空 args)、_purge_one_queue(biz/celery id 双匹配命中、未命中 rpush 保序、全删不 rpush、解析失败保守保留、lrange/重写异常)、 purge_stale_messages_from_queues(空 ids 早退、redis 未安装、连接失败、 close 异常)、revoke_and_purge(逐条 revoke、异常不阻断、空 id 跳过) - 新增 test_persist_celery_id_routes_1714.py(9 测试):ingest_jobs / task_center 重试 / upload helper 的 celery_task_id 持久化正常与异常 吞掉分支、task_enqueue 持久化失败仍入队成功、celery_app 队列配置 异常不阻断启动(用独立模块对象加载,不 reload 污染 task_enqueue)、 仓储 update 落库 celery_task_id - chunked_upload 完成回调去重:改为复用 upload._persist_celery_task_id --- apps/api/app/api/routes/chunked_upload.py | 8 +- .../test_orphan_guard_purge_mocked_1714.py | 360 ++++++++++++++++++ .../test_persist_celery_id_routes_1714.py | 259 +++++++++++++ 3 files changed, 621 insertions(+), 6 deletions(-) create mode 100644 tests/unit/test_orphan_guard_purge_mocked_1714.py create mode 100644 tests/unit/test_persist_celery_id_routes_1714.py diff --git a/apps/api/app/api/routes/chunked_upload.py b/apps/api/app/api/routes/chunked_upload.py index cb97ff06b..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 @@ -382,12 +383,7 @@ async def complete_chunked_upload( ) ) 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 + _persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", "")) # Update metadata status meta["status"] = "completed" 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_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" -- 2.54.0 From f11b71361d8bbda5bc6399860448077cc02346ba Mon Sep 17 00:00:00 2001 From: saas-backend-agent Date: Sat, 5 Sep 2026 20:45:02 +0800 Subject: [PATCH 20/33] =?UTF-8?q?feat(#1718):=20=E5=BE=AE=E4=BF=A1=20state?= =?UTF-8?q?=20=E5=AD=98=E5=82=A8=20Redis=20=E5=8C=96=20+=20=E5=9B=9E?= =?UTF-8?q?=E8=B0=83=20UA=20=E6=97=A5=E5=BF=97=20+=20=E4=B8=AD=E6=96=87?= =?UTF-8?q?=E6=98=B5=E7=A7=B0=20UTF-8=20=E4=BF=AE=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - state store 改 Redis(复用 celery Redis,key 前缀 wechat:state:, TTL 10 分钟):SET NX EX 写入,Lua 脚本原子 GET+DEL 一次性消费 (兼容 Redis <6.2 无 GETDEL);Redis 不可用时自动降级内存,登录不中断; 容器重启/多实例后 state 不丢,修复 worker 扩容后回调 state 失效 - /wechat/callback 加可观测日志:User-Agent(识别 MicroMessenger 微信内置浏览器)、state 校验结果、失败上下文,便于排查回调停滞 - 修复微信中文昵称乱码:sns/oauth2/access_token 与 sns/userinfo 响应在 .json() 前显式 encoding=utf-8(微信响应头不带 charset, requests 默认 ISO-8859-1 解码导致中文乱码) - 15 个新单测(全 mock/fake,CI 无 redis 也覆盖):Redis state 存取/一次性消费/eval 降级/异常降级内存/ping 失败降级、中文昵称 UTF-8 解析、errcode 透传、callback 路由日志分支、工厂降级分支 --- apps/api/app/api/routes/auth.py | 22 +- .../application/auth/wechat_oauth_service.py | 95 ++++++- .../unit/test_wechat_callback_logging_1718.py | 115 ++++++++ tests/unit/test_wechat_state_redis_1718.py | 262 ++++++++++++++++++ 4 files changed, 492 insertions(+), 2 deletions(-) create mode 100644 tests/unit/test_wechat_callback_logging_1718.py create mode 100644 tests/unit/test_wechat_state_redis_1718.py diff --git a/apps/api/app/api/routes/auth.py b/apps/api/app/api/routes/auth.py index 38438f931..5e8234484 100755 --- a/apps/api/app/api/routes/auth.py +++ b/apps/api/app/api/routes/auth.py @@ -13,7 +13,7 @@ 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 @@ -426,6 +426,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 +434,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) diff --git a/packages/application/auth/wechat_oauth_service.py b/packages/application/auth/wechat_oauth_service.py index d8cd3d125..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: @@ -206,9 +285,23 @@ class WechatOAuthService: _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 服务单例(state store 跨请求共享)""" global _oauth_service_singleton if _oauth_service_singleton is None: - _oauth_service_singleton = WechatOAuthService() + _oauth_service_singleton = WechatOAuthService(state_store=_build_default_state_store()) return _oauth_service_singleton 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_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) -- 2.54.0 From 5ca64898b7db8bae2960e70576e9e4c113a4cc78 Mon Sep 17 00:00:00 2001 From: saas-backend-agent Date: Sat, 5 Sep 2026 21:31:44 +0800 Subject: [PATCH 21/33] =?UTF-8?q?feat(#1719):=20=E5=BE=AE=E4=BF=A1?= =?UTF-8?q?=E8=B4=A6=E5=8F=B7=E7=BB=91=E5=AE=9A/=E8=A7=A3=E7=BB=91?= =?UTF-8?q?=E4=B8=89=E6=8E=A5=E5=8F=A3=EF=BC=88GET=20bind/url=E3=80=81POST?= =?UTF-8?q?=20bind=E3=80=81DELETE=20bind=EF=BC=89+=20/me=20=E8=BF=94?= =?UTF-8?q?=E5=9B=9E=20wechat=5Fbound?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - WechatBindUseCase:绑定到当前登录账号不建新号;openid/unionid 已绑其他账号 409;已绑同一微信幂等 - WechatUnbindUseCase:解绑前守卫——必须有已验证手机或真实已验证邮箱(随机密码 hash/占位邮箱不算兜底,口径同 binding_complete) - /auth/me 增加 wechat_bound 字段,供设置页判断绑定状态 - state 复用 RedisStateStore CSRF 校验;22 个新单测,diff coverage 100% --- apps/api/app/api/routes/auth.py | 118 ++++++++ .../application/auth/wechat_bind_use_case.py | 115 ++++++++ tests/unit/test_wechat_bind_routes_1719.py | 233 ++++++++++++++++ tests/unit/test_wechat_bind_use_case_1719.py | 251 ++++++++++++++++++ 4 files changed, 717 insertions(+) create mode 100644 packages/application/auth/wechat_bind_use_case.py create mode 100644 tests/unit/test_wechat_bind_routes_1719.py create mode 100644 tests/unit/test_wechat_bind_use_case_1719.py diff --git a/apps/api/app/api/routes/auth.py b/apps/api/app/api/routes/auth.py index 5e8234484..0c268d92d 100755 --- a/apps/api/app/api/routes/auth.py +++ b/apps/api/app/api/routes/auth.py @@ -84,6 +84,7 @@ class CurrentUserResponse(BaseModel): phone: str = "" phone_verified: bool = False binding_complete: bool = False + wechat_bound: bool = False class PasswordResetRequestModel(BaseModel): @@ -272,6 +273,7 @@ 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), ) @@ -492,6 +494,122 @@ async def wechat_callback( ) +# ==================== 微信账号绑定/解绑(已登录用户) ==================== + + +class WechatBindUrlResponse(BaseModel): + auth_url: str + state: str + + +class WechatBindCompleteRequest(BaseModel): + code: str + state: str = "" + + +class WechatBindUserProfile(BaseModel): + """绑定/解绑后返回的用户信息(字段对齐 /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 + + +class WechatBindCompleteResponse(BaseModel): + success: bool + user: WechatBindUserProfile + + +class WechatUnbindResponse(BaseModel): + success: bool + + +def _wechat_user_profile(user) -> WechatBindUserProfile: + binding_complete = bool( + user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email + ) + return WechatBindUserProfile( + 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), + ) + + +@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=_wechat_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/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/tests/unit/test_wechat_bind_routes_1719.py b/tests/unit/test_wechat_bind_routes_1719.py new file mode 100644 index 000000000..37df5fada --- /dev/null +++ b/tests/unit/test_wechat_bind_routes_1719.py @@ -0,0 +1,233 @@ +"""#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, + ) + 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, + ) + + 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 -- 2.54.0 From cdcb032e452aacba59e6c2e58d1683dc89d53c6a Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 5 Sep 2026 22:52:49 +0800 Subject: [PATCH 22/33] =?UTF-8?q?fix(#1718/#1714):=20=E5=BE=AE=E4=BF=A1?= =?UTF-8?q?=E5=9B=9E=E8=B0=83state=E8=AF=AF=E6=9D=80=E4=BF=AE=E5=A4=8D+?= =?UTF-8?q?=E9=94=99=E8=AF=AF=E9=80=8F=E4=BC=A0=E9=98=B2=E8=BF=9E=E7=82=B9?= =?UTF-8?q?=E3=80=81=E4=B8=8A=E4=BC=A0=E5=A4=B1=E8=B4=A5=E5=AE=8C=E6=95=B4?= =?UTF-8?q?=E5=8F=AF=E8=A7=82=E6=B5=8B=E3=80=81=E5=93=88=E5=B8=8C=E9=98=88?= =?UTF-8?q?=E5=80=BC=E9=99=8D=E8=87=B364MB=E3=80=81=E6=98=B5=E7=A7=B0?= =?UTF-8?q?=E4=B8=8D=E9=A2=84=E5=A1=AB=20(#1723)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- apps/web/src/api/assets/upload.ts | 10 +- apps/web/src/api/assets/uploadDedup.ts | 14 +- apps/web/src/api/errors.ts | 132 ++++++++++++++++++ apps/web/src/pages/assets/assets.css | 14 ++ .../assets/components/UploadQueuePanel.tsx | 27 +++- .../src/pages/assets/hooks/useAssetUpload.ts | 32 +++-- apps/web/src/pages/auth/Login.tsx | 20 ++- .../web/src/pages/auth/WechatBindCallback.tsx | 16 +-- apps/web/src/pages/auth/WechatCallback.tsx | 28 ++-- apps/web/src/pages/auth/WechatOnboarding.tsx | 3 +- apps/web/src/test/api/uploadDedup.test.ts | 33 ++++- .../test/pages/assets/useAssetUpload.test.tsx | 28 ++++ .../pages/auth/WechatBindCallback.test.tsx | 105 ++++++++++++++ .../test/pages/auth/WechatCallback.test.tsx | 55 ++++++-- .../test/pages/auth/WechatOnboarding.test.tsx | 6 + 15 files changed, 468 insertions(+), 55 deletions(-) create mode 100644 apps/web/src/api/errors.ts create mode 100644 apps/web/src/test/pages/auth/WechatBindCallback.test.tsx diff --git a/apps/web/src/api/assets/upload.ts b/apps/web/src/api/assets/upload.ts index ab450cd69..3b08b2914 100644 --- a/apps/web/src/api/assets/upload.ts +++ b/apps/web/src/api/assets/upload.ts @@ -126,7 +126,15 @@ export const prepareDirectUploadHandle = async (data: { /** 本次逻辑上传的幂等 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, diff --git a/apps/web/src/api/assets/uploadDedup.ts b/apps/web/src/api/assets/uploadDedup.ts index 67f31f7de..f30ad6472 100644 --- a/apps/web/src/api/assets/uploadDedup.ts +++ b/apps/web/src/api/assets/uploadDedup.ts @@ -13,10 +13,10 @@ * 重试复用同一 ID,重新入队才生成新 ID) */ -/** 大文件抽样阈值:超过此大小只哈希头尾片段,避免上传前长时间卡 UI */ -export const HASH_FULL_READ_LIMIT = 256 * 1024 * 1024 // 256MB -/** 抽样读取的头尾片段大小(各 8MB) */ -export const HASH_SAMPLE_CHUNK = 8 * 1024 * 1024 +/** 全量哈希阈值:≤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" @@ -95,10 +95,10 @@ function toHex(buffer: ArrayBuffer): string { /** * 计算文件内容 SHA-256(hex,64 字符,与后端 file_hash 字段长度一致)。 - * - ≤256MB:全量哈希,内容一致必然一致 - * - >256MB:哈希「头部 8MB + 尾部 8MB + 文件大小」,视频素材体积大、 + * - ≤64MB:全量哈希,内容一致必然一致 + * - >64MB:哈希「头部 16MB + 尾部 16MB + 文件大小」,视频素材体积大、 * 头部含 moov 元数据、尾部含 mdat 结尾,抽样碰撞概率可忽略, - * 且避免上传前对 2GB 文件全量读取造成长时间卡顿 + * 且避免 100~256MB 视频被整文件读进内存导致页面卡死/崩溃 * * 运行环境不支持 crypto.subtle(非安全上下文/老浏览器)时返回空字符串, * 调用方据此降级为不传 hash(后端仍有幂等 token + 同文件名兜底去重)。 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/pages/assets/assets.css b/apps/web/src/pages/assets/assets.css index 1c143fb79..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; diff --git a/apps/web/src/pages/assets/components/UploadQueuePanel.tsx b/apps/web/src/pages/assets/components/UploadQueuePanel.tsx index a88f4f9eb..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, @@ -82,8 +90,23 @@ const UploadQueuePanel: React.FC = ({ {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" && ( diff --git a/apps/web/src/pages/assets/hooks/useAssetUpload.ts b/apps/web/src/pages/assets/hooks/useAssetUpload.ts index 0d0a588b2..d322e3d5e 100644 --- a/apps/web/src/pages/assets/hooks/useAssetUpload.ts +++ b/apps/web/src/pages/assets/hooks/useAssetUpload.ts @@ -2,6 +2,7 @@ import { useState, useCallback, useRef, useEffect } from "react" import { useQueryClient } from "@tanstack/react-query" import { message } from "antd" import { prepareDirectUploadHandle, type DirectUploadHandle } from "@/api/assets" +import { getErrorMessage, isErrorMsgShown } from "@/api/errors" import { MAX_FILE_SIZE } from "../constants" import { computeFileHash, @@ -44,9 +45,15 @@ export interface UploadItem { /** 批量直传最大并发数,避免多文件瓜分上行带宽 */ const MAX_CONCURRENT = 3 -/** complete 阶段失败后的错误提示:素材可能已在服务器处理中,重试不会重新上传 */ -const COMPLETE_ERROR_HINT = - "确认请求失败,素材可能已在服务器处理中;点重试将安全确认,不会重新上传文件" +/** complete 阶段失败后的安全提示:素材可能已在后端建成,重试只重发 complete 幂等安全 */ +export const COMPLETE_RETRY_HINT = "素材可能已在服务器处理中,点重试将安全确认,不会重新上传文件" + +/** 失败阶段中文名(toast 提示用,明确失败发生在哪一步) */ +const STAGE_LABEL: Record = { + prepare: "准备上传", + transfer: "文件传输", + complete: "确认入库", +} /** * 素材批量上传 Hook @@ -165,20 +172,25 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) { message.success(`"${item.fileName}" 上传完成,正在转码处理`) } } catch (err: unknown) { - const detail = err instanceof Error ? err.message : "上传失败" + // 完整失败原因:HTTP 状态码 / OSS XML 的 Code+Message / 后端 detail, + // 由 getErrorMessage 统一提取(OSS XHR 错误自带「OSS 直传失败: HTTP xxx ...」明细) + const detail = getErrorMessage(err, "未知错误") console.error("[useAssetUpload] 上传失败:", item.fileName, stage, err) if (stage === "complete") { // complete 失败(超时/5xx/网络):后端记录可能已建成,handle 保留供幂等重试; - // 刷新列表让用户看到可能已创建的「处理中」素材,避免误以为没传上去而重复操作 + // 刷新列表让用户看到可能已创建的「处理中」素材,避免误以为没传上去而重复操作。 + // 卡片同时展示真实错误原因 + 安全重试提示(重试只重发 complete,不重新上传) refreshList() updateItem(item.tempId, { status: "error", failedStage: "complete", - error: COMPLETE_ERROR_HINT, + error: `${detail}\n${COMPLETE_RETRY_HINT}`, hint: undefined, }) - message.error(`"${item.fileName}" ${COMPLETE_ERROR_HINT}`) + if (!isErrorMsgShown(err)) { + message.error(`"${item.fileName}" 确认入库失败:${detail}`) + } } else { // prepare / transfer 失败:后端尚无素材记录,可安全全量重跑 handlesRef.current.delete(item.tempId) @@ -188,7 +200,11 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) { error: detail, hint: undefined, }) - message.error(`"${item.fileName}" 上传失败:${detail}`) + // 拦截器已对后端错误弹过 toast(含真实 detail)时不重复弹; + // OSS XHR 直传错误不走 axios,必须在这里弹 + if (!isErrorMsgShown(err)) { + message.error(`"${item.fileName}" ${STAGE_LABEL[stage]}失败:${detail}`) + } } } }, diff --git a/apps/web/src/pages/auth/Login.tsx b/apps/web/src/pages/auth/Login.tsx index 1b92d9eca..90acff2b2 100644 --- a/apps/web/src/pages/auth/Login.tsx +++ b/apps/web/src/pages/auth/Login.tsx @@ -1,11 +1,12 @@ /** * 登录页面 - V21 完全对标 */ -import React, { useState } from "react" +import React, { useRef, useState } from "react" import { Form, Input, Checkbox, message } from "antd" import { Link, useNavigate } from "react-router-dom" import { useLogin } from "@/hooks/useAuth" import { getWechatAuthUrl } from "@/api/auth" +import { getErrorMessage, isErrorMsgShown } from "@/api/errors" import Button from "@/components/ui/Button" import "./Login.css" @@ -20,6 +21,9 @@ const Login: React.FC = () => { const loginMutation = useLogin() const [form] = Form.useForm() const [wechatLoading, setWechatLoading] = useState(false) + // 同步防连点守卫:state 更新有渲染间隙,连点两次会各自请求授权 URL, + // 后一次的 state 覆盖前一次写入 localStorage 的 state,导致回调校验失败 + const wechatStartingRef = useRef(false) const onFinish = async (values: LoginFormValues) => { try { @@ -36,8 +40,10 @@ const Login: React.FC = () => { } const handleWechatLogin = async () => { + if (wechatStartingRef.current) return + wechatStartingRef.current = true + setWechatLoading(true) try { - setWechatLoading(true) const result = await getWechatAuthUrl() // 保存 state 到 localStorage 用于回调时验证 localStorage.setItem("wechat_state", result.state) @@ -51,11 +57,15 @@ const Login: React.FC = () => { // 跳转到微信授权页 window.location.href = result.auth_url } catch (error) { - if (!(error as { __msgShown?: boolean })?.__msgShown) - message.error("微信登录暂不可用,请稍后重试") - } finally { + // 跳走前才可能回到这里;拦截器已弹过后端 detail 时不重复弹, + // 否则透传真实原因(如微信服务未配置、网络异常) + if (!isErrorMsgShown(error)) { + message.error(`微信登录启动失败:${getErrorMessage(error, "请稍后重试")}`) + } + wechatStartingRef.current = false setWechatLoading(false) } + // 成功时 window.location 跳走,不复位 loading(页面即将卸载) } return ( diff --git a/apps/web/src/pages/auth/WechatBindCallback.tsx b/apps/web/src/pages/auth/WechatBindCallback.tsx index ad013cdd8..61ded7e61 100644 --- a/apps/web/src/pages/auth/WechatBindCallback.tsx +++ b/apps/web/src/pages/auth/WechatBindCallback.tsx @@ -6,6 +6,7 @@ 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" const WechatBindCallback: React.FC = () => { @@ -19,17 +20,13 @@ const WechatBindCallback: React.FC = () => { const state = searchParams.get("state") if (!code || !state) { - setError("无效的回调参数") + setError("无效的回调参数,请回到设置页重新扫码绑定") return } const handleBind = async () => { - // state 校验:绑定场景由设置页生成并落库,前缀 bind: - const savedState = localStorage.getItem("wechat_bind_state") - if (!savedState || savedState !== state) { - setError("安全校验失败,请重新绑定") - return - } + // state 校验由后端 state store 一次性消费兜底(前端不再比对 localStorage, + // 微信内打开/跨浏览器场景本地无 state 会误杀);清理绑定前写入的 state localStorage.removeItem("wechat_bind_state") try { @@ -37,8 +34,9 @@ const WechatBindCallback: React.FC = () => { setUser(normalizeUser(result.user)) // 用 replace 回设置页,query 携带成功标记由设置页提示 navigate("/app/profile?wechat_bind=success", { replace: true }) - } catch { - navigate("/app/profile?wechat_bind=failed", { replace: true }) + } catch (err) { + // 绑定失败直接在本页展示真实原因(如微信已被其他账号绑定),不静默跳走 + setError(`微信绑定失败:${getErrorMessage(err, "请回到设置页重试")}`) } } diff --git a/apps/web/src/pages/auth/WechatCallback.tsx b/apps/web/src/pages/auth/WechatCallback.tsx index 4186c3033..754a4ca70 100644 --- a/apps/web/src/pages/auth/WechatCallback.tsx +++ b/apps/web/src/pages/auth/WechatCallback.tsx @@ -7,6 +7,7 @@ import React, { useEffect, useState } from "react" import { useSearchParams, useNavigate } from "react-router-dom" import { Spin } from "antd" import { wechatCallback, getCurrentUser, normalizeUser, type User } from "@/api/auth" +import { getErrorMessage } from "@/api/errors" import { useAuthStore } from "@/store/authStore" import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh" @@ -21,21 +22,27 @@ const WechatCallback: React.FC = () => { const code = searchParams.get("code") const state = searchParams.get("state") + // 微信重定向出错时(如用户拒绝授权 error=access_denied)直接展示原因 + const wxErrorCode = searchParams.get("error") + const wxErrDesc = searchParams.get("error_description") + if (wxErrorCode || wxErrDesc) { + const reason = [wxErrorCode, wxErrDesc].filter(Boolean).join(":") + setError(`微信授权失败:${reason}`) + setLoading(false) + return + } + if (!code || !state) { - setError("无效的回调参数") + setError("无效的回调参数,请重新扫码登录") setLoading(false) 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) @@ -63,8 +70,9 @@ const WechatCallback: React.FC = () => { const redirect = localStorage.getItem("login_redirect") || "/" localStorage.removeItem("login_redirect") navigate(redirect, { replace: true }) - } catch { - setError("微信登录失败,请重试") + } catch (err) { + // 透传后端真实错误(如 state 过期、code 已消费、接口异常),禁止吞成通用提示 + setError(`微信登录失败:${getErrorMessage(err, "请重试或更换登录方式")}`) setLoading(false) } } diff --git a/apps/web/src/pages/auth/WechatOnboarding.tsx b/apps/web/src/pages/auth/WechatOnboarding.tsx index a4893b846..0b237243f 100644 --- a/apps/web/src/pages/auth/WechatOnboarding.tsx +++ b/apps/web/src/pages/auth/WechatOnboarding.tsx @@ -67,7 +67,8 @@ const WechatOnboarding: React.FC = () => { onFinish={onFinish} autoComplete="off" layout="vertical" - initialValues={{ display_name: user?.display_name || "" }} + // 不预填:新微信用户必须自己输入昵称(user.display_name 可能是微信昵称/系统占位) + initialValues={{ display_name: "" }} > { }) }) -describe("computeFileHash 大文件抽样(>256MB)", () => { +describe("computeFileHash 大文件抽样(>64MB)", () => { it("抽样路径正常返回 64 位 hex,且大小不同则 hash 不同", async () => { // mock 一个「声称」300MB 的 File:slice 返回小 buffer 即可,不真分配 300MB const makeBig = (declaredSize: number, head: number) => { @@ -98,4 +100,31 @@ describe("computeFileHash 大文件抽样(>256MB)", () => { // 声明大小不同 → 写入的 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/pages/assets/useAssetUpload.test.tsx b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx index cda209f5d..4c5004fa0 100644 --- a/apps/web/src/test/pages/assets/useAssetUpload.test.tsx +++ b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx @@ -224,6 +224,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 () => { @@ -334,6 +339,9 @@ describe("useAssetUpload", () => { 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 不重复) @@ -352,4 +360,24 @@ describe("useAssetUpload", () => { 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("签名服务内部错误") + }) }) 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..a2592fc2f --- /dev/null +++ b/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx @@ -0,0 +1,105 @@ +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 }), +})) + +const renderPage = () => + render( + + + , + ) + +describe("WechatBindCallback Page", () => { + afterEach(() => { + cleanup() + }) + + beforeEach(() => { + vi.clearAllMocks() + 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() + }) + }) +}) diff --git a/apps/web/src/test/pages/auth/WechatCallback.test.tsx b/apps/web/src/test/pages/auth/WechatCallback.test.tsx index 01675db9b..5ef33e4c2 100644 --- a/apps/web/src/test/pages/auth/WechatCallback.test.tsx +++ b/apps/web/src/test/pages/auth/WechatCallback.test.tsx @@ -5,7 +5,11 @@ import WechatCallback from "@/pages/auth/WechatCallback" const mockNavigate = vi.fn() const mockSetAuth = vi.fn() -const mockSearchParams = [new URLSearchParams({ code: "test_code", state: "test_state" })] as const + +// 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 导致链式污染) @@ -20,7 +24,7 @@ vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => { let mockCallbackResult: Record = {} let mockCurrentUser: Record = {} -let callbackShouldFail = false +let callbackError: unknown = null vi.mock("react-router-dom", async () => { const actual = await vi.importActual("react-router-dom") @@ -33,7 +37,7 @@ vi.mock("react-router-dom", async () => { vi.mock("@/api/auth", () => ({ wechatCallback: vi.fn(async () => { - if (callbackShouldFail) throw new Error("fail") + if (callbackError) throw callbackError return mockCallbackResult }), getCurrentUser: vi.fn(async () => mockCurrentUser), @@ -63,7 +67,11 @@ describe("WechatCallback Page", () => { beforeEach(() => { vi.clearAllMocks() - callbackShouldFail = 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", @@ -103,20 +111,47 @@ describe("WechatCallback Page", () => { }) }) - it("state 不匹配显示安全错误", async () => { - localStorageStore.wechat_state = "other_state" + it("本地无 wechat_state(微信内打开/跨浏览器场景)不再误杀,正常完成登录", async () => { + delete localStorageStore.wechat_state renderPage() await waitFor(() => { - expect(screen.getByText("安全校验失败,请重新登录")).toBeTruthy() + 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("接口失败显示错误提示", async () => { - callbackShouldFail = true + it("缺少 code/state 参数时提示无效回调", async () => { + mockParams.delete("code") renderPage() await waitFor(() => { - expect(screen.getByText("微信登录失败,请重试")).toBeTruthy() + expect(screen.getByText(/无效的回调参数/)).toBeTruthy() }) }) diff --git a/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx index b93218448..41c4fcfa9 100644 --- a/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx +++ b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx @@ -83,6 +83,12 @@ describe("WechatOnboarding 昵称引导页", () => { expect(screen.queryByText("进入小虾智剪")).toBeNull() }) + it("昵称输入框不预填,必须用户自己输入", () => { + renderPage() + expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy() + expect((screen.getByPlaceholderText("请输入您的昵称") as HTMLInputElement).value).toBe("") + }) + it("新用户可见昵称表单并能提交", async () => { renderPage() expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy() -- 2.54.0 From a83ed588649019d4879e3ffd973bb2014770f046 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 00:31:48 +0800 Subject: [PATCH 23/33] =?UTF-8?q?feat(#1718):=20=E5=BE=AE=E4=BF=A1?= =?UTF-8?q?=E7=99=BB=E5=BD=95/=E7=BB=91=E5=AE=9A=E6=94=B9=E4=B8=BA?= =?UTF-8?q?=E5=BC=B9=E7=AA=97=E5=86=85=E5=B5=8C=E4=BA=8C=E7=BB=B4=E7=A0=81?= =?UTF-8?q?=EF=BC=8C=E4=B8=8D=E5=86=8D=E6=95=B4=E9=A1=B5=E8=B7=B3=E8=BD=AC?= =?UTF-8?q?=20(#1726)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/src/api/auth/wxLogin.ts | 112 ++++++++ .../auth/WechatQrModal/WechatQrModal.css | 74 +++++ .../components/auth/WechatQrModal/index.tsx | 265 ++++++++++++++++++ .../components/auth/WechatQrModal/messages.ts | 71 +++++ apps/web/src/pages/auth/Login.tsx | 68 +++-- .../web/src/pages/auth/WechatBindCallback.tsx | 31 +- apps/web/src/pages/auth/WechatCallback.tsx | 37 ++- apps/web/src/pages/profile/Settings.tsx | 26 +- apps/web/src/test/api/wxLogin.test.ts | 64 +++++ .../test/components/WechatQrModal.test.tsx | 165 +++++++++++ .../pages/auth/WechatBindCallback.test.tsx | 39 +++ .../test/pages/auth/WechatCallback.test.tsx | 61 ++++ 12 files changed, 956 insertions(+), 57 deletions(-) create mode 100644 apps/web/src/api/auth/wxLogin.ts create mode 100644 apps/web/src/components/auth/WechatQrModal/WechatQrModal.css create mode 100644 apps/web/src/components/auth/WechatQrModal/index.tsx create mode 100644 apps/web/src/components/auth/WechatQrModal/messages.ts create mode 100644 apps/web/src/test/api/wxLogin.test.ts create mode 100644 apps/web/src/test/components/WechatQrModal.test.tsx 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/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/pages/auth/Login.tsx b/apps/web/src/pages/auth/Login.tsx index 90acff2b2..a869eecdd 100644 --- a/apps/web/src/pages/auth/Login.tsx +++ b/apps/web/src/pages/auth/Login.tsx @@ -1,13 +1,12 @@ /** * 登录页面 - V21 完全对标 */ -import React, { useRef, useState } from "react" +import React, { useState } from "react" import { Form, Input, Checkbox, message } from "antd" import { Link, useNavigate } from "react-router-dom" import { useLogin } from "@/hooks/useAuth" -import { getWechatAuthUrl } from "@/api/auth" -import { getErrorMessage, isErrorMsgShown } from "@/api/errors" import Button from "@/components/ui/Button" +import WechatQrModal from "@/components/auth/WechatQrModal" import "./Login.css" interface LoginFormValues { @@ -20,10 +19,7 @@ const Login: React.FC = () => { const navigate = useNavigate() const loginMutation = useLogin() const [form] = Form.useForm() - const [wechatLoading, setWechatLoading] = useState(false) - // 同步防连点守卫:state 更新有渲染间隙,连点两次会各自请求授权 URL, - // 后一次的 state 覆盖前一次写入 localStorage 的 state,导致回调校验失败 - const wechatStartingRef = useRef(false) + const [wechatQrOpen, setWechatQrOpen] = useState(false) const onFinish = async (values: LoginFormValues) => { try { @@ -39,33 +35,28 @@ const Login: React.FC = () => { } } - const handleWechatLogin = async () => { - if (wechatStartingRef.current) return - wechatStartingRef.current = true - setWechatLoading(true) - try { - const result = await getWechatAuthUrl() - // 保存 state 到 localStorage 用于回调时验证 - localStorage.setItem("wechat_state", result.state) - // 记录登录前的来源页,登录成功后跳回 - const from = window.location.pathname + window.location.search - if (from !== "/login" && from !== "/register") { - localStorage.setItem("login_redirect", from) - } else { - localStorage.removeItem("login_redirect") - } - // 跳转到微信授权页 - window.location.href = result.auth_url - } catch (error) { - // 跳走前才可能回到这里;拦截器已弹过后端 detail 时不重复弹, - // 否则透传真实原因(如微信服务未配置、网络异常) - if (!isErrorMsgShown(error)) { - message.error(`微信登录启动失败:${getErrorMessage(error, "请稍后重试")}`) - } - wechatStartingRef.current = false - setWechatLoading(false) + const handleWechatLogin = () => { + // 记录登录前的来源页,登录成功后(弹窗回调)跳回 + const from = window.location.pathname + window.location.search + if (from !== "/login" && from !== "/register") { + localStorage.setItem("login_redirect", from) + } else { + localStorage.removeItem("login_redirect") } - // 成功时 window.location 跳走,不复位 loading(页面即将卸载) + setWechatQrOpen(true) + // 弹窗打开期间按钮 disabled;WxLogin 脚本加载失败/超时时弹窗内会自动回退整页跳转 + } + + // 弹窗扫码登录成功:登录态已由弹窗同步,按用户类型跳转 + const handleWechatQrSuccess = (needOnboarding: boolean) => { + setWechatQrOpen(false) + if (needOnboarding) { + navigate("/welcome/wechat", { replace: true }) + return + } + const redirect = localStorage.getItem("login_redirect") || "/" + localStorage.removeItem("login_redirect") + navigate(redirect, { replace: true }) } return ( @@ -136,10 +127,10 @@ const Login: React.FC = () => { type="button" className="xx-btn-wechat" onClick={handleWechatLogin} - disabled={wechatLoading} + disabled={wechatQrOpen} > 💬 - {wechatLoading ? "加载中..." : "微信登录"} + 微信登录
@@ -147,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 index 61ded7e61..641f0d4bd 100644 --- a/apps/web/src/pages/auth/WechatBindCallback.tsx +++ b/apps/web/src/pages/auth/WechatBindCallback.tsx @@ -1,6 +1,11 @@ /** * 微信绑定回调页(已登录用户在设置页发起"绑定微信"扫码后回到这里) * 用 code 调绑定接口把微信关联到当前账号,成功后回设置页 + * + * 两种运行环境: + * - 整页跳转授权(旧流程/兜底):本页整页加载,成功/失败后 navigate 回设置页 + * - 弹窗内嵌二维码(WxLogin self_redirect):本页在同源 iframe 内加载, + * 结果通过 postMessage 通知父窗口弹窗,不做页面导航 */ import React, { useEffect, useState } from "react" import { useSearchParams, useNavigate } from "react-router-dom" @@ -8,19 +13,30 @@ 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) { - setError("无效的回调参数,请回到设置页重新扫码绑定") + fail("无效的回调参数,请回到设置页重新扫码绑定") return } @@ -32,16 +48,23 @@ const WechatBindCallback: React.FC = () => { 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) { - // 绑定失败直接在本页展示真实原因(如微信已被其他账号绑定),不静默跳走 - setError(`微信绑定失败:${getErrorMessage(err, "请回到设置页重试")}`) + // 绑定失败直接在本页展示/上报真实原因(如微信已被其他账号绑定),不静默跳走 + fail(`微信绑定失败:${getErrorMessage(err, "请回到设置页重试")}`) } } handleBind() - }, [searchParams, navigate, setUser]) + }, [searchParams, navigate, setUser, inIframe]) if (error) { return ( diff --git a/apps/web/src/pages/auth/WechatCallback.tsx b/apps/web/src/pages/auth/WechatCallback.tsx index 754a4ca70..2b7ecbaef 100644 --- a/apps/web/src/pages/auth/WechatCallback.tsx +++ b/apps/web/src/pages/auth/WechatCallback.tsx @@ -2,6 +2,11 @@ * 微信登录回调页 * 扫码授权后由微信重定向回来:用 code 换登录态, * 新用户/资料未完善 → 跳昵称引导页;老用户 → 回来源页/首页 + * + * 两种运行环境: + * - 整页跳转授权(旧流程/兜底):本页整页加载,按上述逻辑导航 + * - 弹窗内嵌二维码(WxLogin self_redirect):本页在同源 iframe 内加载, + * 成功/失败均通过 postMessage 通知父窗口弹窗,不做页面导航 */ import React, { useEffect, useState } from "react" import { useSearchParams, useNavigate } from "react-router-dom" @@ -10,6 +15,7 @@ import { wechatCallback, getCurrentUser, normalizeUser, type User } from "@/api/ import { getErrorMessage } from "@/api/errors" import { useAuthStore } from "@/store/authStore" import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh" +import { isInIframe, postWechatQrResult } from "@/components/auth/WechatQrModal/messages" const WechatCallback: React.FC = () => { const [searchParams] = useSearchParams() @@ -17,24 +23,33 @@ const WechatCallback: React.FC = () => { const setAuth = useAuthStore((state) => state.setAuth) const [loading, setLoading] = useState(true) const [error, setError] = useState(null) + const inIframe = isInIframe() useEffect(() => { const code = searchParams.get("code") const state = searchParams.get("state") - // 微信重定向出错时(如用户拒绝授权 error=access_denied)直接展示原因 + 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(":") - setError(`微信授权失败:${reason}`) - setLoading(false) + fail(`微信授权失败:${reason}`) return } if (!code || !state) { - setError("无效的回调参数,请重新扫码登录") - setLoading(false) + fail("无效的回调参数,请重新扫码登录") return } @@ -61,6 +76,13 @@ const WechatCallback: React.FC = () => { // 新用户 或 资料未完善(如上次中断没填昵称)→ 强制昵称引导 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 @@ -72,13 +94,12 @@ const WechatCallback: React.FC = () => { navigate(redirect, { replace: true }) } catch (err) { // 透传后端真实错误(如 state 过期、code 已消费、接口异常),禁止吞成通用提示 - setError(`微信登录失败:${getErrorMessage(err, "请重试或更换登录方式")}`) - setLoading(false) + fail(`微信登录失败:${getErrorMessage(err, "请重试或更换登录方式")}`) } } handleCallback() - }, [searchParams, navigate, setAuth]) + }, [searchParams, navigate, setAuth, inIframe]) if (loading) { return ( diff --git a/apps/web/src/pages/profile/Settings.tsx b/apps/web/src/pages/profile/Settings.tsx index 7513c3bcf..6a7f766ee 100644 --- a/apps/web/src/pages/profile/Settings.tsx +++ b/apps/web/src/pages/profile/Settings.tsx @@ -8,9 +8,10 @@ import { useSearchParams } from "react-router-dom" import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query" import { message } from "antd" import { Button, Input, Modal } from "@/components/ui" -import { getCurrentUser, updateProfile, getWechatBindUrl, unbindWechat } from "@/api/auth" +import { getCurrentUser, updateProfile, unbindWechat } from "@/api/auth" import { useAuthStore } from "@/store/authStore" import PageHead from "@/components/layout/PageHead" +import WechatQrModal from "@/components/auth/WechatQrModal" import "./ProfileSettings.css" const Settings: React.FC = () => { @@ -19,6 +20,7 @@ const Settings: React.FC = () => { const queryClient = useQueryClient() const [searchParams, setSearchParams] = useSearchParams() const [displayName, setDisplayName] = useState(user?.display_name || "") + const [wechatBindOpen, setWechatBindOpen] = useState(false) const bindTipShownRef = useRef(false) // 拉取最新用户信息(微信绑定状态以后端为准) @@ -62,14 +64,11 @@ const Settings: React.FC = () => { }, }) - const handleBindWechat = async () => { - try { - const result = await getWechatBindUrl() - localStorage.setItem("wechat_bind_state", result.state) - window.location.href = result.auth_url - } catch { - message.error("微信绑定暂不可用,请稍后重试") - } + // 弹窗扫码绑定成功:关闭弹窗,刷新用户信息并提示 + const handleBindSuccess = () => { + setWechatBindOpen(false) + queryClient.invalidateQueries({ queryKey: ["currentUser"] }) + message.success("微信绑定成功") } const unbindMutation = useMutation({ @@ -177,13 +176,20 @@ const Settings: React.FC = () => { 解绑 ) : ( - )}
+ + setWechatBindOpen(false)} + onBindSuccess={handleBindSuccess} + />
) } 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/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/auth/WechatBindCallback.test.tsx b/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx index a2592fc2f..c9cadc752 100644 --- a/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx +++ b/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx @@ -41,6 +41,16 @@ 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( @@ -55,6 +65,7 @@ describe("WechatBindCallback Page", () => { beforeEach(() => { vi.clearAllMocks() + mockIsInIframe.mockReturnValue(false) bindError = null Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k)) mockParams.set("code", "bind_code") @@ -102,4 +113,32 @@ describe("WechatBindCallback Page", () => { 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 5ef33e4c2..7aa16f3c4 100644 --- a/apps/web/src/test/pages/auth/WechatCallback.test.tsx +++ b/apps/web/src/test/pages/auth/WechatCallback.test.tsx @@ -53,6 +53,16 @@ vi.mock("@/store/authStore", () => ({ useAuthStore: (selector: (state: unknown) => unknown) => selector({ setAuth: mockSetAuth }), })) +// 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( @@ -67,6 +77,7 @@ describe("WechatCallback Page", () => { beforeEach(() => { vi.clearAllMocks() + mockIsInIframe.mockReturnValue(false) callbackError = null // 默认正常回调参数;用例可改写 mockParams 模拟 error 重定向 Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k)) @@ -159,4 +170,54 @@ describe("WechatCallback Page", () => { 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"), + }) + }) + }) + }) }) -- 2.54.0 From 06b0bacce1ac7d4cdfa63919c9c58cf5d3d4c8c4 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 11:24:48 +0800 Subject: [PATCH 24/33] =?UTF-8?q?fix(#1718):=20=E6=98=B5=E7=A7=B0=E9=A1=B5?= =?UTF-8?q?=E6=8F=90=E4=BA=A4=E9=98=B2=E8=BF=9E=E7=82=B9=20+=20=E8=A1=A5?= =?UTF-8?q?=20/vite.svg=20favicon=20(#1727)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/public/vite.svg | 4 ++ apps/web/src/pages/auth/WechatOnboarding.tsx | 19 ++++++-- .../test/pages/auth/WechatOnboarding.test.tsx | 48 +++++++++++++++++++ 3 files changed, 68 insertions(+), 3 deletions(-) create mode 100644 apps/web/public/vite.svg 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/pages/auth/WechatOnboarding.tsx b/apps/web/src/pages/auth/WechatOnboarding.tsx index 0b237243f..38c62f345 100644 --- a/apps/web/src/pages/auth/WechatOnboarding.tsx +++ b/apps/web/src/pages/auth/WechatOnboarding.tsx @@ -2,11 +2,12 @@ * 微信新用户昵称引导页 * 新微信用户首次登录后强制填写昵称,完成后才进入主界面 */ -import React from "react" +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" @@ -22,6 +23,10 @@ const WechatOnboarding: React.FC = () => { 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 }), @@ -37,6 +42,8 @@ const WechatOnboarding: React.FC = () => { } const onFinish = async (values: OnboardingFormValues) => { + if (submittingRef.current) return + submittingRef.current = true try { const updated = await saveMutation.mutateAsync(values.display_name.trim()) // 后端返回的 profile_completed 以最新资料为准,前端同步标记完善 @@ -45,9 +52,14 @@ const WechatOnboarding: React.FC = () => { const redirect = localStorage.getItem("login_redirect") || "/app/dashboard" localStorage.removeItem("login_redirect") navigate(redirect, { replace: true }) - } catch { - message.error("保存失败,请重试") + } catch (err) { + // 透传后端真实原因(如接口异常/校验失败);拦截器已弹过的不重复弹 + if (!isErrorMsgShown(err)) { + message.error(`昵称保存失败:${getErrorMessage(err, "请稍后重试")}`) + } + submittingRef.current = false } + // 成功时页面跳走,不复位 } return ( @@ -89,6 +101,7 @@ const WechatOnboarding: React.FC = () => { buttonSize="lg" htmlType="submit" loading={saveMutation.isPending} + disabled={saveMutation.isPending} style={{ width: "100%" }} > {saveMutation.isPending ? "保存中..." : "进入小虾智剪"} diff --git a/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx index 41c4fcfa9..548d7f52e 100644 --- a/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx +++ b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx @@ -122,6 +122,54 @@ describe("WechatOnboarding 昵称引导页", () => { ) }) + 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") -- 2.54.0 From ff60fdf956fab508c641cc7c34e58bb82b687e76 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 11:43:59 +0800 Subject: [PATCH 25/33] =?UTF-8?q?feat(#1714):=20=E5=89=8D=E7=AB=AF=20prepa?= =?UTF-8?q?re=20=E7=9F=AD=E8=B7=AF=EF=BC=88=E5=90=8E=E7=AB=AF=20skip=5Ftra?= =?UTF-8?q?nsfer=20=E5=91=BD=E4=B8=AD=E6=97=B6=E8=B7=B3=E8=BF=87=20OSS=20?= =?UTF-8?q?=E7=9B=B4=E4=BC=A0=EF=BC=89=20(#1729)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/src/api/assets/types.ts | 11 ++ apps/web/src/api/assets/upload.ts | 10 ++ .../src/pages/assets/hooks/useAssetUpload.ts | 14 +++ apps/web/src/test/api/assets.test.ts | 106 ++++++++++++++++++ .../test/pages/assets/useAssetUpload.test.tsx | 29 +++++ 5 files changed, 170 insertions(+) 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 3b08b2914..3a452945d 100644 --- a/apps/web/src/api/assets/upload.ts +++ b/apps/web/src/api/assets/upload.ts @@ -179,6 +179,16 @@ export const uploadAssetDirect = async (data: { 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/pages/assets/hooks/useAssetUpload.ts b/apps/web/src/pages/assets/hooks/useAssetUpload.ts index d322e3d5e..528cf32ce 100644 --- a/apps/web/src/pages/assets/hooks/useAssetUpload.ts +++ b/apps/web/src/pages/assets/hooks/useAssetUpload.ts @@ -141,6 +141,20 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) { })) handlesRef.current.set(item.tempId, h) + // prepare 阶段后端 file_hash 命中素材库已有相同文件(skip_transfer / duplicated): + // 立即标记 done、调一次 refreshList 让已存在素材立即显示,跳过 transfer + complete + if (h.prepared.skip_transfer || h.prepared.duplicated) { + updateItem(item.tempId, { + status: "done", + duplicated: true, + assetId: h.prepared.asset_id, + }) + handlesRef.current.delete(item.tempId) + refreshList() + message.info(`"${item.fileName}" 与素材库已有内容相同,已跳过`) + return + } + if (h.prepared.asset_id) { updateItem(item.tempId, { status: "uploading", diff --git a/apps/web/src/test/api/assets.test.ts b/apps/web/src/test/api/assets.test.ts index 71f5788d7..90967f1f5 100644 --- a/apps/web/src/test/api/assets.test.ts +++ b/apps/web/src/test/api/assets.test.ts @@ -259,6 +259,112 @@ 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) + 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/pages/assets/useAssetUpload.test.tsx b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx index 4c5004fa0..2ad8d634c 100644 --- a/apps/web/src/test/pages/assets/useAssetUpload.test.tsx +++ b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx @@ -26,6 +26,8 @@ interface FakeHandle { fields: Record max_size_bytes: number asset_id: string + duplicated?: boolean + skip_transfer?: boolean } transfer: ReturnType complete: ReturnType @@ -46,6 +48,8 @@ const makeFakeHandle = (opts: { duplicated?: boolean failTransfer?: boolean completeAuto?: boolean + /** prepare 阶段就命中去重:prepare 响应 skip_transfer/duplicated=true */ + prepareDedup?: boolean }) => { const h: FakeHandle = { prepared: { @@ -56,6 +60,8 @@ const makeFakeHandle = (opts: { 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(), @@ -380,4 +386,27 @@ describe("useAssetUpload", () => { 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") + }) + }) }) -- 2.54.0 From 528f56254de8a4ba76adff7ae37e93e9e6f2ff80 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 12:06:58 +0800 Subject: [PATCH 26/33] =?UTF-8?q?feat(#1718):=20PATCH=20/auth/me=20?= =?UTF-8?q?=E8=B5=84=E6=96=99=E6=9B=B4=E6=96=B0=E6=8E=A5=E5=8F=A3=20+=20pr?= =?UTF-8?q?ofile=5Fcompleted=20=E5=AD=97=E6=AE=B5=20(#1728)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- .../versions/068_user_profile_completed.py | 26 ++ apps/api/app/api/routes/auth.py | 84 +++++-- packages/adapters/sqlalchemy_impl/models.py | 1 + .../sqlalchemy_impl/user_repository.py | 2 + .../application/auth/wechat_sync_use_case.py | 2 + packages/domain/entities.py | 2 + tests/unit/test_patch_me_profile_1718.py | 222 ++++++++++++++++++ tests/unit/test_wechat_bind_routes_1719.py | 2 + 8 files changed, 322 insertions(+), 19 deletions(-) create mode 100644 alembic/versions/068_user_profile_completed.py create mode 100644 tests/unit/test_patch_me_profile_1718.py 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 0c268d92d..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 @@ -15,7 +16,7 @@ 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, 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 @@ -85,6 +86,22 @@ class CurrentUserResponse(BaseModel): 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): @@ -274,9 +291,51 @@ async def get_current_user_info( 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 @@ -507,34 +566,20 @@ class WechatBindCompleteRequest(BaseModel): state: str = "" -class WechatBindUserProfile(BaseModel): - """绑定/解绑后返回的用户信息(字段对齐 /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 - - class WechatBindCompleteResponse(BaseModel): success: bool - user: WechatBindUserProfile + user: UserProfileResponse class WechatUnbindResponse(BaseModel): success: bool -def _wechat_user_profile(user) -> WechatBindUserProfile: +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 WechatBindUserProfile( + return UserProfileResponse( user_id=user.id, email=user.email, username=user.username, @@ -544,6 +589,7 @@ def _wechat_user_profile(user) -> WechatBindUserProfile: phone_verified=user.phone_verified, binding_complete=binding_complete, wechat_bound=bool(user.wechat_openid), + profile_completed=user.profile_completed, ) @@ -589,7 +635,7 @@ async def wechat_bind( 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=_wechat_user_profile(result.user)) + return WechatBindCompleteResponse(success=True, user=_user_profile(result.user)) @router.delete("/wechat/bind", response_model=WechatUnbindResponse) diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 0ee3e8565..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)) 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_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/domain/entities.py b/packages/domain/entities.py index 8d1060c87..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)) 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_wechat_bind_routes_1719.py b/tests/unit/test_wechat_bind_routes_1719.py index 37df5fada..9f07cc43d 100644 --- a/tests/unit/test_wechat_bind_routes_1719.py +++ b/tests/unit/test_wechat_bind_routes_1719.py @@ -38,6 +38,7 @@ def _auth_user(user_id="u-1", openid=None): display_name="用户", phone="", phone_verified=False, + profile_completed=True, ) return SimpleNamespace(user=user, session_id="s-1", token_type="user_auth") @@ -112,6 +113,7 @@ def test_bind_success_returns_user_with_wechat_bound(): display_name="用户", phone="", phone_verified=False, + profile_completed=True, ) import packages.application.auth.wechat_oauth_service as oauth_mod -- 2.54.0 From 70dde8cbfb82762c1b92a8ed26ee5b37fc62bfe7 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 12:31:38 +0800 Subject: [PATCH 27/33] =?UTF-8?q?feat(#1714):=20prepare=5Fdirect=5Fupload?= =?UTF-8?q?=20=E5=8E=BB=E9=87=8D=20+=20=E9=A2=84=E5=BB=BA=20PROCESSING=20a?= =?UTF-8?q?sset=20=E5=8D=A0=E4=BD=8D=20(#1730)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/api/app/api/routes/upload.py | 151 ++++++++- apps/api/app/schemas/upload.py | 4 + tests/unit/test_prepare_dedup_1714.py | 471 ++++++++++++++++++++++++++ 3 files changed, 614 insertions(+), 12 deletions(-) create mode 100644 tests/unit/test_prepare_dedup_1714.py diff --git a/apps/api/app/api/routes/upload.py b/apps/api/app/api/routes/upload.py index 48b367a0e..279888a78 100644 --- a/apps/api/app/api/routes/upload.py +++ b/apps/api/app/api/routes/upload.py @@ -165,15 +165,44 @@ def _find_duplicate_asset( within_minutes=FALLBACK_DEDUP_WINDOW_MINUTES, file_size=file_size or 0, ) - if existing is not None and getattr(existing, "status", None) in ACTIVE_ASSET_STATUSES: - logger.info( - "素材幂等兜底命中(近期活动同名记录): library=%s name=%s asset=%s status=%s", - library_id, - filename, - getattr(existing, "id", "?"), - getattr(existing, "status", "?"), - ) - return existing + # 兜底去重:按状态区分处理 + # - READY/ERROR:稳定素材,总命中(避免重复创建) + # - PROCESSING/UPLOADING:预建或 complete 占位,仅当 hash 一致才命中 + # - 占位无 hash(旧客户端 complete 建的)→ 命中 + # - 占位有 hash 且与当前请求 hash 一致 → 命中 + # - 占位有 hash 且与当前请求 hash 不同 → 跳过(内容不同) + if existing is not None: + status = getattr(existing, "status", None) + existing_hash = getattr(existing, "file_hash", "") or "" + if status in (AssetStatus.READY, AssetStatus.ERROR): + logger.info( + "素材幂等兜底命中(近期同名稳定记录): library=%s name=%s asset=%s status=%s", + library_id, + filename, + getattr(existing, "id", "?"), + status, + ) + return existing + elif status in ACTIVE_ASSET_STATUSES: + if existing_hash and file_hash and existing_hash != file_hash: + logger.debug( + "素材兜底去重跳过(占位 hash 不同): library=%s name=%s asset=%s hash=%s req_hash=%s", + library_id, + filename, + getattr(existing, "id", "?"), + existing_hash, + file_hash, + ) + existing = None + else: + logger.info( + "素材幂等兜底命中(近期同名活动记录): library=%s name=%s asset=%s status=%s", + library_id, + filename, + getattr(existing, "id", "?"), + status, + ) + return existing return None @@ -187,8 +216,44 @@ def _create_pending_asset( 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, @@ -199,6 +264,7 @@ def _create_pending_asset( 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) @@ -243,9 +309,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: @@ -264,8 +336,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( @@ -284,6 +387,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"]), @@ -291,6 +415,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, ) diff --git a/apps/api/app/schemas/upload.py b/apps/api/app/schemas/upload.py index bc606649c..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,6 +26,9 @@ 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): 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 -- 2.54.0 From 9bca7e53e3f177af918e8a55bc2b9c9f1f3bd617 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 12:44:05 +0800 Subject: [PATCH 28/33] =?UTF-8?q?feat(#1714):=20complete=20=E8=AF=B7?= =?UTF-8?q?=E6=B1=82=E6=90=BA=E5=B8=A6=20file=5Fsize=EF=BC=8C=E4=BF=AE?= =?UTF-8?q?=E5=A4=8D=E5=90=8C=E5=90=8D=E5=85=9C=E5=BA=95=E8=AF=AF=E6=9D=80?= =?UTF-8?q?=E6=96=B0=E8=A7=86=E9=A2=91=20(#1731)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/web/src/api/assets/upload.ts | 4 ++++ apps/web/src/test/api/assets.test.ts | 16 ++++++++++++++++ 2 files changed, 20 insertions(+) diff --git a/apps/web/src/api/assets/upload.ts b/apps/web/src/api/assets/upload.ts index 3a452945d..490b08df7 100644 --- a/apps/web/src/api/assets/upload.ts +++ b/apps/web/src/api/assets/upload.ts @@ -32,6 +32,8 @@ export const completeDirectUpload = async (data: { file_hash?: string /** 前端上传幂等 token(与 prepare 一致),同一次上传重发 complete 不重复建记录 */ client_upload_id?: string + /** 文件字节数;后端同名兜底去重需用它做大小校验,缺失(=0)时同名记录一律不判重 */ + file_size?: number }): Promise => { // complete 内含 OSS 存在性检查 + 建库 + 派单,放宽到 60s; // 超时不代表失败(记录可能已建成),调用方禁止超时后盲目重传整个文件 @@ -156,6 +158,8 @@ export const prepareDirectUploadHandle = async (data: { storage_key: prepared.storage_key, file_hash: data.fileHash, client_upload_id: data.clientUploadId, + // 透传文件字节数:后端同名兜底去重依赖大小校验,缺省会导致同名新视频被误判重复 + file_size: data.file.size, }), } } diff --git a/apps/web/src/test/api/assets.test.ts b/apps/web/src/test/api/assets.test.ts index 90967f1f5..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")) @@ -354,6 +368,8 @@ describe("assets API", () => { ([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) { -- 2.54.0 From 9a289e1e1fc8d44030a0eacaa02f052aac05b26d Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 13:35:35 +0800 Subject: [PATCH 29/33] =?UTF-8?q?fix:=20=E5=8F=91=E7=89=88=E5=90=8E?= =?UTF-8?q?=E6=87=92=E5=8A=A0=E8=BD=BD=20chunk=20=E5=A4=B1=E6=95=88?= =?UTF-8?q?=E7=99=BD=E5=B1=8F=E2=80=94=E2=80=94ErrorBoundary=20=E8=87=AA?= =?UTF-8?q?=E5=8A=A8=E5=88=B7=E6=96=B0=20+=20lazy=20=E9=87=8D=E8=AF=95=20(?= =?UTF-8?q?#1732)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- .../components/common/ChunkErrorBoundary.tsx | 85 +++++++++++ apps/web/src/main.tsx | 5 +- apps/web/src/router/appRoutes.tsx | 141 ++++-------------- apps/web/src/router/lazyRoute.ts | 40 +++++ .../components/ChunkErrorBoundary.test.tsx | 79 ++++++++++ apps/web/src/test/router/lazyRoute.test.ts | 44 ++++++ .../web/src/test/utils/chunkLoadError.test.ts | 80 ++++++++++ apps/web/src/utils/chunkLoadError.ts | 84 +++++++++++ 8 files changed, 445 insertions(+), 113 deletions(-) create mode 100644 apps/web/src/components/common/ChunkErrorBoundary.tsx create mode 100644 apps/web/src/router/lazyRoute.ts create mode 100644 apps/web/src/test/components/ChunkErrorBoundary.test.tsx create mode 100644 apps/web/src/test/router/lazyRoute.test.ts create mode 100644 apps/web/src/test/utils/chunkLoadError.test.ts create mode 100644 apps/web/src/utils/chunkLoadError.ts 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/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/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/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 = "/" +} -- 2.54.0 From 54916aff86b59e95191719965b89f7dffbaf361f Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 13:41:26 +0800 Subject: [PATCH 30/33] =?UTF-8?q?fix(#1714):=20=E5=90=8C=E5=90=8D=E5=85=9C?= =?UTF-8?q?=E5=BA=95=E5=8E=BB=E9=87=8D=E8=AF=AF=E6=9D=80=E6=96=B0=E8=A7=86?= =?UTF-8?q?=E9=A2=91=20+=20ingest=20=E9=93=BE=E8=B7=AF=E5=AD=A4=E5=84=BF?= =?UTF-8?q?=E6=B8=85=E7=90=86/=E5=90=AF=E5=8A=A8=E6=81=A2=E5=A4=8D=20(#173?= =?UTF-8?q?4)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/api/app/api/routes/upload.py | 72 ++-- apps/worker/worker_app/celery_app.py | 13 + apps/worker/worker_app/tasks/_startup.py | 29 ++ apps/worker/worker_app/tasks/cleanup.py | 56 ++++ .../sqlalchemy_impl/asset_repository.py | 9 +- packages/application/ingest_orphan_cleanup.py | 310 ++++++++++++++++++ .../test_asset_repo_fallback_dedup_1714.py | 93 ++++++ tests/unit/test_cleanup_ingest_beat_1714.py | 75 +++++ tests/unit/test_ingest_orphan_cleanup_1714.py | 262 +++++++++++++++ .../test_upload_complete_idempotency_1714.py | 94 +++++- 10 files changed, 963 insertions(+), 50 deletions(-) create mode 100644 packages/application/ingest_orphan_cleanup.py create mode 100644 tests/unit/test_asset_repo_fallback_dedup_1714.py create mode 100644 tests/unit/test_cleanup_ingest_beat_1714.py create mode 100644 tests/unit/test_ingest_orphan_cleanup_1714.py diff --git a/apps/api/app/api/routes/upload.py b/apps/api/app/api/routes/upload.py index 279888a78..08aed4e7b 100644 --- a/apps/api/app/api/routes/upload.py +++ b/apps/api/app/api/routes/upload.py @@ -108,9 +108,8 @@ def _infer_mime_type_from_storage_key(storage_key: str) -> str: return "video/mp4" # default -# 兜底去重:无 file_hash / client_upload_id 时,同库同名近期活动记录视为重复 +# 兜底去重:无 file_hash / client_upload_id 且大小已知时,同库同名同大小近期活动记录视为重复 FALLBACK_DEDUP_WINDOW_MINUTES = 30 -ACTIVE_ASSET_STATUSES = (AssetStatus.UPLOADING, AssetStatus.PROCESSING) def _find_duplicate_asset( @@ -126,8 +125,12 @@ def _find_duplicate_asset( 1. client_upload_id(客户端幂等 token,同一次上传的重试保持一致) 2. file_hash(内容哈希,不同上传只要内容相同即去重) - 3. 兜底:同库 + 同文件名(+同大小)且 30 分钟内仍处 uploading/processing - 的记录——旧客户端不传 hash/token 时,防止 complete 超时重试反复建占位。 + 3. 兜底(严格模式,宁可漏判不可误杀):file_hash 与 client_upload_id + 均缺失、且 file_size > 0 时,同库 + 同文件名 + **同大小** 且 30 分钟内 + 仍处 uploading/processing 的记录才判重。 + - file_hash 非空时跳过兜底(hash 已代表内容;同名但内容全新的视频 + 如 iPhone 的 IMG_xxxx.MOV 绝不能被同名占位误杀) + - file_size=0(未知)时不允许仅凭同名 + processing 判重,直接放行 全部为鸭子类型调用:旧仓储无对应方法时静默跳过,不破坏既有实现。 """ @@ -156,53 +159,35 @@ def _find_duplicate_asset( existing.id, ) return existing - if filename: + # 同名兜底去重(最后防线,严格模式): + # - 仅当 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 or 0, + file_size=file_size, ) - # 兜底去重:按状态区分处理 - # - READY/ERROR:稳定素材,总命中(避免重复创建) - # - PROCESSING/UPLOADING:预建或 complete 占位,仅当 hash 一致才命中 - # - 占位无 hash(旧客户端 complete 建的)→ 命中 - # - 占位有 hash 且与当前请求 hash 一致 → 命中 - # - 占位有 hash 且与当前请求 hash 不同 → 跳过(内容不同) if existing is not None: - status = getattr(existing, "status", None) - existing_hash = getattr(existing, "file_hash", "") or "" - if status in (AssetStatus.READY, AssetStatus.ERROR): - logger.info( - "素材幂等兜底命中(近期同名稳定记录): library=%s name=%s asset=%s status=%s", - library_id, - filename, - getattr(existing, "id", "?"), - status, - ) - return existing - elif status in ACTIVE_ASSET_STATUSES: - if existing_hash and file_hash and existing_hash != file_hash: - logger.debug( - "素材兜底去重跳过(占位 hash 不同): library=%s name=%s asset=%s hash=%s req_hash=%s", - library_id, - filename, - getattr(existing, "id", "?"), - existing_hash, - file_hash, - ) - existing = None - else: - logger.info( - "素材幂等兜底命中(近期同名活动记录): library=%s name=%s asset=%s status=%s", - library_id, - filename, - getattr(existing, "id", "?"), - status, - ) - return existing + 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 @@ -487,6 +472,7 @@ async def complete_direct_upload( 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( diff --git a/apps/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py index 75477f583..34d40ad09 100755 --- a/apps/worker/worker_app/celery_app.py +++ b/apps/worker/worker_app/celery_app.py @@ -17,6 +17,12 @@ 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", @@ -48,4 +54,11 @@ celery_app.conf.beat_schedule = { "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": 540}, + }, } diff --git a/apps/worker/worker_app/tasks/_startup.py b/apps/worker/worker_app/tasks/_startup.py index 2d3f80b7c..afe329704 100644 --- a/apps/worker/worker_app/tasks/_startup.py +++ b/apps/worker/worker_app/tasks/_startup.py @@ -257,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 07a7de8dc..ff29840dc 100644 --- a/apps/worker/worker_app/tasks/cleanup.py +++ b/apps/worker/worker_app/tasks/cleanup.py @@ -16,6 +16,12 @@ from worker_app.tasks._startup import ( 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__) @@ -69,3 +75,53 @@ def scheduled_cleanup_stale_running(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_M 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/packages/adapters/sqlalchemy_impl/asset_repository.py b/packages/adapters/sqlalchemy_impl/asset_repository.py index ba334d498..fad219727 100755 --- a/packages/adapters/sqlalchemy_impl/asset_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_repository.py @@ -488,20 +488,25 @@ class SQLAlchemyAssetRepository: 用于旧客户端未传 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, ) - if file_size and file_size > 0: - query = query.filter(AssetModel.file_size == file_size) model = query.order_by(AssetModel.created_at.desc()).first() if model is None: return None 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/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_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_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_upload_complete_idempotency_1714.py b/tests/unit/test_upload_complete_idempotency_1714.py index 1effb6e78..c0436b065 100644 --- a/tests/unit/test_upload_complete_idempotency_1714.py +++ b/tests/unit/test_upload_complete_idempotency_1714.py @@ -73,6 +73,9 @@ class StubAssetRepository: 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 @@ -81,7 +84,7 @@ class StubAssetRepository: and a.name == name and a.status in (AssetStatus.UPLOADING, AssetStatus.PROCESSING) and a.created_at >= cutoff - and (not file_size or a.file_size == file_size) + and a.file_size == file_size ] return max(candidates, key=lambda a: a.created_at) if candidates else None @@ -234,14 +237,21 @@ class TestDirectCompleteIdempotency: 不应再建第二条。 """ client, asset_repo, ingest_repo, _ = _client() - # 第一次 complete(旧客户端无 token/hash) - r1 = client.post("/api/v1/direct/complete", json=COMPLETE_BODY) + # 第一次 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),且近期 + # 前端重试传的是同一个 File),且近期;同大小才允许兜底命中 r2 = client.post( "/api/v1/direct/complete", - json={**COMPLETE_BODY, "storage_key": "uploads/retry/IMG_2282.MOV", "file_size": 0}, + 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 @@ -249,6 +259,80 @@ class TestDirectCompleteIdempotency: 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( -- 2.54.0 From 2ce3a5efd3188f30cbf8e18a86f03bc1ce7e095a Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 14:05:44 +0800 Subject: [PATCH 31/33] =?UTF-8?q?fix:=20SPA=20=E8=B7=AF=E7=94=B1=E5=9B=9E?= =?UTF-8?q?=E9=80=80=E7=9A=84=20HTML=20=E8=A1=A5=20no-cache=20=E5=A4=B4?= =?UTF-8?q?=EF=BC=88=E9=85=8D=E5=90=88=20#1732=20=E7=99=BD=E5=B1=8F?= =?UTF-8?q?=E4=BF=AE=E5=A4=8D=EF=BC=89=20(#1735)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- deploy/configs/nginx-production.conf | 4 ++++ deploy/configs/nginx-staging.conf | 4 ++++ infra/docker/nginx-production.conf | 4 ++++ infra/docker/nginx-staging.conf | 4 ++++ infra/docker/nginx.conf | 4 ++++ 5 files changed, 20 insertions(+) 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/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 -- 2.54.0 From dddc1cd08123d42ed20db93a2afea20b3903cf9b Mon Sep 17 00:00:00 2001 From: saas-backend-agent Date: Sun, 6 Sep 2026 14:13:43 +0800 Subject: [PATCH 32/33] =?UTF-8?q?fix(#1714):=20beat=20=E8=B0=83=E5=BA=A6?= =?UTF-8?q?=E6=96=87=E4=BB=B6=E6=94=B9=E7=94=A8=20/tmp=20=E8=B7=AF?= =?UTF-8?q?=E5=BE=84=EF=BC=8C=E4=BF=AE=E5=A4=8D=20celery=20=E7=94=A8?= =?UTF-8?q?=E6=88=B7=E6=97=A0=20CWD=20=E5=86=99=E6=9D=83=E9=99=90=E5=AF=BC?= =?UTF-8?q?=E8=87=B4=20beat=20=E5=B4=A9=E6=BA=83=E3=80=81=E5=AE=9A?= =?UTF-8?q?=E6=97=B6=E5=B7=A1=E6=A3=80=EF=BC=88ingest=20=E5=AD=A4=E5=84=BF?= =?UTF-8?q?=E6=B8=85=E7=90=86=EF=BC=89=E4=BB=8E=E6=9C=AA=E6=89=A7=E8=A1=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- infra/docker/entrypoint-worker.sh | 1 + 1 file changed, 1 insertion(+) diff --git a/infra/docker/entrypoint-worker.sh b/infra/docker/entrypoint-worker.sh index 2f20644b2..0a74b474a 100755 --- a/infra/docker/entrypoint-worker.sh +++ b/infra/docker/entrypoint-worker.sh @@ -36,6 +36,7 @@ celery \ worker \ --loglevel=info \ "-B" \ + -s /tmp/celerybeat-schedule \ -Q generation \ "--concurrency=${GEN_CONCURRENCY}" \ "--max-tasks-per-child=${MAX_TASKS}" \ -- 2.54.0 From e148f995a86adaa25ae824e94b0640d9fd826037 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 6 Sep 2026 14:47:03 +0800 Subject: [PATCH 33/33] =?UTF-8?q?fix(#1737):=20=E7=94=9F=E6=88=90=E9=A1=B5?= =?UTF-8?q?=E6=A0=87=E9=A2=98=E6=A1=86=E8=81=9A=E7=84=A6=E5=8D=B3=E5=B1=95?= =?UTF-8?q?=E5=BC=80=E6=A0=87=E9=A2=98=E5=BA=93=E5=88=97=E8=A1=A8=20+=20?= =?UTF-8?q?=E4=B8=8B=E6=8B=89=E7=AE=AD=E5=A4=B4=E6=8F=90=E7=A4=BA=20(#1739?= =?UTF-8?q?)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- .../components/Step4TitleSettings.tsx | 30 ++--- .../title/TitleLibraryAutoComplete.tsx | 75 +++++++++++ .../title-library-autocomplete.test.tsx | 126 ++++++++++++++++++ 3 files changed, 210 insertions(+), 21 deletions(-) create mode 100644 apps/web/src/pages/generate/components/title/TitleLibraryAutoComplete.tsx create mode 100644 apps/web/src/test/pages/generate/title-library-autocomplete.test.tsx diff --git a/apps/web/src/pages/generate/components/Step4TitleSettings.tsx b/apps/web/src/pages/generate/components/Step4TitleSettings.tsx index 064bc167a..6f8500519 100644 --- a/apps/web/src/pages/generate/components/Step4TitleSettings.tsx +++ b/apps/web/src/pages/generate/components/Step4TitleSettings.tsx @@ -9,12 +9,13 @@ * - 标题样式(字体/颜色/位置/大小/粗斜描边/预设):全局统一 */ import React, { useMemo, useState } from "react" -import { AutoComplete, Input, message } from "antd" +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" @@ -201,21 +202,14 @@ const Step4TitleSettings: React.FC = (props) => {
- { t.updateTitle(val || "") onPreviewTitlesChange?.([val || ""]) }} options={titleOptions} - filterOption={(inputValue, option) => { - const title = (option?.label || option?.value || "") as string - return title.toLowerCase().includes((inputValue || "").toLowerCase()) - }} />
@@ -269,17 +263,11 @@ const Step4TitleSettings: React.FC = (props) => { {Array.from({ length: previewCount }, (_, i) => (
- updateVariantTitle(i, val || "")} + updateVariantTitle(i, val)} options={titleOptions} - filterOption={(inputValue, option) => { - const title = (option?.label || option?.value || "") as string - return title.toLowerCase().includes((inputValue || "").toLowerCase()) - }} />
))} diff --git a/apps/web/src/pages/generate/components/title/TitleLibraryAutoComplete.tsx b/apps/web/src/pages/generate/components/title/TitleLibraryAutoComplete.tsx new file mode 100644 index 000000000..72cdd8ed6 --- /dev/null +++ b/apps/web/src/pages/generate/components/title/TitleLibraryAutoComplete.tsx @@ -0,0 +1,75 @@ +/** + * 标题库 AutoComplete(Issue #1737) + * + * 原生 antd AutoComplete(combobox 模式)的两个行为不符合产品预期: + * 1. combobox 默认 showAction=[],输入框聚焦时下拉不展开——用户必须先打字才能看到标题库, + * 且组件无下拉箭头,视觉上是"纯输入框",不知道标题库里已有标题可选。 + * 2. 空态聚焦不展示任何标题库内容。 + * + * 本组件封装修复: + * - 受控 open:聚焦(且标题库非空)即展开,展示全部标题;失焦/选中/Esc 关闭 + * (rc-select 失焦会主动 onToggleOpen(false),onOpenChange 同步状态即可,不会死循环) + * - suffixIcon 加下拉三角,视觉提示"可选择";有值时 allowClear 的清除按钮照常出现 + * - 输入文字时由 filterOption 过滤(空串展示全部) + * - 保留 combobox 自由输入能力:用户可输入标题库之外的自定义标题 + */ +import React, { useState } from "react" +import { AutoComplete } from "antd" +import { DownOutlined } from "@ant-design/icons" +import type { AutoCompleteProps } from "antd" + +export interface TitleOption { + label: string + value: string +} + +interface TitleLibraryAutoCompleteProps { + value: string + onChange: (val: string) => void + options: TitleOption[] + placeholder?: string + allowClear?: boolean + maxLength?: number + style?: React.CSSProperties +} + +const TitleLibraryAutoComplete: React.FC = ({ + value, + onChange, + options, + placeholder = "输入或从标题库选择", + allowClear = true, + maxLength = 50, + style, +}) => { + const [open, setOpen] = useState(false) + const hasTitles = options.length > 0 + + const filterOption: AutoCompleteProps["filterOption"] = (inputValue, option) => { + const title = (option?.label || option?.value || "") as string + return title.toLowerCase().includes((inputValue || "").toLowerCase()) + } + + return ( + onChange(val || "")} + options={options} + filterOption={filterOption} + open={open} + onOpenChange={setOpen} + onFocus={() => { + // 标题库为空时不展开(避免弹出"暂无数据"空壳) + if (hasTitles) setOpen(true) + }} + onSelect={() => setOpen(false)} + suffixIcon={} + placeholder={placeholder} + allowClear={allowClear} + maxLength={maxLength} + style={{ width: "100%", ...style }} + /> + ) +} + +export default TitleLibraryAutoComplete 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("已有标题") + }) +}) -- 2.54.0