Release: 同步 develop 到 main(含前端素材库自动创建 + 数据库 migration) #1629

Open
xiaoxia wants to merge 29 commits from develop into main
139 changed files with 12966 additions and 1600 deletions
+2
View File
@@ -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..."
@@ -0,0 +1,34 @@
"""add client_upload_id to assets and asset_id to ingest_jobs
Issue #1714:上传 complete 幂等 + worker 转码回写关联。
- assets.client_upload_id:客户端幂等 tokencomplete 去重)
- ingest_jobs.asset_idcomplete 阶段创建的占位 asset idworker 回写关联,
防止 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")
@@ -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")
@@ -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")
+186 -2
View File
@@ -1,4 +1,5 @@
"""
from __future__ import annotations
Canonical authentication API routes.
The route layer is intentionally thin: repository construction lives in
@@ -13,9 +14,9 @@ import jwt
from app.auth import AuthenticatedUser, blacklist_token, get_current_user
from app.config import settings
from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository
from fastapi import APIRouter, Depends, Header, HTTPException, status
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, EmailStr
from pydantic import BaseModel, EmailStr, field_validator
from packages.adapters.redis import NoopSessionStore
from packages.adapters.smtp import NoopEmailService
@@ -84,6 +85,23 @@ class CurrentUserResponse(BaseModel):
phone: str = ""
phone_verified: bool = False
binding_complete: bool = False
wechat_bound: bool = False
profile_completed: bool = True
class UserProfileResponse(BaseModel):
"""用户资料负载(PATCH /me、绑定/解绑接口复用;字段与 GET /auth/me 一致,前端 normalizeUser 直接消费)"""
user_id: str
email: str
username: str
display_name: str
email_verified: bool
phone: str = ""
phone_verified: bool = False
binding_complete: bool = False
wechat_bound: bool = False
profile_completed: bool = True
class PasswordResetRequestModel(BaseModel):
@@ -272,9 +290,52 @@ async def get_current_user_info(
phone=user.phone or "",
phone_verified=user.phone_verified,
binding_complete=binding_complete,
wechat_bound=bool(user.wechat_openid),
profile_completed=user.profile_completed,
)
class UpdateProfileRequest(BaseModel):
"""更新个人资料请求(当前仅支持昵称)"""
display_name: str
@field_validator("display_name")
@classmethod
def _validate_display_name(cls, v: str) -> str:
name = (v or "").strip()
if not name:
raise ValueError("昵称不能为空白")
if len(name) > 20:
raise ValueError("昵称长度需在 1-20 个字符之间")
return name
class UpdateProfileResponse(BaseModel):
"""更新资料响应:前端 normalizeUser(response.user) 直接消费"""
user: UserProfileResponse
@router.patch("/me", response_model=UpdateProfileResponse)
async def update_current_user_profile(
request: UpdateProfileRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> UpdateProfileResponse:
"""更新当前登录用户昵称(微信新用户首次设置昵称后置 profile_completed=True)。"""
user = current_user.user
user.display_name = request.display_name # 已 stripvalidator
if not user.profile_completed:
user.profile_completed = True
user_repository.save(user)
logger.info("[资料更新] 用户 %s 更新昵称,profile_completed=%s", user.id, user.profile_completed)
# 重新读取,确保返回的是持久化后的最新状态
fresh = user_repository.find_by_id(user.id) or user
return UpdateProfileResponse(user=_user_profile(fresh))
class _NoopSessionStore(NoopSessionStore):
pass
@@ -426,6 +487,7 @@ async def get_wechat_auth_url() -> WechatAuthUrlResponse:
@router.post("/wechat/callback", response_model=WechatLoginResponse)
async def wechat_callback(
request: WechatCallbackRequest,
http_request: Request,
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatLoginResponse:
"""微信登录回调处理"""
@@ -433,11 +495,30 @@ async def wechat_callback(
from packages.application.auth.wechat_sync_use_case import WechatSyncRequest as SyncRequest
from packages.application.auth.wechat_sync_use_case import WechatSyncUseCase
# 回调可观测性:记录 UA(区分微信内置浏览器 MicroMessenger)与 state
# 便于排查"停留 open.weixin.qq.com / 回调失败"类问题(#1718
user_agent = http_request.headers.get("User-Agent", "")
is_wechat_browser = "MicroMessenger" in user_agent
logger.info(
"[微信回调] 收到回调: state=%s code_len=%d UA=%r 微信内置浏览器=%s",
(request.state or "")[:8],
len(request.code or ""),
user_agent[:200],
is_wechat_browser,
)
# 1. 用 code 换微信用户信息
oauth_service = get_wechat_oauth_service()
wechat_user, err = oauth_service.handle_callback(request.code, request.state)
if err:
# state 校验失败 / 微信 errcode 等错误原文已在 service 内 log,这里带上 UA 上下文
logger.warning("[微信回调] 处理失败: err=%s 微信内置浏览器=%s", err, is_wechat_browser)
raise HTTPException(status_code=400, detail=err)
logger.info(
"[微信回调] state 校验通过,微信用户信息获取成功: openid=%s unionid=%s",
wechat_user.openid[:8] if wechat_user.openid else "",
bool(wechat_user.unionid),
)
# 2. 同步登录/注册(复用 wechat-sync 逻辑)
use_case = WechatSyncUseCase(user_repository=user_repository)
@@ -472,6 +553,109 @@ async def wechat_callback(
)
# ==================== 微信账号绑定/解绑(已登录用户) ====================
class WechatBindUrlResponse(BaseModel):
auth_url: str
state: str
class WechatBindCompleteRequest(BaseModel):
code: str
state: str = ""
class WechatBindCompleteResponse(BaseModel):
success: bool
user: UserProfileResponse
class WechatUnbindResponse(BaseModel):
success: bool
def _user_profile(user) -> UserProfileResponse:
binding_complete = bool(
user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email
)
return UserProfileResponse(
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
email_verified=user.email_verified,
phone=user.phone or "",
phone_verified=user.phone_verified,
binding_complete=binding_complete,
wechat_bound=bool(user.wechat_openid),
profile_completed=user.profile_completed,
)
@router.get("/wechat/bind/url", response_model=WechatBindUrlResponse)
async def get_wechat_bind_url(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> WechatBindUrlResponse:
"""获取微信绑定授权链接(已登录用户场景)。state 经 Redis 存储做 CSRF 校验。"""
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
oauth_service = get_wechat_oauth_service()
auth_url, state = oauth_service.generate_auth_url()
logger.info("[微信绑定] 用户 %s 请求绑定授权链接", current_user.user.id)
return WechatBindUrlResponse(auth_url=auth_url, state=state)
@router.post("/wechat/bind", response_model=WechatBindCompleteResponse)
async def wechat_bind(
request: WechatBindCompleteRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatBindCompleteResponse:
"""微信绑定完成:扫码回调后用 code 换 openid,绑定到当前登录账号(不创建新用户)。"""
from packages.application.auth.wechat_bind_use_case import WechatBindRequest, WechatBindUseCase
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
oauth_service = get_wechat_oauth_service()
wechat_user, err = oauth_service.handle_callback(request.code, request.state)
if err:
logger.warning("[微信绑定] 用户 %s 换取微信信息失败: %s", current_user.user.id, err)
raise HTTPException(status_code=400, detail=err)
use_case = WechatBindUseCase(user_repository=user_repository)
result, error, http_status = use_case.bind(
WechatBindRequest(
user_id=current_user.user.id,
openid=wechat_user.openid,
unionid=wechat_user.unionid or "",
)
)
if error:
logger.warning("[微信绑定] 用户 %s 绑定失败: %s", current_user.user.id, error)
raise HTTPException(status_code=http_status, detail=error)
logger.info("[微信绑定] 用户 %s 绑定成功 openid=%s", current_user.user.id, wechat_user.openid[:8])
return WechatBindCompleteResponse(success=True, user=_user_profile(result.user))
@router.delete("/wechat/bind", response_model=WechatUnbindResponse)
async def wechat_unbind(
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatUnbindResponse:
"""解绑微信:需账号仍有其他登录方式(密码/手机/真实邮箱),否则拒绝。"""
from packages.application.auth.wechat_bind_use_case import WechatUnbindUseCase
use_case = WechatUnbindUseCase(user_repository=user_repository)
result, error, http_status = use_case.unbind(current_user.user.id)
if error:
logger.warning("[微信解绑] 用户 %s 解绑失败: %s", current_user.user.id, error)
raise HTTPException(status_code=http_status, detail=error)
logger.info("[微信解绑] 用户 %s 解绑成功", current_user.user.id)
return WechatUnbindResponse(success=True)
# ==================== 验证码 & 绑定 ====================
+3 -1
View File
@@ -14,6 +14,7 @@ from typing import Any
from uuid import uuid4
from app.api.routes._helpers import require_project_and_library
from app.api.routes.upload import _persist_celery_task_id
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service
@@ -381,7 +382,8 @@ async def complete_chunked_upload(
file_hash=request.file_hash,
)
)
celery_app.send_task("worker.ingest_asset", args=[job.id])
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
# Update metadata status
meta["status"] = "completed"
+246 -123
View File
@@ -14,6 +14,7 @@ from app.core.task_enqueue import (
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task,
)
from app.dependencies import (
@@ -23,6 +24,7 @@ from app.dependencies import (
get_generation_task_repository,
)
from app.schemas.generation_task import (
BatchPreviewGenerationTaskResponse,
CreatePreviewGenerationTaskRequest,
PreviewGenerationTaskResponse,
)
@@ -193,11 +195,19 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
if started_at and completed_at:
generate_duration = (completed_at - started_at).total_seconds()
title_cfg = getattr(task, "title_config", None)
title_cfg = title_cfg if isinstance(title_cfg, dict) else {}
extra_meta = getattr(task, "extra_meta", None)
extra_meta = extra_meta if isinstance(extra_meta, dict) else {}
voice_library_id = getattr(task, "voice_library_id", "") or ""
if not isinstance(voice_library_id, str):
voice_library_id = str(voice_library_id) if voice_library_id else ""
return PreviewGenerationTaskResponse(
task_id=task.id,
status=task.status.value if hasattr(task.status, "value") else str(task.status),
progress=float(task.progress or 0.0),
is_preview=bool(getattr(task, "is_preview", True)),
variant_index=int(extra_meta.get("variant_index", 0) or 0),
resolution=getattr(task, "resolution", "") or "",
video_url=video_url,
duration=duration,
@@ -206,6 +216,8 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
transition_count=transition_count,
material_usage=material_usage,
error_message=task.error_message or "",
title_text=str(title_cfg.get("text", "") or ""),
voice_library_id=voice_library_id,
created_at=task.created_at,
started_at=started_at,
finished_at=completed_at,
@@ -213,50 +225,100 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
)
@router.post("/preview", response_model=PreviewGenerationTaskResponse, status_code=201)
def _resolve_preview_edit_plan_id(
*,
request: CreatePreviewGenerationTaskRequest,
task,
db: Session,
user_id: str,
) -> str:
"""确定任务关联的编辑计划ID:优先前端传入,否则按 template_id+user 兜底查找。"""
if task.source_edit_plan_id:
return task.source_edit_plan_id
if not request.template_id:
return ""
try:
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
)
_plan_repo = SQLAlchemyEditPlanRepository(db)
_plans = _plan_repo.list_by_template(request.template_id, limit=20)
for _p in _plans:
if (_p.created_by_user_id or "") == user_id:
logger.info(
"[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s",
task.id,
_p.id,
)
return _p.id
except Exception:
logger.warning(
"[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s",
task.id,
exc_info=True,
)
return ""
def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
"""从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。"""
if not values:
return fallback
if len(values) == 1:
return values[0]
return values[index] if index < len(values) else fallback
@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201)
def create_preview_generation_task(
request: CreatePreviewGenerationTaskRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generation_task_repository=Depends(get_generation_task_repository),
db: Session = Depends(get_db_session),
asset_repo=Depends(get_asset_repository),
) -> PreviewGenerationTaskResponse:
"""创建预览生成任务。
) -> BatchPreviewGenerationTaskResponse:
"""创建预览生成任务(支持批量)
预览渲染品质与正式生成一致(1080p, CRF 23, medium preset),确认生成时可直接复用预览产物。
Args:
request: 预览任务创建请求(template_id + asset_ids 等)
preview_count=1 时行为与旧版完全一致(创建 1 个任务);
preview_count=N 时一次创建 N 个独立变体任务:
- 每个变体克隆独立编辑计划(独立 clips、独立随机素材起点),N 个预览内容互不相同
- 每个变体拥有独立 task_id / 状态 / 预览视频 URL,前端按 task_id 分别轮询
- 标题样式(font/color/position 等)全局共用;标题文字/配音/封面可按变体独立
titles[] / voice_library_ids[] / cover_urls[],长度1=共用,长度N=独立)
Returns:
201 + 预览任务详情
201 + 变体任务数组 {items: [...], total: N}
"""
user_id = authenticated_user.user.id
count = max(1, request.preview_count)
logger.info(
"[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d, preview_count=%d",
user_id,
request.template_id,
len(request.asset_ids),
request.preview_count,
count,
)
# 预检查队列限流
# 预检查队列限流(按变体总数计)
try:
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending + 1 > USER_PENDING_LIMIT:
raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending + 1, limit=USER_PENDING_LIMIT)
if global_pending + 1 > GLOBAL_PENDING_LIMIT:
raise GlobalQueueFull(pending_count=global_pending + 1, limit=GLOBAL_PENDING_LIMIT)
if user_pending + count > USER_PENDING_LIMIT:
raise UserPendingLimitExceeded(
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
)
if global_pending + count > GLOBAL_PENDING_LIMIT:
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
except UserPendingLimitExceeded as e:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交",
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
) from e
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
) from e
# 确定视频比例:优先前端传入,否则从模板 mode 推断
@@ -273,14 +335,11 @@ def create_preview_generation_task(
w, h = int(parts[0]), int(parts[1])
base = 1920
if w < h:
# 竖屏
output_width = round(base * w / h)
output_height = base
else:
# 横屏
output_width = base
output_height = round(base * h / w)
# 对齐到偶数
output_width = output_width - output_width % 2
output_height = output_height - output_height % 2
except (ValueError, ZeroDivisionError):
@@ -289,42 +348,71 @@ def create_preview_generation_task(
logger.info(
"[预览生成] 分辨率: video_ratio=%s%s (%dx%d)",
video_ratio, resolution, output_width, output_height,
video_ratio,
resolution,
output_width,
output_height,
)
# 从模板读取 editing_mode / mode 作为 strategy_id(渲染 pipeline 的 mode 参数)
strategy_id = _resolve_strategy_id_from_template(request.template_id, db, user_id)
title_config = request.title_config or {}
base_title_config = request.title_config or {}
use_case = CreateGenerationTaskUseCase(generation_task_repository)
# ── 预创建第一个任务,仅用于解析源编辑计划(不落库为最终任务)──
# 先创建一个临时任务拿到 task 对象上下文,实际 N 个任务在循环中统一创建;
# 为保持与旧版一致的源 plan 解析逻辑,先创建任务0、解析源 plan,
# 再预克隆 N 个变体 plan,最后重建任务关联。
# 简化实现:直接创建全部任务,plan 关联在创建后、入队前完成。
created_tasks: list = []
variant_plan_ids: list[str] = [] # 每个变体最终关联的 plan_id(按变体顺序)
try:
task = use_case.execute(
CreateGenerationTaskCommand(
project_id="",
asset_library_id="",
strategy_id=strategy_id,
voice_library_id=request.voice_library_id,
template_id=request.template_id,
asset_ids=list(request.asset_ids),
title_ids=list(request.title_ids),
voice_ids=list(request.voice_ids),
created_by_user_id=user_id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode="",
batch_id="",
video_title=request.video_title,
resolution=resolution,
bgm_config=request.bgm_config or {},
auto_retry_enabled=False,
auto_retry_max=0,
is_preview=True,
title_config=title_config,
output_width=output_width,
output_height=output_height,
for variant_index in range(count):
# 变体独立标题文字:titles[] 覆盖 title_config.text
variant_title_text = _variant_value(request.titles, variant_index, "")
variant_title_config = dict(base_title_config)
if variant_title_text.strip():
variant_title_config["text"] = variant_title_text.strip()
# 变体独立配音
variant_voice_library_id = _variant_value(
request.voice_library_ids, variant_index, request.voice_library_id
)
)
task = use_case.execute(
CreateGenerationTaskCommand(
project_id="",
asset_library_id="",
strategy_id=strategy_id,
voice_library_id=variant_voice_library_id,
template_id=request.template_id,
asset_ids=list(request.asset_ids),
title_ids=list(request.title_ids),
voice_ids=list(request.voice_ids),
created_by_user_id=user_id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode="",
batch_id="",
video_title=request.video_title,
resolution=resolution,
bgm_config=request.bgm_config or {},
auto_retry_enabled=False,
auto_retry_max=0,
is_preview=True,
title_config=variant_title_config,
output_width=output_width,
output_height=output_height,
)
)
task.extra_meta["variant_index"] = variant_index
# 解析源编辑计划(前端传入或按模板兜底查找)
source_plan_id = _resolve_preview_edit_plan_id(request=request, task=task, db=db, user_id=user_id)
task.source_edit_plan_id = source_plan_id
generation_task_repository.update(task)
created_tasks.append(task)
except ValueError as e:
logger.warning("[预览生成] 创建失败: %s", e)
raise HTTPException(status_code=400, detail=str(e)) from e
@@ -332,93 +420,128 @@ def create_preview_generation_task(
logger.error("[预览生成] 创建失败: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e
# 关联编辑计划:如果前端未传 source_edit_plan_id,通过 template_id + user_id 查找
if not task.source_edit_plan_id and request.template_id:
try:
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
)
_plan_repo = SQLAlchemyEditPlanRepository(db)
_plans = _plan_repo.list_by_template(request.template_id, limit=20)
for _p in _plans:
if (_p.created_by_user_id or "") == user_id:
task.source_edit_plan_id = _p.id
generation_task_repository.update(task)
logger.info(
"[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s",
task.id,
_p.id,
)
break
except Exception:
logger.warning(
"[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s",
task.id,
exc_info=True,
)
# 每条预览都关联独立克隆 plan:多预览前端为 N 次并发调用,若共用同一 plan
# 则 N 条预览片段完全相同;克隆时片段起点按持久化历史区间重算(含受控复用),
# 保证各预览版本内容不同
if task.source_edit_plan_id:
# ── 克隆独立变体 plan:N 个预览全部克隆(预览不污染源 plan)──
# 源 plan 不存在(无编辑历史)时各任务走自身随机选片流程,不克隆。
source_plan_id = created_tasks[0].source_edit_plan_id if created_tasks else ""
if source_plan_id:
try:
from app.services.edit_plan_service import EditPlanService
_plan_svc = EditPlanService(db)
_preview_plan = _plan_svc.clone_plan_for_variant(
task.source_edit_plan_id,
created_by_user_id=user_id,
name_suffix="预览变体",
)
task.source_edit_plan_id = _preview_plan.id
generation_task_repository.update(task)
logger.info(
"[预览生成] 预览关联独立克隆 plan: task_id=%s clone_plan_id=%s",
task.id,
_preview_plan.id,
)
except Exception as clone_err:
# 不退回共用原 plan(否则多条预览内容相同,违反去重诉求):
# 标记任务失败并中断,前端可重新发起预览
logger.error(
"[预览生成] 克隆预览变体 plan 失败,任务标记失败: task_id=%s error=%s",
task.id,
clone_err,
exc_info=True,
)
_mark_task_failed(generation_task_repository, task, "预览变体计划创建失败")
for variant_index in range(count):
last_err: Exception | None = None
variant_plan = None
for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动
try:
variant_plan = _plan_svc.clone_plan_for_variant(
source_plan_id,
created_by_user_id=user_id,
name_suffix=f"预览变体{variant_index + 1}" if count > 1 else "预览变体",
)
break
except Exception as clone_err: # noqa: PERF203
last_err = clone_err
logger.warning(
"[预览生成] 克隆变体 plan 失败(尝试%d/2): variant=%d error=%s",
_attempt + 1,
variant_index,
clone_err,
exc_info=True,
)
if variant_plan is None:
logger.error(
"[预览生成] 克隆预览变体 plan 重试仍失败: variant=%d source=%s",
variant_index,
source_plan_id,
exc_info=last_err,
)
# 标记已创建任务失败
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体计划创建失败")
raise HTTPException(
status_code=500,
detail="创建预览任务失败:无法生成独立剪辑计划,请重试",
) from last_err
variant_plan_ids.append(variant_plan.id)
except HTTPException:
raise
except Exception as e:
logger.error("[预览生成] 克隆变体 plan 异常: %s", e, exc_info=True)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体计划创建失败")
raise HTTPException(
status_code=500,
detail="创建预览任务失败:无法生成独立剪辑计划,请重试",
) from clone_err
) from e
# 入队执行;若入队失败则标记任务为 failed 避免僵尸数据
try:
if not safe_enqueue_generation_task(
task,
generation_task_repository,
user_id=user_id,
log_prefix="[预览生成]",
log_task_status=True,
):
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队失败")
raise HTTPException(status_code=500, detail="任务入队失败,请稍后重试")
except UserPendingLimitExceeded as e:
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交",
) from None
except GlobalQueueFull:
_mark_task_failed(generation_task_repository, task, "系统队列已满")
# 关联变体 plan 并回写标题配置
for variant_index, task in enumerate(created_tasks):
if variant_plan_ids:
task.source_edit_plan_id = variant_plan_ids[variant_index]
generation_task_repository.update(task)
# 回写变体标题到 plan configworker 渲染时从 plan 读取 title 配置)
if task.source_edit_plan_id and (task.title_config or {}).get("text", "").strip():
try:
from app.api.routes.generation_tasks import _writeback_edit_plan_config
_writeback_edit_plan_config(
plan_id=task.source_edit_plan_id,
task_id=task.id,
title_config=task.title_config,
db=db,
)
except Exception:
logger.warning(
"[预览生成] 回写标题配置失败(不影响主流程): task_id=%s",
task.id,
exc_info=True,
)
# ── 入队 ──
responses: list[PreviewGenerationTaskResponse] = []
rate_limit_exc: Exception | None = None # 记录首个限流异常,全部失败时返回结构化提示
for variant_index, task in enumerate(created_tasks):
try:
enqueued = safe_enqueue_generation_task(
task,
generation_task_repository,
user_id=user_id,
log_prefix=f"[预览生成][变体{variant_index + 1}]",
log_task_status=True,
)
if not enqueued:
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队失败")
except UserPendingLimitExceeded as e:
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
rate_limit_exc = rate_limit_exc or e
except GlobalQueueFull as e:
_mark_task_failed(generation_task_repository, task, "系统队列已满")
rate_limit_exc = rate_limit_exc or e
except Exception:
logger.exception("[预览生成] 入队异常: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队异常")
# enqueue 会原地更新 task 状态/进度,直接用 task 构造响应
responses.append(_to_preview_response(task))
# 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败"
if all(r.status == "failed" for r in responses) and rate_limit_exc is not None:
if isinstance(rate_limit_exc, UserPendingLimitExceeded):
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="user"),
)
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="global"),
)
return _to_preview_response(task)
logger.info(
"[预览生成] 创建完成: %d 个变体任务, task_ids=%s",
len(responses),
[r.task_id for r in responses],
)
return BatchPreviewGenerationTaskResponse(items=responses, total=len(responses))
@router.get("/preview/{task_id}", response_model=PreviewGenerationTaskResponse)
+53 -20
View File
@@ -10,6 +10,7 @@ from app.core.task_enqueue import (
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task,
)
from app.dependencies import (
@@ -47,6 +48,15 @@ logger = logging.getLogger(__name__)
router = APIRouter()
def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
"""从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。"""
if not values:
return fallback
if len(values) == 1:
return values[0]
return values[index] if index < len(values) else fallback
def _to_generation_task_response(task) -> GenerationTaskResponse:
return GenerationTaskResponse(
id=task.id,
@@ -408,12 +418,12 @@ def create_generation_task(
except UserPendingLimitExceeded as e:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交",
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
) from e
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
) from e
# 画中画已下线:strategy_id 中的 pip/voice_pip 统一映射为 one_take
@@ -472,12 +482,21 @@ def create_generation_task(
if task_index > 0 and variant_plan_ids:
effective_plan_id = variant_plan_ids[task_index - 1]
# 变体级独立配置:titles[]/voice_library_ids[]/cover_urls[]
# 长度1=所有变体共用,长度=count=每个变体独立,空数组=回退单值字段
variant_title_text = _variant_value(request.titles, task_index, "")
variant_title_config = dict(request.title_config or {})
if variant_title_text.strip():
variant_title_config["text"] = variant_title_text.strip()
variant_voice_library_id = _variant_value(request.voice_library_ids, task_index, request.voice_library_id)
variant_cover_url = _variant_value(request.cover_urls, task_index, request.cover_url)
task = use_case.execute(
CreateGenerationTaskCommand(
project_id=project_id,
asset_library_id=asset_library_id,
strategy_id=effective_strategy_id,
voice_library_id=request.voice_library_id,
voice_library_id=variant_voice_library_id,
template_id=request.template_id,
asset_ids=resolved_asset_ids,
title_ids=request.title_ids,
@@ -495,10 +514,12 @@ def create_generation_task(
source_task_id=request.source_task_id,
output_width=request.output_width,
output_height=request.output_height,
cover_url=request.cover_url,
title_config=request.title_config or {},
cover_url=variant_cover_url,
title_config=variant_title_config,
)
)
# 变体序号写入 extra_meta(响应/排查时可辨识)
task.extra_meta["variant_index"] = task_index
try:
# 兜底关联编辑计划:前端未传 source_edit_plan_id 时,
# 通过 template_id + user_id 在 DB 层直接查找最新的 plan。
@@ -533,13 +554,13 @@ def create_generation_task(
# 回写 plan.config:必须在 enqueue 之前执行,
# 确保 worker 读取 plan 时 config 中已包含 generation_task_id。
# 只在首个任务时回写一次,避免批量生成时循环覆盖
# 批量场景下每个变体关联独立 plan,需各自回写自己的变体标题配置
_effective_plan_id = task.source_edit_plan_id
if _effective_plan_id and len(created_tasks) == 0:
if _effective_plan_id:
_writeback_edit_plan_config(
plan_id=_effective_plan_id,
task_id=task.id,
title_config=request.title_config,
title_config=variant_title_config,
db=db,
)
@@ -559,7 +580,7 @@ def create_generation_task(
if not created_tasks:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
) from _e
break
except GlobalQueueFull as _e:
@@ -567,7 +588,7 @@ def create_generation_task(
if not created_tasks:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
) from _e
break
except HTTPException:
@@ -693,15 +714,15 @@ def confirm_generation(
log_task_status=True,
):
logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id)
except UserPendingLimitExceeded:
except UserPendingLimitExceeded as _e:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
) from None
except GlobalQueueFull:
except GlobalQueueFull as _e:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
) from None
return BatchGenerationTaskResponse(
@@ -783,12 +804,24 @@ def retry_generation_task(
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
detail=build_rate_limit_detail(
UserPendingLimitExceeded(
user_id=user_id,
pending_count=user_pending,
limit=USER_PENDING_LIMIT,
),
generation_task_repository,
scope="user",
),
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
detail=build_rate_limit_detail(
GlobalQueueFull(pending_count=global_pending, limit=GLOBAL_PENDING_LIMIT),
generation_task_repository,
scope="global",
),
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
@@ -823,15 +856,15 @@ def retry_generation_task(
log_task_status=True,
):
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded:
except UserPendingLimitExceeded as _e:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
) from None
except GlobalQueueFull:
except GlobalQueueFull as _e:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
) from None
return _to_generation_task_response(retried)
+7 -1
View File
@@ -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,
+7 -1
View File
@@ -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",
+162 -55
View File
@@ -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,16 +198,29 @@ 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)
def _persist_celery_task_id(repo: Any, job: Any, celery_task_id: str) -> None:
"""记录 celery 消息 ID 到任务行,供孤儿清理时 revoke/清除队列消息(#1714)。"""
if not celery_task_id:
return
try:
job.celery_task_id = celery_task_id
repo.update(job)
except Exception: # noqa: BLE001 — 记录失败不影响主流程(执行前状态守卫兜底)
pass
def _submit_ingest_job(
project_id: str,
library_id: str,
storage_key: str,
ingest_job_repository: Any,
file_hash: str = "",
asset_id: str = "",
) -> Any:
use_case = SubmitIngestJobUseCase(ingest_job_repository)
job = use_case.execute(
@@ -129,9 +229,11 @@ def _submit_ingest_job(
library_id=library_id,
storage_key=storage_key,
file_hash=file_hash,
asset_id=asset_id,
)
)
celery_app.send_task("worker.ingest_asset", args=[job.id])
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
return job
@@ -202,7 +304,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 +314,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 +348,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 +359,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 +368,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 +388,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 +400,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 +453,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 +462,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(
+7 -3
View File
@@ -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,
+8
View File
@@ -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
+131 -3
View File
@@ -8,27 +8,145 @@ logger = logging.getLogger(__name__)
# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ──
USER_PENDING_LIMIT = 3 # 单用户 pending 上限
GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限
WORKER_CONCURRENCY = 4 # worker 渲染并发数(infra/docker/compose.yml WORKER_CONCURRENCY 默认值)
# 限流错误码:前端据此区分"排队等待"与"创建失败"
ERROR_CODE_USER_QUEUE_FULL = "USER_QUEUE_FULL" # 429:用户自己的任务排队中
ERROR_CODE_SYSTEM_QUEUE_FULL = "SYSTEM_QUEUE_FULL" # 503:系统整体繁忙
class UserPendingLimitExceeded(Exception):
"""用户 pending 任务数超限,返回 429。"""
def __init__(self, user_id: str, pending_count: int, limit: int):
def __init__(
self,
user_id: str,
pending_count: int,
limit: int,
*,
running_count: int = 0,
requested_count: int = 1,
queue_ahead: int = 0,
estimated_wait_seconds: int = 0,
):
self.user_id = user_id
self.pending_count = pending_count
self.limit = limit
# 排队上下文(用于 429 结构化提示,前端展示"排队中"而非"创建失败"
self.running_count = running_count
self.requested_count = requested_count
self.queue_ahead = queue_ahead
self.estimated_wait_seconds = estimated_wait_seconds
super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}")
class GlobalQueueFull(Exception):
"""全局限流,返回 503。"""
def __init__(self, pending_count: int, limit: int):
def __init__(
self,
pending_count: int,
limit: int,
*,
running_count: int = 0,
queue_ahead: int = 0,
estimated_wait_seconds: int = 0,
):
self.pending_count = pending_count
self.limit = limit
self.running_count = running_count
self.queue_ahead = queue_ahead
self.estimated_wait_seconds = estimated_wait_seconds
super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}")
def _estimate_wait_seconds(queue_ahead: int, generation_task_repository: Any) -> int:
"""根据排队任务数 + worker 并发数 + 历史平均任务耗时估算等待秒数。
估算公式:ceil(排队任务数 / 并发数) × 平均单任务耗时。
拿不到历史数据时仓储层返回默认 120 秒。
"""
import math
if queue_ahead <= 0:
return 0
try:
estimator = getattr(generation_task_repository, "estimate_avg_duration_seconds", None)
avg_seconds = estimator() if estimator is not None else 120.0
except Exception:
avg_seconds = 120.0
return int(math.ceil(queue_ahead / WORKER_CONCURRENCY) * avg_seconds)
def build_rate_limit_detail(
exc: Exception,
generation_task_repository: Any,
*,
scope: str = "user",
) -> dict:
"""构造结构化限流响应体(HTTPException 的 detail)。
前端按 detail.code 判断场景:
- USER_QUEUE_FULL (429):用户自己的任务在排队,应提示"等待/继续排队",不是创建失败
- SYSTEM_QUEUE_FULL (503):系统繁忙,稍后重试
detail 字段:
- code: 错误码
- message: 可读中文提示(可直接展示)
- queued_count: 当前排队(pending)任务数
- running_count: 当前渲染中(running)任务数
- queue_ahead: 前方排队任务数(预计等待批次依据)
- estimated_wait_seconds: 预计等待秒数
- limit: 对应限流上限
"""
if scope == "user" and isinstance(exc, UserPendingLimitExceeded):
running = exc.running_count
if not running:
try:
counter = getattr(generation_task_repository, "count_running_by_user", None)
running = counter(exc.user_id) if counter is not None else 0
except Exception:
running = 0
queue_ahead = exc.queue_ahead or max(exc.pending_count, 0)
wait = exc.estimated_wait_seconds or _estimate_wait_seconds(queue_ahead, generation_task_repository)
wait_minutes = max(1, round(wait / 60))
message = (
f"您有 {exc.pending_count} 个任务正在排队、{running} 个正在渲染,"
f"同一时间最多提交 {exc.limit} 个任务。请等待约 {wait_minutes} 分钟后再提交"
)
return {
"code": ERROR_CODE_USER_QUEUE_FULL,
"message": message,
"queued_count": exc.pending_count,
"running_count": running,
"queue_ahead": queue_ahead,
"estimated_wait_seconds": wait,
"limit": exc.limit,
}
# 全局繁忙
pending = getattr(exc, "pending_count", 0)
running = getattr(exc, "running_count", 0)
if not running:
try:
counter = getattr(generation_task_repository, "count_running_total", None)
running = counter() if counter is not None else 0
except Exception:
running = 0
queue_ahead = getattr(exc, "queue_ahead", 0) or pending
wait = getattr(exc, "estimated_wait_seconds", 0) or _estimate_wait_seconds(queue_ahead, generation_task_repository)
wait_minutes = max(1, round(wait / 60))
return {
"code": ERROR_CODE_SYSTEM_QUEUE_FULL,
"message": f"系统繁忙:当前 {pending} 个任务排队中、{running} 个渲染中,预计等待约 {wait_minutes} 分钟,请稍后再试",
"queued_count": pending,
"running_count": running,
"queue_ahead": queue_ahead,
"estimated_wait_seconds": wait,
"limit": getattr(exc, "limit", GLOBAL_PENDING_LIMIT),
}
def check_queue_limits(
user_id: str,
generation_task_repository: Any,
@@ -161,7 +279,17 @@ def safe_enqueue_generation_task(
# ── 发送 Celery 任务 ──
try:
celery_app.send_task("worker.generate_video", args=[task.id])
celery_result = celery_app.send_task("worker.generate_video", args=[task.id])
# 记录 celery 消息 ID:孤儿清理/超时作废时据此 revoke + 清除队列消息(#1714
celery_task_id = getattr(celery_result, "id", "")
if celery_task_id:
try:
task.celery_task_id = celery_task_id
generation_task_repository.update(task)
except Exception as persist_err: # noqa: BLE001
logger.warning(
"%s 持久化 celery_task_id 失败(不影响主流程): task_id=%s err=%s", log_prefix, task.id, persist_err
)
except Exception as e:
logger.error(
"%s 入队失败,标记为失败: task_id=%s error=%s",
+69 -2
View File
@@ -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
+7 -5
View File
@@ -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_idduplicated=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_idduplicated=true 时返回")
duplicated: bool = Field(default=False, description="是否为重复素材/重复提交(命中幂等去重)")
asset_id: str = Field(default="", description="素材 asset_id重复提交时返回已存在记录")
+18 -24
View File
@@ -185,6 +185,13 @@ test.describe("Core generation flow", () => {
await expect(page.locator(".xx-choice-item.selected")).toBeVisible()
await page.getByRole("button", { name: "下一步" }).click()
// Step1 下一步弹出数量选择弹窗(Issue #1677 固定6步:模板→素材→配音→标题→确认生成→封面)
// 单视频流程:默认 1 个,点击「生成 1 个视频」进入步骤2
await expect(page.getByRole("heading", { name: "要生成几个视频?" })).toBeVisible({
timeout: 10_000,
})
await page.getByRole("button", { name: "生成 1 个视频" }).click()
// Step 2: select material (card grid UI)
await expect(page.getByRole("heading", { name: /选择素材/ })).toBeVisible()
const librarySelect = page.locator("select").first()
@@ -255,32 +262,19 @@ test.describe("Core generation flow", () => {
expect(genData.items.length).toBeGreaterThan(0)
expect(genData.items[0].id).toBeTruthy()
// Step 5: 确认生成页 — 任务创建成功后自动跳转,展示渲染进度
await expect(page.getByRole("heading", { name: /确认生成/ })).toBeVisible({
timeout: 15_000,
// 单视频(N=1):点击「确认生成视频」后跳 Step 5确认生成,展示实时渲染进度
await expect(page.getByRole("heading", { name: "🎬 确认生成" })).toBeVisible({
timeout: 30_000,
})
// Step 5 → Step 6:等待渲染终态
// - 完成:页面出现「视频生成完成」,步骤5「下一步」按钮解锁,点击进入封面
// - 失败:出现「生成失败」,停在确认生成页也算向导流程走通
// - 超时未终态(测试环境 worker 可能不处理任务):进度仍在轮询,同样算走通
const renderSucceeded = await page
.getByText("视频生成完成", { exact: false })
.waitFor({ timeout: 180_000 })
.then(() => true)
.catch(() => false)
if (renderSucceeded) {
// 渲染完成:手动点「下一步」进入封面步骤(渲染完不自动跳转)
await page.getByRole("button", { name: "下一步" }).click()
// Step 6: 封面(最后一步,无主按钮),仅验证页面渲染
await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({
timeout: 15_000,
})
} else {
// 失败或超时:仍在确认生成页(进度展示或失败提示),向导流程已完整走通
await expect(page.getByRole("heading", { name: /确认生成/ })).toBeVisible()
console.log("[E2E] 渲染任务失败或未在 180s 内完成,冒烟测试仍通过(已达确认生成页)")
}
// 等待渲染完成:进度卡变为「视频生成完成」(最长等待 3 分钟)
await expect(page.getByText("视频生成完成")).toBeVisible({ timeout: 180_000 })
// 全部完成后「下一步:选择封面」解锁,点击进入 Step 6
await page.getByRole("button", { name: /下一步:选择封面/ }).click()
await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({
timeout: 30_000,
})
} else {
console.log(`[E2E] Generate API returned ${genResp.status()}, wizard flow test still passes`)
// 创建失败时停留在标题页并展示错误提示
+4
View File
@@ -0,0 +1,4 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">
<rect width="64" height="64" rx="14" fill="#3b82f6"/>
<text x="32" y="44" font-size="34" text-anchor="middle">🦐</text>
</svg>

After

Width:  |  Height:  |  Size: 194 B

+11
View File
@@ -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
}
/** 直传完成确认返回 */
+54 -4
View File
@@ -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<DirectUploadPrepareResult> => {
const response = await apiClient.post("/upload/direct/prepare", data)
// prepare 单独放宽到 30s(全局 axios 实例只有 10sstaging 抖动时易超时)
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<DirectUploadCompleteResult> => {
const response = await apiClient.post("/upload/direct/complete", data)
// complete 内含 OSS 存在性检查 + 建库 + 派单,放宽到 60s;
// 超时不代表失败(记录可能已建成),调用方禁止超时后盲目重传整个文件
const response = await apiClient.post("/upload/direct/complete", data, { timeout: 60_000 })
return response.data
}
@@ -109,8 +121,20 @@ export interface DirectUploadHandle {
export const prepareDirectUploadHandle = async (data: {
file: File
library_id: string
/** 前端算好的文件内容哈希(SHA-256 hex),prepare/complete 均携带 */
fileHash?: string
/** 本次逻辑上传的幂等 tokenprepare/complete 一致、重试复用 */
clientUploadId?: string
}): Promise<DirectUploadHandle> => {
const project = await getOrCreateDefaultProject()
// 默认项目初始化失败(项目列表接口异常/自动创建失败)给出独立、明确的提示,
// 不与 prepare 的签名接口错误混在一起
let project: Awaited<ReturnType<typeof getOrCreateDefaultProject>>
try {
project = await getOrCreateDefaultProject()
} catch (err) {
const reason = err instanceof Error ? err.message : "网络异常"
throw new Error(`初始化默认项目失败,无法开始上传:${reason}`)
}
const prepared = await prepareDirectUpload({
project_id: project.id,
@@ -118,6 +142,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 +154,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 +165,30 @@ export const uploadAssetDirect = async (data: {
file: File
library_id: string
onProgress?: (percent: number) => void
/** 文件内容哈希;未传时自动补算(配音/封面/克隆等非队列链路统一受益) */
fileHash?: string
/** 幂等 token;未传时自动生成 */
clientUploadId?: string
}): Promise<DirectUploadCompleteResult> => {
const handle = await prepareDirectUploadHandle({ file: data.file, library_id: data.library_id })
// 自动补算哈希与幂等 token:确保 file_hash 去重闸门对所有上传链路生效
const fileHash = data.fileHash ?? (await computeFileHash(data.file))
const clientUploadId = data.clientUploadId ?? makeClientUploadId()
const handle = await prepareDirectUploadHandle({
file: data.file,
library_id: data.library_id,
fileHash,
clientUploadId,
})
// prepare 阶段后端 file_hash 命中素材库已有相同文件:跳过 transfer + complete
if (handle.prepared.skip_transfer || handle.prepared.duplicated) {
return {
storage_key: handle.prepared.storage_key,
ingest_job_id: "",
url: "",
duplicated: true,
asset_id: handle.prepared.asset_id,
}
}
await handle.transfer(data.onProgress)
return handle.complete()
}
+143
View File
@@ -0,0 +1,143 @@
/**
* 上传去重 / 幂等工具(Issue #1714
*
* 背景:同一文件被反复入队、complete 超时后盲目重传,导致后端创建大量重复
* PROCESSING 素材记录。本模块提供两类纯函数:
*
* 1. 文件指纹:
* - makeFileFingerprint():文件名+大小+lastModified,入队去重用(同步、零开销)
* - computeFileHash()SHA-256 内容哈希(小文件全量、大文件抽样头尾),
* prepare/complete 时发给后端打开 file_hash 去重闸门
* 2. 队列去重:findDuplicateInQueue() 判断文件是否已在队列中
* 3. 幂等 tokenmakeClientUploadId() 生成上传幂等 ID(每次"一次逻辑上传"一个,
* 重试复用同一 ID,重新入队才生成新 ID)
*/
/** 全量哈希阈值:≤64MB 全量读入计算;超过即走头尾抽样,避免 100~256MB 视频被整文件读进内存卡死页面 */
export const HASH_FULL_READ_LIMIT = 64 * 1024 * 1024 // 64MB
/** 抽样读取的头尾片段大小(各 16MB) */
export const HASH_SAMPLE_CHUNK = 16 * 1024 * 1024
/** 计算指纹时,文件在队列中已存在的状态(已失败的可以重试,不算重复) */
export type DedupExcludeStatus = "error" | "done"
/**
* 文件入队指纹:同库 + 文件名 + 大小 + 修改时间。
* 同一文件(File 对象由 <input> 重选或拖拽重复触发时三个字段均一致)稳定复现;
* 不同文件极小概率碰撞时可由后端 file_hash 内容去重兜底。
*/
export function makeFileFingerprint(file: Pick<File, "name" | "size" | "lastModified">): string {
return `${file.name}::${file.size}::${file.lastModified}`
}
/**
* 在现有队列项中查找同一文件的在途记录。
* 已失败(error)的项允许重试路径复用、已完成(done)的可跳过;
* 处于 preparing/uploading/ingesting 的在途项一律视为重复,禁止重复入队。
*
* 返回命中的队列项 id(tempId),未命中返回 null。
*/
export function findDuplicateInQueue<T extends { fileKey: string; status: string }>(
queue: T[],
fileKey: string,
excludeStatuses: DedupExcludeStatus[] = [],
): T | null {
const exclude = new Set<string>(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<ArrayBuffer> {
if (typeof blob.arrayBuffer === "function") {
return blob.arrayBuffer()
}
return new Promise<ArrayBuffer>((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<ArrayBuffer> {
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-256hex64 字符,与后端 file_hash 字段长度一致)。
* - ≤64MB:全量哈希,内容一致必然一致
* - >64MB:哈希「头部 16MB + 尾部 16MB + 文件大小」,视频素材体积大、
* 头部含 moov 元数据、尾部含 mdat 结尾,抽样碰撞概率可忽略,
* 且避免 100~256MB 视频被整文件读进内存导致页面卡死/崩溃
*
* 运行环境不支持 crypto.subtle(非安全上下文/老浏览器)时返回空字符串,
* 调用方据此降级为不传 hash(后端仍有幂等 token + 同文件名兜底去重)。
*/
export async function computeFileHash(file: File): Promise<string> {
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 ""
}
}
+14 -3
View File
@@ -12,13 +12,18 @@ export type {
UserResponse,
WechatAuthUrlResponse,
WechatCallbackResponse,
WechatBindUrlResponse,
WechatBindCompleteResponse,
WechatUnbindResponse,
UpdateProfileRequest,
UpdateProfileResponse,
SendVerificationCodeRequest,
BindContactRequest,
BindContactResponse,
} from "./types"
// 用户工具函数
export { normalizeUser } from "./user"
export { normalizeUser, updateProfile } from "./user"
// 登录/注册/登出/刷新
export { login, refreshAccessToken, register, logout } from "./login"
@@ -32,8 +37,14 @@ export { requestPasswordReset, resetPassword } from "./password"
// 邮箱验证
export { verifyEmail } from "./email"
// 微信登录
export { getWechatAuthUrl, wechatCallback } from "./wechat"
// 微信登录 / 绑定
export {
getWechatAuthUrl,
wechatCallback,
getWechatBindUrl,
bindWechat,
unbindWechat,
} from "./wechat"
// 联系方式
export { sendVerificationCode, bindContact } from "./contact"
+44
View File
@@ -34,6 +34,17 @@ export interface User {
is_email_verified: boolean
email_verified: boolean
created_at?: string
/** 微信是否已绑定 */
wechat_bound?: boolean
/** 微信昵称(绑定后展示) */
wechat_nickname?: string
/** 头像 URL(微信头像等) */
avatar_url?: string
/** 手机号 */
phone?: string
phone_verified?: boolean
/** 资料是否完善(微信新用户首次登录为 false,需填昵称引导) */
profile_completed?: boolean
}
export interface UserResponse {
@@ -45,6 +56,12 @@ export interface UserResponse {
is_email_verified?: boolean
email_verified?: boolean
created_at?: string
wechat_bound?: boolean
wechat_nickname?: string
avatar_url?: string
phone?: string
phone_verified?: boolean
profile_completed?: boolean
}
export interface WechatAuthUrlResponse {
@@ -80,3 +97,30 @@ export interface BindContactResponse {
success: boolean
user: User
}
/** 更新个人资料请求 */
export interface UpdateProfileRequest {
display_name?: string
}
/** 更新个人资料响应(返回最新用户信息) */
export interface UpdateProfileResponse {
user: UserResponse
}
/** 微信绑定授权链接响应 */
export interface WechatBindUrlResponse {
auth_url: string
state: string
}
/** 微信绑定完成响应 */
export interface WechatBindCompleteResponse {
success: boolean
user: UserResponse
}
/** 微信解绑响应 */
export interface WechatUnbindResponse {
success: boolean
}
+16 -1
View File
@@ -1,4 +1,5 @@
import type { User, UserResponse } from "./types"
import apiClient from "../client"
import type { User, UserResponse, UpdateProfileRequest, UpdateProfileResponse } from "./types"
/**
* 规范化用户数据,兼容不同后端返回格式
@@ -16,5 +17,19 @@ export const normalizeUser = (data: UserResponse): User => {
is_email_verified: emailVerified,
email_verified: emailVerified,
created_at: data.created_at,
wechat_bound: data.wechat_bound,
wechat_nickname: data.wechat_nickname,
avatar_url: data.avatar_url,
phone: data.phone,
phone_verified: data.phone_verified,
profile_completed: data.profile_completed,
}
}
/**
* 更新个人资料(昵称等)
*/
export const updateProfile = async (data: UpdateProfileRequest): Promise<User> => {
const response = await apiClient.patch<UpdateProfileResponse>("/auth/me", data)
return normalizeUser(response.data.user)
}
+35 -2
View File
@@ -1,8 +1,14 @@
import apiClient from "../client"
import type { WechatAuthUrlResponse, WechatCallbackResponse } from "./types"
import type {
WechatAuthUrlResponse,
WechatCallbackResponse,
WechatBindUrlResponse,
WechatBindCompleteResponse,
WechatUnbindResponse,
} from "./types"
/**
* 获取微信授权链接
* 获取微信授权链接(登录场景)
*/
export const getWechatAuthUrl = async (): Promise<WechatAuthUrlResponse> => {
const response = await apiClient.get("/auth/wechat/url")
@@ -19,3 +25,30 @@ export const wechatCallback = async (
const response = await apiClient.post("/auth/wechat/callback", { code, state })
return response.data
}
/**
* 获取微信绑定授权链接(已登录用户绑定场景)
*/
export const getWechatBindUrl = async (): Promise<WechatBindUrlResponse> => {
const response = await apiClient.get("/auth/wechat/bind/url")
return response.data
}
/**
* 微信绑定完成(扫码回调后用 code 绑定到当前登录账号)
*/
export const bindWechat = async (
code: string,
state: string,
): Promise<WechatBindCompleteResponse> => {
const response = await apiClient.post("/auth/wechat/bind", { code, state })
return response.data
}
/**
* 解绑微信
*/
export const unbindWechat = async (): Promise<WechatUnbindResponse> => {
const response = await apiClient.delete("/auth/wechat/bind")
return response.data
}
+112
View File
@@ -0,0 +1,112 @@
/**
* 微信扫码登录 WxLogin JS-SDK 动态加载与授权参数解析
*
* 微信官网嵌入式二维码方案:页面引入 https://res.wx.qq.com/connect/zh_CN/htmledition/js/wxLogin.js
* 后挂载全局 window.WxLoginnew 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<WxLoginConstructor> | null = null
/**
* 动态加载微信 WxLogin JS(单例:并发调用复用同一个 promise)。
* 加载失败或超时会 reject,调用方应回退到整页跳转授权方式。
*/
export function loadWxLoginScript(): Promise<WxLoginConstructor> {
if (window.WxLogin) return Promise.resolve(window.WxLogin)
if (loadPromise) return loadPromise
loadPromise = new Promise<WxLoginConstructor>((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
}
}
+132
View File
@@ -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<string, unknown>
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<ErrorBody>
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)
}
+25 -4
View File
@@ -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
}
+8
View File
@@ -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 */
@@ -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;
}
@@ -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<WechatQrScene, string> = {
login: "wechat-qr-login-container",
bind: "wechat-qr-bind-container",
}
const STATE_STORAGE_KEY: Record<WechatQrScene, string> = {
login: "wechat_state",
bind: "wechat_bind_state",
}
/**
* 等待二维码容器挂载到 DOM。antd Modal 内容通过 portal 渲染且带进场动画,
* 父组件 effect 首次执行时容器可能尚未出现在 document 中。
*/
function waitForContainer(id: string, timeoutMs = 3000): Promise<HTMLElement | null> {
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<WechatQrModalProps> = ({
open,
scene,
onClose,
onLoginSuccess,
onBindSuccess,
}) => {
const setAuth = useAuthStore((state) => state.setAuth)
const setUser = useAuthStore((state) => state.setUser)
const [status, setStatus] = useState<QrStatus>("loading")
const [errorMsg, setErrorMsg] = useState("")
/** 刷新二维码计数:变化时重新请求授权链接并重渲染 */
const [renderSeq, setRenderSeq] = useState(0)
/** 最新授权链接,用于"整页打开"兜底 */
const authUrlRef = useRef<string | null>(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 (
<Modal
title={isLogin ? "微信扫码登录" : "绑定微信"}
open={open}
onCancel={onClose}
footer={null}
width={380}
maskClosable={false}
destroyOnHidden
>
<div className="xx-wechat-qr-modal">
{/* 二维码容器常驻:WxLogin 在 loading 阶段就会把 iframe 渲染进来,
不能按 status 条件渲染,否则 effect 里永远找不到容器 */}
<div
id={CONTAINER_ID[scene]}
className="xx-wechat-qr-container"
style={{ visibility: status === "qrcode" ? "visible" : "hidden" }}
/>
{status === "loading" && (
<div className="xx-wechat-qr-overlay">
<Spin size="large" />
<p>...</p>
</div>
)}
{status === "qrcode" && (
<p className="xx-wechat-qr-tip">使{isLogin ? "登录" : "绑定账号"}</p>
)}
{status === "error" && (
<div className="xx-wechat-qr-overlay xx-wechat-qr-error">
<p className="xx-wechat-qr-error-msg">{errorMsg}</p>
<div className="xx-wechat-qr-error-actions">
<Button buttonType="primary" buttonSize="md" onClick={handleRefresh}>
</Button>
{authUrlRef.current && (
<button
type="button"
className="xx-wechat-qr-fallback"
onClick={handleFullPageRedirect}
>
使
</button>
)}
</div>
</div>
)}
</div>
</Modal>
)
}
export default WechatQrModal
@@ -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<WechatQrMessageData>
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)
}
@@ -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 = () => {
{/* 上传区域 */}
<AssetUploadZone
uploading={uploading}
transferActive={transferActive}
activeCount={activeCount}
pendingCount={pendingCount}
onUpload={enqueueUploads}
@@ -214,6 +217,7 @@ const AssetLibrary: React.FC = () => {
selectedIds={selectedIds}
diagnosingId={diagnosingId}
uploadProgressMap={uploadProgressMap}
stalledAssetIds={stalledAssetIds}
onRetry={refetchAssets}
onToggleSelect={toggleSelect}
onDiagnose={handleDiagnose}
+36
View File
@@ -831,6 +831,20 @@
color: #ef4444;
}
.xx-upload-queue-error-detail {
margin-top: 4px;
font-size: 12px;
line-height: 1.5;
color: #ef4444;
word-break: break-word;
white-space: normal;
}
.xx-upload-queue-error-hint {
margin-top: 2px;
color: #b45309;
}
.xx-upload-queue-actions {
display: flex;
gap: 6px;
@@ -1041,3 +1055,25 @@
background: #fef2f2;
color: #dc2626;
}
/* 上传入口禁用态(直传进行中,防重复提交,Issue #1714 */
.xx-asset-upload-btn:disabled {
opacity: 0.6;
cursor: not-allowed;
}
.xx-asset-upload-btn:disabled:hover {
opacity: 0.6;
}
.xx-asset-upload-btn:disabled:active {
transform: none;
}
/* 处理超时遮罩:创建超过 10 分钟仍在处理中(疑似后端卡住),停止转圈并警示 */
.xx-asset-thumb-stalled {
background: rgba(217, 119, 6, 0.28);
color: #fde68a;
backdrop-filter: blur(2px);
}
.xx-asset-thumb-stalled :first-child {
font-size: var(--font-size-2xl);
}
@@ -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<AssetCardProps> = ({
selected,
diagnosing,
uploadProgress,
stalled,
onToggle,
onDiagnose,
onPlay,
@@ -67,11 +70,15 @@ const AssetCard: React.FC<AssetCardProps> = ({
</div>
)}
{/* 转码/处理中遮罩 */}
{/* 转码/处理中遮罩(卡死超过 10 分钟时停止转圈,提示超时) */}
{asset.loading && !isUploading && (
<div className="xx-asset-thumb-overlay xx-asset-thumb-processing">
<LoadingOutlined />
<span></span>
<div
className={`xx-asset-thumb-overlay ${
stalled ? "xx-asset-thumb-stalled" : "xx-asset-thumb-processing"
}`}
>
{stalled ? <CloseCircleOutlined /> : <LoadingOutlined />}
<span>{stalled ? "处理超时,可重试上传" : "转码处理中"}</span>
</div>
)}
@@ -129,7 +136,7 @@ const AssetCard: React.FC<AssetCardProps> = ({
</p>
<div className="xx-asset-meta">
<span className="xx-asset-meta-status">
<StatusPill status={asset.status} label={asset.statusLabel} />
<StatusPill status={asset.status} label={stalled ? "处理超时" : asset.statusLabel} />
</span>
{asset.duration && <span className="xx-asset-meta-duration">{asset.duration}</span>}
</div>
@@ -19,6 +19,8 @@ export interface AssetGridSectionProps {
selectedIds: Set<string>
diagnosingId: string | null
uploadProgressMap?: UploadProgressMap
/** 创建超过 10 分钟仍在处理中的素材 id(疑似后端卡住),卡片提示处理超时 */
stalledAssetIds?: Set<string>
onRetry?: () => void
onToggleSelect: (id: string) => void
onDiagnose: (asset: AssetItem) => void
@@ -34,6 +36,7 @@ export const AssetGridSection: React.FC<AssetGridSectionProps> = ({
selectedIds,
diagnosingId,
uploadProgressMap,
stalledAssetIds,
onRetry,
onToggleSelect,
onDiagnose,
@@ -76,6 +79,7 @@ export const AssetGridSection: React.FC<AssetGridSectionProps> = ({
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)}
@@ -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<AssetUploadZoneProps> = ({
uploading,
transferActive,
activeCount,
pendingCount,
onUpload,
@@ -27,6 +31,11 @@ export const AssetUploadZone: React.FC<AssetUploadZoneProps> = ({
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<AssetUploadZoneProps> = ({
<button
type="button"
className="xx-asset-upload-btn"
onClick={() => inputRef.current?.click()}
disabled={transferActive}
title={transferActive ? "文件上传中,暂不能添加新文件" : undefined}
onClick={() => {
if (transferActive) {
message.warning("文件正在上传中,请等待当前上传完成后再添加")
return
}
inputRef.current?.click()
}}
>
<PlusOutlined />
{transferActive ? "上传中…" : "上传素材"}
</button>
<span className="xx-asset-upload-status">
{uploading ? (
@@ -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<UploadItem["status"], string> = {
error: "上传失败",
}
/** 失败阶段中文名:让用户一眼看到失败发生在哪一步 */
const FAIL_STAGE_TEXT: Record<UploadFailStage, string> = {
prepare: "准备上传阶段",
transfer: "文件传输阶段",
complete: "确认入库阶段",
}
const UploadQueuePanel: React.FC<UploadQueuePanelProps> = ({
items,
onRetry,
@@ -80,16 +88,34 @@ const UploadQueuePanel: React.FC<UploadQueuePanelProps> = ({
) : null}
<div className="xx-upload-queue-status">
{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]}`
: ""}
</div>
{it.status === "error" && it.error ? (
<div className="xx-upload-queue-error-detail" title={it.error}>
{it.error.split("\n").map((line, idx) =>
line === COMPLETE_RETRY_HINT ? (
<div key={idx} className="xx-upload-queue-error-hint">
{line}
</div>
) : (
<div key={idx}>{line}</div>
),
)}
</div>
) : null}
</div>
<span className="xx-upload-queue-actions">
{it.status === "error" && (
<button
type="button"
className="xx-upload-queue-btn"
title="重试"
title={
it.failedStage === "complete" ? "安全重试(只确认,不重新上传)" : "重试上传"
}
onClick={() => onRetry(it.tempId)}
>
<ReloadOutlined />
+218 -33
View File
@@ -2,34 +2,65 @@ 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,
findDuplicateInQueue,
makeClientUploadId,
makeFileFingerprint,
} from "@/api/assets/uploadDedup"
/** 单文件上传状态机 */
export type UploadItemStatus = "preparing" | "uploading" | "ingesting" | "done" | "error"
/** 失败发生的阶段:complete 阶段失败时记录可能已在后端建成,禁止盲目重传整个文件 */
export type UploadFailStage = "prepare" | "transfer" | "complete"
export interface UploadItem {
/** 前端临时 idprepare 前无 asset_id 时用) */
/** 前端临时 idprepare 前无 asset_id 时用),同时作为队列项 key */
tempId: string
file: File
fileName: string
/** 进度 0~100(仅直传阶段有真实进度) */
progress: number
status: UploadItemStatus
/** 文件指纹(name+size+lastModified),入队去重用 */
fileKey: string
/** 本次逻辑上传的幂等 token:重试复用、重新入队才换新 */
clientUploadId: string
/** 上传前算好的文件内容哈希(SHA-256),prepare/complete 都带上 */
fileHash?: string
/** 后端 prepare 预建的 asset id(旧后端可能为空) */
assetId?: string
/** 去重命中:complete 返回 duplicated,标记完成但不产生新素材 */
duplicated?: boolean
/** 失败发生的阶段;complete 阶段失败点重试只重发 complete,不重新上传文件 */
failedStage?: UploadFailStage
/** 状态行补充提示(如"正在计算文件指纹…" */
hint?: string
error?: string
}
/** 批量直传最大并发数,避免多文件瓜分上行带宽 */
const MAX_CONCURRENT = 3
/** complete 阶段失败后的安全提示:素材可能已在后端建成,重试只重发 complete 幂等安全 */
export const COMPLETE_RETRY_HINT = "素材可能已在服务器处理中,点重试将安全确认,不会重新上传文件"
/** 失败阶段中文名(toast 提示用,明确失败发生在哪一步) */
const STAGE_LABEL: Record<UploadFailStage, string> = {
prepare: "准备上传",
transfer: "文件传输",
complete: "确认入库",
}
/**
* 素材批量上传 Hook
* - prepare 阶段后端预建 status=uploading 的 asset,前端拿到 asset_id 立即刷新列表
* - 入队按文件指纹(name+size+lastModified)去重:同一文件已在队列/上传中/处理中时不重复入队
* - 上传前计算文件 SHA-256prepare/complete 携带 file_hash + 幂等 tokenclientUploadId
* - complete 超时/失败不盲目重传:复用 handle 只重发 complete(幂等),prepare/transfer 失败才全量重跑
* - OSS 直传并发限制为 3,其余排队;每个文件独立进度/状态
* - complete 后素材进入转码(ingesting/processing),由列表轮询反映
* - 失败卡片支持重试/移除
*/
export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
@@ -39,6 +70,13 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
const itemsRef = useRef<UploadItem[]>([])
itemsRef.current = items
/**
* prepare 成功后的 handle 按 tempId 留存:
* complete 阶段失败(超时/网络)时 OSS 文件已存在、后端记录也可能已建成,
* 重试必须复用同一 handle 只重发 complete,绝不能重新 prepare+直传。
*/
const handlesRef = useRef<Map<string, DirectUploadHandle>>(new Map())
const updateItem = useCallback((tempId: string, patch: Partial<UploadItem>) => {
setItems((prev) => prev.map((it) => (it.tempId === tempId ? { ...it, ...patch } : it)))
}, [])
@@ -52,14 +90,71 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
}, [queryClient, effectiveLibId])
/** 执行单个文件的完整上传流程(prepare→transfer→complete */
/**
* 执行单个文件的完整上传流程。
* @param completeOnly complete 阶段失败后的重试:跳过 hash/prepare/transfer,只重发 complete
* (OSS 文件已传完,重发由 file_hash + clientUploadId 保证幂等)
*/
const runUpload = useCallback(
async (item: UploadItem, handle?: DirectUploadHandle) => {
async (item: UploadItem, handle?: DirectUploadHandle, completeOnly = false) => {
let stage: UploadFailStage = "prepare"
try {
// 1. prepare(重试时复用已准备的 handle 也行,但签名可能过期,重新 prepare 最稳)
const h =
handle ??
(await prepareDirectUploadHandle({ file: item.file, library_id: effectiveLibId }))
let h = handle
if (completeOnly && h) {
// ── complete 重试:文件已在 OSS,直接幂等重发确认 ──
stage = "complete"
updateItem(item.tempId, {
status: "ingesting",
progress: 100,
error: undefined,
failedStage: undefined,
hint: undefined,
})
const result = await h.complete()
refreshList()
handlesRef.current.delete(item.tempId)
if (result.duplicated) {
updateItem(item.tempId, { status: "done", duplicated: true, assetId: result.asset_id })
message.info(`"${item.fileName}" 与素材库已有内容相同,已跳过`)
} else {
updateItem(item.tempId, { status: "done", assetId: result.asset_id || item.assetId })
message.success(`"${item.fileName}" 上传完成,正在转码处理`)
}
return
}
// 1. 计算文件内容哈希(失败不阻塞,降级为不传 hash;后端仍有幂等 token 兜底)
updateItem(item.tempId, { hint: "正在计算文件指纹…" })
const fileHash = item.fileHash || (await computeFileHash(item.file))
updateItem(item.tempId, { fileHash, hint: undefined })
// 2. prepare(携带 file_hash + 幂等 token;重试时复用同一 clientUploadId
stage = "prepare"
h =
h ??
(await prepareDirectUploadHandle({
file: item.file,
library_id: effectiveLibId,
fileHash,
clientUploadId: item.clientUploadId,
}))
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",
@@ -72,26 +167,59 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
updateItem(item.tempId, { status: "uploading", progress: 0 })
}
// 2. OSS 直传(真实进度)
// 3. OSS 直传(真实进度)
stage = "transfer"
await h.transfer((pct) => updateItem(item.tempId, { progress: pct }))
// 3. complete:后端创建 ingest job,素材进入转码
// 4. complete:后端确认入库并创建 ingest jobfile_hash + 幂等 token 已在 handle 闭包中)
stage = "complete"
updateItem(item.tempId, { status: "ingesting", progress: 100 })
const result = await h.complete()
refreshList()
handlesRef.current.delete(item.tempId)
if (result.duplicated) {
updateItem(item.tempId, { status: "done", duplicated: true, assetId: result.asset_id })
message.info(`"${item.fileName}" 与素材库已有内容相同,已跳过`)
} else {
updateItem(item.tempId, { status: "done" })
updateItem(item.tempId, { status: "done", assetId: result.asset_id || item.assetId })
message.success(`"${item.fileName}" 上传完成,正在转码处理`)
}
} catch (err: unknown) {
const detail = err instanceof Error ? err.message : "上传失败"
console.error("[useAssetUpload] 上传失败:", item.fileName, err)
updateItem(item.tempId, { status: "error", error: detail })
message.error(`"${item.fileName}" 上传失败:${detail}`)
// 完整失败原因: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: `${detail}\n${COMPLETE_RETRY_HINT}`,
hint: undefined,
})
if (!isErrorMsgShown(err)) {
message.error(`"${item.fileName}" 确认入库失败:${detail}`)
}
} else {
// prepare / transfer 失败:后端尚无素材记录,可安全全量重跑
handlesRef.current.delete(item.tempId)
updateItem(item.tempId, {
status: "error",
failedStage: stage === "transfer" ? "transfer" : "prepare",
error: detail,
hint: undefined,
})
// 拦截器已对后端错误弹过 toast(含真实 detail)时不重复弹;
// OSS XHR 直传错误不走 axios,必须在这里弹
if (!isErrorMsgShown(err)) {
message.error(`"${item.fileName}" ${STAGE_LABEL[stage]}失败:${detail}`)
}
}
}
},
[effectiveLibId, refreshList, updateItem],
@@ -113,7 +241,9 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
if (!next) return
claimedRef.current.add(next.tempId)
inFlightRef.current += 1
void runUpload(next).finally(() => {
// complete 阶段失败的重试:复用留存的 handle,只重发 complete
const existingHandle = handlesRef.current.get(next.tempId)
void runUpload(next, existingHandle, existingHandle !== undefined).finally(() => {
inFlightRef.current -= 1
claimedRef.current.delete(next.tempId)
// 一个任务结束(成功/失败)后继续拉起排队任务
@@ -126,7 +256,11 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
pumpRef.current()
}, [items])
/** 入队一个或多个文件 */
/**
* 入队一个或多个文件(按文件指纹去重):
* - 同一文件已在队列且 preparing/uploading/ingesting/done → 跳过,不重复入队
* - 同一文件此前失败(error)→ 重新激活原队列项(复用 clientUploadId,保持幂等语义)
*/
const enqueueUploads = useCallback(
(files: File[]) => {
if (!effectiveLibId) {
@@ -143,37 +277,86 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
}
if (valid.length === 0) return
const newItems: UploadItem[] = valid.map((file, idx) => ({
tempId: `${Date.now()}-${idx}-${Math.random().toString(36).slice(2, 8)}`,
file,
fileName: file.name,
progress: 0,
status: "preparing",
}))
setItems((prev) => [...prev, ...newItems])
let skipped = 0
let rearmed = 0
const newItems: UploadItem[] = []
for (const file of valid) {
const fileKey = makeFileFingerprint(file)
// error 项允许重新激活;其余状态(preparing/uploading/ingesting/done)都算重复
const dup = findDuplicateInQueue([...itemsRef.current, ...newItems], fileKey, ["error"])
if (dup) {
skipped += 1
continue
}
// 失败项重新激活:复用 tempId/clientUploadId,由 pump 按留存 handle 决定重试方式
const failed = itemsRef.current.find(
(it) => it.fileKey === fileKey && it.status === "error",
)
if (failed) {
rearmed += 1
updateItem(failed.tempId, {
status: "preparing",
progress: 0,
error: undefined,
failedStage: undefined,
hint: undefined,
})
continue
}
newItems.push({
tempId: `${Date.now()}-${newItems.length}-${Math.random().toString(36).slice(2, 8)}`,
file,
fileName: file.name,
progress: 0,
status: "preparing",
fileKey,
clientUploadId: makeClientUploadId(),
})
}
if (newItems.length > 0) {
setItems((prev) => [...prev, ...newItems])
}
if (skipped > 0) {
message.warning(`已跳过 ${skipped} 个重复文件(已在上传队列、处理中或本页已上传)`)
}
if (rearmed > 0) {
message.info(`已重新加入 ${rearmed} 个此前失败的文件`)
}
},
[effectiveLibId],
[effectiveLibId, updateItem],
)
/** 重试失败任务 */
/**
* 重试失败任务(仅限 status=error):
* - complete 阶段失败:复用留存 handle 只重发 complete(幂等,不重新上传)
* - prepare/transfer 阶段失败:全量重跑(后端尚无记录,安全)
*/
const retryUpload = useCallback(
(tempId: string) => {
const target = itemsRef.current.find((it) => it.tempId === tempId)
if (!target) return
updateItem(tempId, { status: "preparing", progress: 0, error: undefined })
// 状态更新后由 useEffect 触发 pump
if (!target || target.status !== "error") return
updateItem(tempId, { status: "preparing", progress: 0, error: undefined, hint: undefined })
// 状态更新后由 useEffect 触发 pumppump 会按 handlesRef 自动选择 completeOnly / 全量
},
[updateItem],
)
/** 从上传列表移除(已进入转码的由素材网格管理;这里只移除上传面板记录) */
const removeUpload = useCallback((tempId: string) => {
handlesRef.current.delete(tempId)
setItems((prev) => prev.filter((it) => it.tempId !== tempId))
}, [])
/** 清空已完成/去重记录 */
const clearFinished = useCallback(() => {
setItems((prev) => prev.filter((it) => it.status !== "done"))
setItems((prev) => {
for (const it of prev) {
if (it.status === "done") handlesRef.current.delete(it.tempId)
}
return prev.filter((it) => it.status !== "done")
})
}, [])
const activeCount = items.filter(
@@ -188,8 +371,10 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
retryUpload,
removeUpload,
clearFinished,
/** 是否有进行中的上传(用于上传区文案) */
/** 是否有进行中的上传(用于上传区文案/禁用入口 */
uploading: hasActive,
/** 是否有文件正在本地处理或直传(用于禁用上传入口,防重复提交) */
transferActive: activeCount > 0,
activeCount,
pendingCount,
}
@@ -10,6 +10,26 @@ import {
import { getOrCreateDefaultProject } from "@/api/projects"
import { mapLibrary, mapAsset, type AssetItem, type LibraryItem } from "../types"
/** 处理中素材快速轮询(3s)的最大持续时间:超过后停止快轮询,避免孤儿任务永久转圈 */
const PROCESSING_POLL_MAX_MS = 10 * 60 * 1000 // 10 分钟
const isProcessingStatus = (st?: string | null): boolean =>
st === "uploading" || st === "ingesting" || st === "processing" || st === "pending"
/** 判断列表是否存在「创建超过 maxMs 仍在处理中」的卡死素材 */
const hasStalledProcessing = (
list: ApiAssetItem[],
maxMs: number = PROCESSING_POLL_MAX_MS,
): boolean => {
const now = Date.now()
return list.some((a) => {
if (!isProcessingStatus(a.status ?? "")) return false
if (!a.created_at) return false
const created = new Date(a.created_at).getTime()
return Number.isFinite(created) && now - created > maxMs
})
}
/**
* 素材库数据 Hook
* 封装视频库列表、素材列表的数据查询,以及筛选、搜索状态管理
@@ -63,15 +83,15 @@ export function useAssetsData() {
}),
enabled: !!effectiveLibId,
staleTime: 30_000,
// 列表中存在上传中/转码中素材时每 3s 轮询;全部就绪后自动停止
// 列表中存在上传中/转码中素材时每 3s 轮询;全部就绪后自动停止
// 但若处理中素材创建已超过 10 分钟仍未就绪(疑似后端卡住/孤儿任务),
// 停止快轮询避免无限转圈——卡死素材在网格中显示「处理超时」提示。
refetchInterval: (query) => {
const data = query.state.data as { items: ApiAssetItem[] } | undefined
const items = data?.items ?? []
const processing = items.some((a) => {
const st = a.status ?? ""
return st === "uploading" || st === "ingesting" || st === "processing" || st === "pending"
})
return processing ? 3000 : false
const processing = items.some((a) => isProcessingStatus(a.status ?? ""))
if (!processing) return false
return hasStalledProcessing(items) ? false : 3000
},
})
@@ -80,6 +100,21 @@ export function useAssetsData() {
[apiAssets],
)
/** 创建超过 10 分钟仍在处理中的素材(后端可能卡住),网格提示「处理超时」 */
const stalledAssetIds = useMemo(() => {
const list = Array.isArray(apiAssets?.items) ? apiAssets.items : []
const ids = new Set<string>()
const now = Date.now()
for (const a of list) {
if (!isProcessingStatus(a.status ?? "") || !a.created_at) continue
const created = new Date(a.created_at).getTime()
if (Number.isFinite(created) && now - created > PROCESSING_POLL_MAX_MS) {
ids.add(a.id)
}
}
return ids
}, [apiAssets])
/* ── 筛选状态 ── */
const [searchText, setSearchText] = useState("")
const [filterType, setFilterType] = useState<string>("all")
@@ -129,6 +164,8 @@ export function useAssetsData() {
assetsError,
assetsErrorObj,
refetchAssets,
stalledAssetIds,
hasStalledAssets: stalledAssetIds.size > 0,
// 筛选
searchText,
setSearchText,
+32 -24
View File
@@ -5,8 +5,8 @@ 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 Button from "@/components/ui/Button"
import WechatQrModal from "@/components/auth/WechatQrModal"
import "./Login.css"
interface LoginFormValues {
@@ -19,7 +19,7 @@ const Login: React.FC = () => {
const navigate = useNavigate()
const loginMutation = useLogin()
const [form] = Form.useForm()
const [wechatLoading, setWechatLoading] = useState(false)
const [wechatQrOpen, setWechatQrOpen] = useState(false)
const onFinish = async (values: LoginFormValues) => {
try {
@@ -35,27 +35,28 @@ const Login: React.FC = () => {
}
}
const handleWechatLogin = async () => {
try {
setWechatLoading(true)
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) {
if (!(error as { __msgShown?: boolean })?.__msgShown)
message.error("微信登录暂不可用,请稍后重试")
} finally {
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")
}
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 (
@@ -126,10 +127,10 @@ const Login: React.FC = () => {
type="button"
className="xx-btn-wechat"
onClick={handleWechatLogin}
disabled={wechatLoading}
disabled={wechatQrOpen}
>
<span className="xx-wechat-icon">💬</span>
{wechatLoading ? "加载中..." : "微信登录"}
</button>
</div>
@@ -137,6 +138,13 @@ const Login: React.FC = () => {
<Link to="/register"></Link>
</div>
</div>
<WechatQrModal
open={wechatQrOpen}
scene="login"
onClose={() => setWechatQrOpen(false)}
onLoginSuccess={handleWechatQrSuccess}
/>
</div>
)
}
@@ -0,0 +1,118 @@
/**
* 微信绑定回调页(已登录用户在设置页发起"绑定微信"扫码后回到这里)
* 用 code 调绑定接口把微信关联到当前账号,成功后回设置页
*
* 两种运行环境:
* - 整页跳转授权(旧流程/兜底):本页整页加载,成功/失败后 navigate 回设置页
* - 弹窗内嵌二维码(WxLogin self_redirect):本页在同源 iframe 内加载,
* 结果通过 postMessage 通知父窗口弹窗,不做页面导航
*/
import React, { useEffect, useState } from "react"
import { useSearchParams, useNavigate } from "react-router-dom"
import { Spin } from "antd"
import { bindWechat, normalizeUser } from "@/api/auth"
import { getErrorMessage } from "@/api/errors"
import { useAuthStore } from "@/store/authStore"
import { isInIframe, postWechatQrResult } from "@/components/auth/WechatQrModal/messages"
const WechatBindCallback: React.FC = () => {
const [searchParams] = useSearchParams()
const navigate = useNavigate()
const setUser = useAuthStore((state) => state.setUser)
const [error, setError] = useState<string | null>(null)
const inIframe = isInIframe()
useEffect(() => {
const code = searchParams.get("code")
const state = searchParams.get("state")
const fail = (message: string) => {
if (inIframe) {
// 弹窗模式:把真实原因上报父窗口在 Modal 内展示
postWechatQrResult("bind", false, { detail: message })
return
}
setError(message)
}
if (!code || !state) {
fail("无效的回调参数,请回到设置页重新扫码绑定")
return
}
const handleBind = async () => {
// state 校验由后端 state store 一次性消费兜底(前端不再比对 localStorage
// 微信内打开/跨浏览器场景本地无 state 会误杀);清理绑定前写入的 state
localStorage.removeItem("wechat_bind_state")
try {
const result = await bindWechat(code, state)
setUser(normalizeUser(result.user))
if (inIframe) {
// 弹窗模式:通知父窗口关闭弹窗并刷新绑定状态
postWechatQrResult("bind", true)
return
}
// 用 replace 回设置页,query 携带成功标记由设置页提示
navigate("/app/profile?wechat_bind=success", { replace: true })
} catch (err) {
// 绑定失败直接在本页展示/上报真实原因(如微信已被其他账号绑定),不静默跳走
fail(`微信绑定失败:${getErrorMessage(err, "请回到设置页重试")}`)
}
}
handleBind()
}, [searchParams, navigate, setUser, inIframe])
if (error) {
return (
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<p style={{ color: "#ef4444", fontSize: 16, marginBottom: 16 }}>{error}</p>
<button
onClick={() => navigate("/app/profile")}
style={{
padding: "8px 24px",
background: "var(--primary-color, #3b82f6)",
color: "white",
border: "none",
borderRadius: 6,
cursor: "pointer",
}}
>
</button>
</div>
</div>
)
}
return (
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<Spin size="large" />
<p style={{ marginTop: 16, color: "#666" }}>...</p>
</div>
</div>
)
}
export default WechatBindCallback
+83 -77
View File
@@ -1,40 +1,63 @@
/**
* 微信登录回调页
* 扫码授权后由微信重定向回来:用 code 换登录态,
* 新用户/资料未完善 → 跳昵称引导页;老用户 → 回来源页/首页
*
* 两种运行环境:
* - 整页跳转授权(旧流程/兜底):本页整页加载,按上述逻辑导航
* - 弹窗内嵌二维码(WxLogin self_redirect):本页在同源 iframe 内加载,
* 成功/失败均通过 postMessage 通知父窗口弹窗,不做页面导航
*/
import React, { useEffect, useState } from "react"
import { useSearchParams, useNavigate } from "react-router-dom"
import { Spin, message } from "antd"
import { Spin } from "antd"
import { wechatCallback, getCurrentUser, normalizeUser, type User } from "@/api/auth"
import { getErrorMessage } from "@/api/errors"
import { useAuthStore } from "@/store/authStore"
import BindContactModal from "@/components/auth/BindContactModal"
import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh"
import { isInIframe, postWechatQrResult } from "@/components/auth/WechatQrModal/messages"
const WechatCallback: React.FC = () => {
const [searchParams] = useSearchParams()
const navigate = useNavigate()
const setAuth = useAuthStore((state) => state.setAuth)
const [loading, setLoading] = useState(true)
const [showBindModal, setShowBindModal] = useState(false)
const [error, setError] = useState<string | null>(null)
const inIframe = isInIframe()
useEffect(() => {
const code = searchParams.get("code")
const state = searchParams.get("state")
if (!code || !state) {
setError("无效的回调参数")
const fail = (message: string) => {
if (inIframe) {
// 弹窗模式:把真实原因上报父窗口在 Modal 内展示,本页保持"处理中"即可
postWechatQrResult("login", false, { detail: message })
return
}
setError(message)
setLoading(false)
}
// 微信重定向出错时(如用户拒绝授权 error=access_denied)直接展示/上报原因
const wxErrorCode = searchParams.get("error")
const wxErrDesc = searchParams.get("error_description")
if (wxErrorCode || wxErrDesc) {
const reason = [wxErrorCode, wxErrDesc].filter(Boolean).join("")
fail(`微信授权失败:${reason}`)
return
}
if (!code || !state) {
fail("无效的回调参数,请重新扫码登录")
return
}
const handleCallback = async () => {
try {
// 校验 state,防止 CSRF
const savedState = localStorage.getItem("wechat_state")
if (!savedState || savedState !== state) {
setError("安全校验失败,请重新登录")
setLoading(false)
return
}
// state CSRF 校验由后端 state store 一次性消费兜底(前端不再比对
// localStorage——微信内打开、跨浏览器等场景本地没有 state,会误杀正常回调);
// 清理登录前写入的 state,避免残留
localStorage.removeItem("wechat_state")
const result = await wechatCallback(code, state)
@@ -49,41 +72,34 @@ const WechatCallback: React.FC = () => {
const userData = await getCurrentUser()
const user: User = normalizeUser(userData)
setAuth(user, result.access_token, result.refresh_token)
scheduleProactiveRefresh()
if (result.binding_complete) {
// 已绑定,跳转到登录前页面或首页
message.success("登录成功")
const redirect = localStorage.getItem("login_redirect") || "/"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
} else {
// 未绑定,显示绑定弹窗
setLoading(false)
setShowBindModal(true)
// 新用户 或 资料未完善(如上次中断没填昵称)→ 强制昵称引导
const needOnboarding = result.is_new_user || user.profile_completed === false
if (inIframe) {
// 弹窗模式:token 已写入同源 localStorage,通知父窗口同步登录态并跳转
postWechatQrResult("login", true, { needOnboarding })
return
}
if (needOnboarding) {
navigate("/welcome/wechat", { replace: true })
return
}
// 老用户:回登录前页面或首页
const redirect = localStorage.getItem("login_redirect") || "/"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
} catch (err) {
setError("登录失败,请重试")
setLoading(false)
// 透传后端真实错误(如 state 过期、code 已消费、接口异常),禁止吞成通用提示
fail(`微信登录失败:${getErrorMessage(err, "请重试或更换登录方式")}`)
}
}
handleCallback()
}, [searchParams, navigate, setAuth])
const handleBindSuccess = (user: User) => {
const setUser = useAuthStore.getState().setUser
setUser(user)
setShowBindModal(false)
message.success("绑定成功")
const redirect = localStorage.getItem("login_redirect") || "/"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
}
const handleBindCancel = () => {
setShowBindModal(false)
navigate("/login")
}
}, [searchParams, navigate, setAuth, inIframe])
if (loading) {
return (
@@ -98,49 +114,39 @@ const WechatCallback: React.FC = () => {
>
<div style={{ textAlign: "center" }}>
<Spin size="large" />
<p style={{ marginTop: 16, color: "#666" }}>...</p>
</div>
</div>
)
}
if (error) {
return (
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<p style={{ color: "#ef4444", fontSize: 16, marginBottom: 16 }}>{error}</p>
<button
onClick={() => navigate("/login")}
style={{
padding: "8px 24px",
background: "var(--primary-color, #3b82f6)",
color: "white",
border: "none",
borderRadius: 6,
cursor: "pointer",
}}
>
</button>
<p style={{ marginTop: 16, color: "#666" }}>...</p>
</div>
</div>
)
}
return (
<BindContactModal
open={showBindModal}
onSuccess={handleBindSuccess}
onCancel={handleBindCancel}
/>
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<p style={{ color: "#ef4444", fontSize: 16, marginBottom: 16 }}>{error}</p>
<button
onClick={() => navigate("/login")}
style={{
padding: "8px 24px",
background: "var(--primary-color, #3b82f6)",
color: "white",
border: "none",
borderRadius: 6,
cursor: "pointer",
}}
>
</button>
</div>
</div>
)
}
@@ -0,0 +1,116 @@
/**
* 微信新用户昵称引导页
* 新微信用户首次登录后强制填写昵称,完成后才进入主界面
*/
import React, { useRef } from "react"
import { Form, Input, message } from "antd"
import { Navigate, useNavigate } from "react-router-dom"
import { useMutation } from "@tanstack/react-query"
import { updateProfile } from "@/api/auth"
import { getErrorMessage, isErrorMsgShown } from "@/api/errors"
import { useAuthStore } from "@/store/authStore"
import Button from "@/components/ui/Button"
import "./Login.css"
interface OnboardingFormValues {
display_name: string
}
const WechatOnboarding: React.FC = () => {
const navigate = useNavigate()
const setUser = useAuthStore((state) => state.setUser)
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
const user = useAuthStore((state) => state.user)
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
const [form] = Form.useForm<OnboardingFormValues>()
// 同步防连点守卫:antd loading 要等 React 重渲染后才禁用按钮,
// 连点两次时第一次的 mutation 刚触发、重渲染未发生,第二次 click 仍会进来
// (截图里 PATCH /me 405 出现两次就是连点导致的重复提交)
const submittingRef = useRef(false)
const saveMutation = useMutation({
mutationFn: (displayName: string) => updateProfile({ display_name: displayName }),
})
// 已登录且资料已完善的用户不该停留在引导页
if (isAuthenticated && hasAccessToken && user?.profile_completed === true) {
return <Navigate to="/app/dashboard" replace />
}
// 未登录(如手动输入 URL)回登录页
if (!isAuthenticated || !hasAccessToken) {
return <Navigate to="/login" replace />
}
const onFinish = async (values: OnboardingFormValues) => {
if (submittingRef.current) return
submittingRef.current = true
try {
const updated = await saveMutation.mutateAsync(values.display_name.trim())
// 后端返回的 profile_completed 以最新资料为准,前端同步标记完善
setUser({ ...updated, profile_completed: true })
message.success("欢迎加入小虾智剪!")
const redirect = localStorage.getItem("login_redirect") || "/app/dashboard"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
} catch (err) {
// 透传后端真实原因(如接口异常/校验失败);拦截器已弹过的不重复弹
if (!isErrorMsgShown(err)) {
message.error(`昵称保存失败:${getErrorMessage(err, "请稍后重试")}`)
}
submittingRef.current = false
}
// 成功时页面跳走,不复位
}
return (
<div className="xx-auth-page">
<div className="xx-auth-card">
<div className="xx-auth-header">
<div className="xx-auth-brand">
<span className="xx-auth-logo">🦐</span>
<span className="xx-auth-brand-name"></span>
</div>
<p>使</p>
</div>
<Form
form={form}
name="wechat-onboarding"
onFinish={onFinish}
autoComplete="off"
layout="vertical"
// 不预填:新微信用户必须自己输入昵称(user.display_name 可能是微信昵称/系统占位)
initialValues={{ display_name: "" }}
>
<Form.Item
name="display_name"
label="昵称"
rules={[
{ required: true, message: "请输入昵称" },
{ whitespace: true, message: "昵称不能为空白" },
{ min: 1, max: 20, message: "昵称长度需在 1-20 个字符之间" },
]}
extra="昵称将展示在您的作品和账户中,之后可在个人设置中修改"
>
<Input placeholder="请输入您的昵称" size="large" maxLength={20} showCount />
</Form.Item>
<Form.Item>
<Button
buttonType="primary"
buttonSize="lg"
htmlType="submit"
loading={saveMutation.isPending}
disabled={saveMutation.isPending}
style={{ width: "100%" }}
>
{saveMutation.isPending ? "保存中..." : "进入小虾智剪"}
</Button>
</Form.Item>
</Form>
</div>
</div>
)
}
export default WechatOnboarding
+245 -67
View File
@@ -1,12 +1,12 @@
/**
* 智能剪辑页面 — 前端实时预览架构
* 6 步向导:选择模板 → 素材 → 配音 → 标题(含预览) → 确认生成 → 选择封面
* 智能剪辑页面Issue #1677 多视频批量生成,修正版)
* 固定 6 步向导:模板(弹数量) → 素材 → 配音 → 标题 → 确认生成 → 封面,单视频与批量完全一致
*
* 架构:
* - 步骤 4 右侧显示 FrontendPreviewPlayer 实时预览
* - 步骤 5 右侧内联播放生成中的/最终视频
* - 步骤 6 封面从最终成片中智能选帧(MediaKit)
* - 点"确认生成"时调用 createGenerationTask 创建一次服务器渲染任务
* - 预览全部为纯前端 Canvas 实时播放(FrontendPreviewPlayer),不调任何后端渲染接口:
* N=1 单播放器;N>1 CanvasPreviewGridvariantSeed 让素材排布/起始点不同,画面有差异)
* - 步骤5确认生成:正式生成接口(count + titles[]/voice_library_ids[]/cover_urls[]),
* 批量时逐任务独立进度/失败重试(BatchGenerationGrid
*/
import React, { useMemo, useState, useEffect, useRef, useCallback } from "react"
import { message } from "antd"
@@ -17,12 +17,15 @@ import { useCloneProgress } from "@/hooks/useCloneProgress"
import CloneModal from "@/components/voice/CloneModal"
import GenerateHeader from "./components/GenerateHeader"
import FrontendPreviewPlayer from "./components/FrontendPreviewPlayer"
import CanvasPreviewGrid from "./components/CanvasPreviewGrid"
import PreviewCountModal from "./components/PreviewCountModal"
import GenerateStepsBar from "./components/GenerateStepsBar"
import GenerateStepContent from "./components/GenerateStepContent"
import GenerateStepActions from "./components/GenerateStepActions"
import { useGenerateFormState } from "./hooks/useGenerateFormState"
import { useStepNavigation } from "./hooks/useStepNavigation"
import { useGenerateVideo } from "./hooks/useGenerateVideo"
import { usePreviewAssets } from "./hooks/usePreviewAssets"
import { useTitleStyleUpdaters } from "./hooks/useStep4Title/useTitleStyleUpdaters"
import { getAssetsByKind } from "@/api/assets"
@@ -53,10 +56,8 @@ const GeneratePage: React.FC = () => {
selectedVoice,
setSelectedVoice,
voiceMode,
setVoiceMode,
selectedClonedVoice,
setSelectedClonedVoice,
presetVoices,
cloneModalOpen,
setCloneModalOpen,
videoRatio,
@@ -72,8 +73,51 @@ const GeneratePage: React.FC = () => {
setStoredSourceEditPlanId,
serverClips,
setServerClips,
previewCount,
setPreviewCount,
previewTitles,
setPreviewTitles,
voiceModePerVideo,
setVoiceModePerVideo,
voiceLibraryIds,
setVoiceLibraryIds,
previewCovers,
setPreviewCovers,
selectedVariantIds,
setSelectedVariantIds,
} = formState
const isBatch = previewCount > 1
/* ── 配音选择同步:共用配音 ↔ 变体数组 ── */
// 触发场景:①共用配音变化 ②批量模式进入/退出 ③独立→共用切换(需把所有变体刷成共用配音)
// 独立模式下:仅同步变体[0](其选择器绑定共用配音),用户单独选择的其他变体不覆盖
const prevVoiceSyncRef = useRef({
voice: selectedVoice,
batch: isBatch,
perVideo: voiceModePerVideo,
})
useEffect(() => {
const prev = prevVoiceSyncRef.current
const voiceChanged = prev.voice !== selectedVoice
const modeChanged = prev.batch !== isBatch || prev.perVideo !== voiceModePerVideo
prevVoiceSyncRef.current = { voice: selectedVoice, batch: isBatch, perVideo: voiceModePerVideo }
if (!voiceChanged && !modeChanged) return
if (!isBatch) return
if (!voiceModePerVideo) {
// 共用模式(含刚从独立切回):所有变体跟随共用配音,未选择的补默认值
setVoiceLibraryIds((prevIds) => (prevIds || []).map((id) => id || selectedVoice))
} else if (voiceChanged) {
// 独立模式下共用配音变化:仅同步变体[0](与共用选择器绑定),其余不覆盖
setVoiceLibraryIds((prevIds) =>
(prevIds || []).map((id, i) => (i === 0 ? selectedVoice : id)),
)
}
}, [selectedVoice, isBatch, voiceModePerVideo, setVoiceLibraryIds])
/* ── 数量选择弹窗 ── */
const [countModalOpen, setCountModalOpen] = useState(false)
/* ── 标题样式回调 ── */
const styleUpdaters = useTitleStyleUpdaters({
titleSettings,
@@ -88,6 +132,8 @@ const GeneratePage: React.FC = () => {
const [previewVoiceAudioUrl, setPreviewVoiceAudioUrl] = useState<string | null>(null)
const ttsAbortRef = useRef<AbortController | null>(null)
// TTS 试听文案:批量跟随变体0标题(仅取首项,避免编辑其他变体标题触发多余 TTS 请求)
const variant0Title = isBatch ? previewTitles?.[0] || "" : ""
useEffect(() => {
const voiceAsset = voiceMaterials.find((m) => m.id === selectedVoice)
@@ -96,8 +142,10 @@ const GeneratePage: React.FC = () => {
return
}
// 批量模式下 TTS 文案跟随变体0标题;单视频跟随主标题
const ttsTitle = isBatch ? variant0Title || "" : titleSettings.title
const voiceId = selectedClonedVoice || selectedVoice
if (!voiceId || !titleSettings.title) {
if (!voiceId || !ttsTitle) {
setPreviewVoiceAudioUrl(null)
return
}
@@ -107,7 +155,7 @@ const GeneratePage: React.FC = () => {
ttsAbortRef.current = controller
let cancelled = false
previewTts({ text: titleSettings.title, voice_id: voiceId })
previewTts({ text: ttsTitle, voice_id: voiceId })
.then((res) => {
if (!cancelled && res.audio_url) {
setPreviewVoiceAudioUrl(res.audio_url)
@@ -124,10 +172,17 @@ const GeneratePage: React.FC = () => {
cancelled = true
controller.abort()
}
}, [selectedVoice, selectedClonedVoice, titleSettings.title, voiceMaterials])
}, [
selectedVoice,
selectedClonedVoice,
titleSettings.title,
variant0Title,
isBatch,
voiceMaterials,
])
/* ── 克隆声音 ── */
const { clones: clonedVoices, addClone, hasProcessing } = useCloneProgress()
const { addClone } = useCloneProgress()
const handleCloneSuccess = (voice: VoiceClone) => {
addClone(voice)
@@ -156,17 +211,25 @@ const GeneratePage: React.FC = () => {
[bgm, currentTemplate],
)
/* ── 加载素材详情(供前端预览播放器使用 + 配音时长校验 ── */
/* ── 加载素材详情(供前端预览播放器使用) ── */
const previewAssetsEnabled = previewAssetIds.length > 0
const { assets: previewAssets, ready: previewAssetsReady } = usePreviewAssets(
previewAssetIds,
previewAssetsEnabled,
)
/* ── 预览就绪:素材已加载,且有模板 ── */
const previewReady = useMemo(
() => previewAssetsReady && !!currentTemplate,
[previewAssetsReady, currentTemplate],
/* ── 预览就绪:纯前端 Canvas 预览,素材详情加载完即可秒开(单视频/批量一致) ── */
const previewReady = previewAssetsReady && !!currentTemplate
/* ── 勾选变体 ── */
const toggleVariantSelect = useCallback(
(index: number) => {
setSelectedVariantIds((prev) => {
const list = prev || []
return list.includes(index) ? list.filter((i) => i !== index) : [...list, index].sort()
})
},
[setSelectedVariantIds],
)
/* ── 视频生成核心逻辑 ── */
@@ -176,8 +239,10 @@ const GeneratePage: React.FC = () => {
generated,
generateError,
generatedVideos,
batchTasks,
generate: handleGenerate,
retry: handleRetryGenerate,
retryBatchTask: handleRetryBatchTask,
dismissError: handleDismissError,
download: handleDownload,
share: handleShare,
@@ -199,27 +264,92 @@ const GeneratePage: React.FC = () => {
sourceEditPlanId: storedSourceEditPlanId || sourceEditPlanId,
previewTaskId,
bgmConfig,
previewCount,
variantTitles: previewTitles,
variantVoiceLibraryIds: voiceLibraryIds,
voiceModePerVideo,
variantCoverUrls: previewCovers,
selectedVariantIndexes: isBatch ? selectedVariantIds : undefined,
onGenerationSuccess: () => {
setPreviewTaskId(null)
setStoredSourceEditPlanId(null)
},
})
/* ── 步骤4「确认生成视频」:校验标题/预览 → 创建最终渲染任务 → 成功后进入步骤5 ── */
/* ── 数量弹窗确认:设置数量 + 同步批量数组长度 + 进入步骤2 ── */
const handleCountConfirm = useCallback(
(count: number) => {
setPreviewCount(count)
setCountModalOpen(false)
// 同步批量数组长度
setPreviewTitles((prev) => {
const list = prev || []
const base = list[0] || titleSettings.title || ""
return Array.from({ length: count }, (_, i) => list[i] ?? (i === 0 ? base : ""))
})
setVoiceLibraryIds((prev) => {
const list = prev || []
return Array.from({ length: count }, (_, i) => list[i] ?? selectedVoice ?? "")
})
setPreviewCovers((prev) => {
const list = prev || []
return Array.from({ length: count }, (_, i) => list[i] ?? "")
})
setSelectedVariantIds(Array.from({ length: count }, (_, i) => i))
setCurrentStep(2)
},
[
setPreviewCount,
setPreviewTitles,
setVoiceLibraryIds,
setPreviewCovers,
setSelectedVariantIds,
setCurrentStep,
titleSettings.title,
selectedVoice,
],
)
/* ── 步骤4「确认生成视频」:校验通过 → 创建正式生成任务 → 跳步骤5看实时进展 ── */
const handleConfirmGenerate = useCallback(async () => {
if (!titleSettings.title.trim()) {
message.warning("请选择或输入标题")
return
// 标题校验:批量只校验已勾选的变体;单视频校验主标题
if (isBatch) {
if (selectedVariantIds.length === 0) {
message.warning("请至少勾选一个视频")
return
}
const missing = selectedVariantIds.some((i) => !previewTitles[i]?.trim())
if (missing) {
message.warning("请为每个勾选的视频输入标题")
return
}
} else if (!titleSettings.title?.trim()) {
// 与 buildPayload.validateGenerateInputs 一致:AI 自动选标题模式(aiAutoSelect
// 允许空标题由后端生成;手动模式必须填写,避免提交空标题
if (!titleSettings.aiAutoSelect) {
message.warning("请先选择或输入标题")
return
}
}
if (!previewReady) {
message.warning("预览视频正在加载,请稍候")
message.warning("预览素材正在加载,请稍候")
return
}
const ok = await handleGenerate()
if (ok) {
// 单视频与批量一致:任务创建成功后进入步骤5「确认生成」看实时渲染进展
setCurrentStep(5)
}
}, [titleSettings.title, previewReady, handleGenerate, setCurrentStep])
}, [
isBatch,
selectedVariantIds,
previewTitles,
titleSettings.aiAutoSelect,
titleSettings.title,
previewReady,
handleGenerate,
setCurrentStep,
])
/* ── 步骤导航 ── */
const { goNext, goPrev } = useStepNavigation({
@@ -230,13 +360,21 @@ const GeneratePage: React.FC = () => {
selectedMaterials,
smartSelectedIds,
titleSettings,
previewReady,
generated,
onOpenCountModal: () => setCountModalOpen(true),
})
/* ── 最终成片(步骤5/6 右侧播放) ── */
/* ── 最终成片(单视频右侧播放) ── */
const finalVideo = generatedVideos[0]
/* ── 布局 class:步骤4标题页=预览+标题侧栏;步骤5/6批量=整行宽;步骤1~3=整行宽 ── */
const layoutClassName = useMemo(() => {
if (currentStep < 4) return "xx-generate-layout full-width"
if (currentStep === 4) return "xx-generate-layout step4-layout"
// 步骤5/6:批量网格需要整行宽度;单视频保持 表单+右侧成片 两栏
return isBatch ? "xx-generate-layout full-width" : "xx-generate-layout"
}, [currentStep, isBatch])
/* ================================================================
渲染
================================================================ */
@@ -247,8 +385,61 @@ const GeneratePage: React.FC = () => {
<GenerateStepsBar currentStep={currentStep} onStepClick={setCurrentStep} />
<div className={`xx-generate-layout${currentStep < 4 ? " full-width" : ""}`}>
{/* ════ 左侧:表单区 ════ */}
<div className={layoutClassName}>
{/* ════ 步骤4:左侧预览大区域(纯前端 Canvas 实时预览) ════ */}
{currentStep === 4 && !!currentTemplate && (
<div className="xx-generate-preview-col">
{!isBatch ? (
/* 单视频:前端 Canvas 实时预览(与旧版一致,零回归) */
<FrontendPreviewPlayer
assets={previewAssets}
template={currentTemplate}
videoRatio={videoRatio}
ready={previewAssets.length > 0}
serverClips={serverClips}
voiceAudioUrl={previewVoiceAudioUrl || undefined}
titleSettings={{
title: titleSettings.title,
size: titleSettings.size,
font: titleSettings.font,
color: titleSettings.color,
position: titleSettings.position as "top" | "center" | "bottom" | "custom",
bold: titleSettings.bold,
italic: titleSettings.italic,
stroke: titleSettings.stroke,
shadow: titleSettings.shadow,
posX: titleSettings.posX,
posY: titleSettings.posY,
}}
onTitlePositionChange={styleUpdaters.updateTitlePosition}
/>
) : (
/* 批量:N 个前端 Canvas 预览网格(不调任何后端渲染接口,秒开) */
<div className="xx-form-section">
<div className="xx-preview-header">
<h3>🎬 {previewCount} </h3>
<span style={{ fontSize: 13, color: "var(--text-secondary, #666)" }}>
</span>
</div>
<CanvasPreviewGrid
count={previewCount}
assets={previewAssets}
template={currentTemplate}
videoRatio={videoRatio}
titles={previewTitles}
titleSettings={titleSettings}
voiceAudioUrl={previewVoiceAudioUrl || undefined}
selectedIds={selectedVariantIds}
onToggleSelect={toggleVariantSelect}
selectable={!generating}
/>
</div>
)}
</div>
)}
{/* ════ 右侧:步骤1~3 表单 / 步骤4 标题边栏 / 步骤5 确认生成进度 / 步骤6 封面 ════ */}
<div className="xx-generate-form">
<GenerateStepContent
currentStep={currentStep}
@@ -280,23 +471,25 @@ const GeneratePage: React.FC = () => {
selectedVoice={selectedVoice}
onSelectedVoiceChange={setSelectedVoice}
onServerClipsChange={setServerClips}
voiceMode={voiceMode}
onVoiceModeChange={setVoiceMode}
selectedClonedVoice={selectedClonedVoice}
onSelectedClonedVoiceChange={setSelectedClonedVoice}
clonedVoices={clonedVoices}
addClone={addClone}
hasProcessing={hasProcessing}
cloneModalOpen={cloneModalOpen}
onCloneModalOpenChange={setCloneModalOpen}
generating={generating}
generated={generated}
generateError={generateError}
progress={progress}
generatedVideos={generatedVideos}
onRetry={handleRetryGenerate}
onRetryBatchTask={handleRetryBatchTask}
onDismissError={handleDismissError}
presetVoices={presetVoices}
batchTasks={batchTasks}
previewCount={previewCount}
previewTitles={previewTitles}
onPreviewTitlesChange={setPreviewTitles}
voiceModePerVideo={voiceModePerVideo}
onVoiceModePerVideoChange={setVoiceModePerVideo}
voiceLibraryIds={voiceLibraryIds}
onVoiceLibraryIdsChange={setVoiceLibraryIds}
previewCovers={previewCovers}
onPreviewCoversChange={setPreviewCovers}
selectedVariantIds={selectedVariantIds}
/>
<GenerateStepActions
@@ -307,36 +500,13 @@ const GeneratePage: React.FC = () => {
generating={generating}
generated={generated}
generateError={generateError}
selectedCount={isBatch ? selectedVariantIds.length : 1}
/>
</div>
{/* ════ 右侧:步骤4实时预览,步骤5/6最终视频 ════ */}
<div className="xx-generate-right-col">
{currentStep === 4 && !!currentTemplate && (
<FrontendPreviewPlayer
assets={previewAssets}
template={currentTemplate}
videoRatio={videoRatio}
ready={previewAssets.length > 0}
serverClips={serverClips}
voiceAudioUrl={previewVoiceAudioUrl || undefined}
titleSettings={{
title: titleSettings.title,
size: titleSettings.size,
font: titleSettings.font,
color: titleSettings.color,
position: titleSettings.position as "top" | "center" | "bottom" | "custom",
bold: titleSettings.bold,
italic: titleSettings.italic,
stroke: titleSettings.stroke,
shadow: titleSettings.shadow,
posX: titleSettings.posX,
posY: titleSettings.posY,
}}
onTitlePositionChange={styleUpdaters.updateTitlePosition}
/>
)}
{currentStep >= 5 && generated && finalVideo && (
{/* ════ 步骤5/6(单视频):右侧成片播放器 ════ */}
{currentStep >= 5 && !isBatch && generated && finalVideo && (
<div className="xx-generate-right-col">
<div className="xx-inline-video-player">
<video
src={finalVideo.download_url || finalVideo.file_url}
@@ -360,10 +530,18 @@ const GeneratePage: React.FC = () => {
</button>
</div>
</div>
)}
</div>
</div>
)}
</div>
{/* 数量选择弹窗 */}
<PreviewCountModal
open={countModalOpen}
defaultCount={1}
onConfirm={handleCountConfirm}
onCancel={() => setCountModalOpen(false)}
/>
{/* 音色克隆弹窗 */}
<CloneModal
open={cloneModalOpen}
@@ -0,0 +1,100 @@
/**
* 第5步「确认生成」— 批量渲染进度网格(Issue #1677
*
* N 个正式生成任务各自独立卡片:进度条 / 成功成片播放 / 失败原因 + 单独重试。
* 数据来自 useGenerateVideo 的 batchTasksuseGenerationPolling 实时回传)。
*/
import React from "react"
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
import type { GeneratedVideo } from "@/api/template-editor"
interface BatchGenerationGridProps {
tasks: BatchTaskState[]
/** 变体标题(按变体序号取) */
titles: string[]
/** 失败任务重试 */
onRetryTask: (taskId: string) => void
}
const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
tasks,
titles,
onRetryTask,
}) => {
const sorted = [...tasks].sort(
(a, b) => Number(a.variantIndex || 0) - Number(b.variantIndex || 0),
)
return (
<div className="xx-form-section">
<div className="xx-preview-header">
<h3>🎬 {tasks.length} </h3>
<span style={{ fontSize: 13, color: "var(--text-secondary, #666)" }}>
{tasks.filter((t) => t.status === "completed").length} / {tasks.length}
</span>
</div>
<div className="xx-batch-gen-grid">
{sorted.map((task) => {
const title = titles[task.variantIndex] || `视频 ${task.variantIndex + 1}`
const video = (task.videos?.[0] || null) as GeneratedVideo | null
return (
<div key={task.taskId} className={`xx-batch-gen-card status-${task.status}`}>
<div className="xx-batch-gen-card-head">
<span className="xx-batch-gen-card-title" title={title}>
{task.status === "completed" ? (
<CheckCircleFilled style={{ color: "#52c41a", marginRight: 6 }} />
) : task.status === "failed" ? (
<CloseCircleOutlined style={{ color: "#ef4444", marginRight: 6 }} />
) : (
<LoadingOutlined style={{ color: "#1677ff", marginRight: 6 }} />
)}
{task.variantIndex + 1}{title}
</span>
</div>
<div className="xx-batch-gen-card-body">
{task.status === "running" && (
<>
<div className="xx-gen-progress-bar">
<div
className="xx-gen-progress-bar-fill"
style={{ width: `${Math.min(task.progress, 100)}%` }}
/>
</div>
<div className="xx-batch-gen-card-pct">{Math.round(task.progress)}%</div>
</>
)}
{task.status === "completed" && video && (
<video
src={video.download_url || video.file_url}
controls
style={{ width: "100%", borderRadius: 8, background: "#000", maxHeight: 280 }}
poster={video.thumbnail_url}
/>
)}
{task.status === "completed" && !video && (
<div className="xx-batch-gen-card-done"> </div>
)}
{task.status === "failed" && (
<div className="xx-batch-gen-card-failed">
<div className="xx-batch-gen-card-err">{task.error || "生成失败"}</div>
<button
type="button"
className="xx-btn xx-btn-primary xx-btn-sm"
onClick={() => onRetryTask(task.taskId)}
>
🔄
</button>
</div>
)}
</div>
</div>
)
})}
</div>
</div>
)
}
export default BatchGenerationGrid
@@ -0,0 +1,97 @@
/**
* 批量前端 Canvas 实时预览网格(Issue #1677 修正方案)
*
* N 个 FrontendPreviewPlayer 网格排列:
* - 纯前端 Canvas + video 元素实时播放素材片段,不调任何后端渲染接口
* - variantSeed 让每个变体素材排布/起始点不同,画面有可见差异
* - 各自叠加独立标题浮层(variantTitle),标题样式全局共用
* - 勾选框决定提交时生成哪些变体
*/
import React from "react"
import type { AssetItem } from "@/api/assets"
import type { EditingTemplate } from "@/api/editing-planner"
import type { TitleSettings } from "../types"
import FrontendPreviewPlayer from "./FrontendPreviewPlayer"
interface CanvasPreviewGridProps {
count: number
assets: AssetItem[]
template: EditingTemplate | null
videoRatio: string
titles: string[]
titleSettings: TitleSettings
/** 共用配音预览音频(仅第 1 个变体播放,避免多路音频重叠) */
voiceAudioUrl?: string
/** 勾选的变体序号 */
selectedIds: number[]
onToggleSelect: (index: number) => void
/** 生成中禁止勾选 */
selectable?: boolean
}
const CanvasPreviewGrid: React.FC<CanvasPreviewGridProps> = ({
count,
assets,
template,
videoRatio,
titles,
titleSettings,
voiceAudioUrl,
selectedIds,
onToggleSelect,
selectable = true,
}) => {
// count 上限已在源头 PreviewCountModal 的数量选择(1~MAX_PREVIEW_COUNT=10clamp
// 这里完整渲染所有变体,保证每个变体都有勾选/预览入口,UI 与数据不脱节
return (
<div className="xx-canvas-grid">
{Array.from({ length: count }, (_, i) => {
const checked = selectedIds.includes(i)
return (
<div
key={i}
className={`xx-canvas-grid-card${checked ? " selected" : ""}`}
data-variant={i}
>
<div className="xx-canvas-grid-card-bar">
<label className="xx-canvas-grid-check">
<input
type="checkbox"
checked={checked}
disabled={!selectable}
onChange={() => onToggleSelect(i)}
/>
<span> {i + 1}</span>
</label>
</div>
<FrontendPreviewPlayer
assets={assets}
template={template}
videoRatio={videoRatio}
ready={assets.length > 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,
}}
/>
</div>
)
})}
</div>
)
}
export default CanvasPreviewGrid
@@ -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<FrontendPreviewPlayerProps> = ({
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<FrontendPreviewPlayerProps> = ({
// ── 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<FrontendPreviewPlayerProps> = ({
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<FrontendPreviewPlayerProps> = ({
? { 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<FrontendPreviewPlayerProps> = ({
: undefined,
}}
>
{titleSettings.title.split(/[/]/).map((part, i) => (
{(effectiveTitle || "").split(/[/]/).map((part, i) => (
<span key={i}>
{i > 0 && <br />}
{part}
@@ -1,9 +1,9 @@
/**
* GeneratePage 步骤底部操作按钮
* GeneratePage 步骤底部操作按钮Issue #1677 修正:固定 6 步)
*
* 步骤 1~3:上一步 / 下一步
* 步骤 4标题+预览):上一步 / 确认生成视频(点击后直接创建最终渲染任务,成功后跳步骤5
* 步骤 5(确认生成):上一步 / 下一步(渲染中禁用,渲染完成后可进入封面)
* 步骤 4选择标题):「✨ 确认生成视频 / 确认生成 N 个视频」→ 创建正式生成任务,成功后跳步骤5
* 步骤 5(确认生成):渲染进度页,全部完成后「下一步:选择封面」;仅上一步
* 步骤 6(选择封面):仅上一步
*/
import React from "react"
@@ -12,14 +12,16 @@ export interface GenerateStepActionsProps {
currentStep: number
onPrev: () => void
onNext: () => void
/** 步骤4:确认生成视频(校验 + 创建渲染任务 + 成功后进入步骤5 */
/** 步骤4:确认生成视频(校验 + 创建渲染任务) */
onConfirmGenerate: () => void | Promise<void>
generating: boolean
generated: boolean
generateError: string | null
/** 批量模式下勾选的视频数量(N=1 时为1) */
selectedCount?: number
}
export const GenerateStepActions: React.FC<GenerateStepActionsProps> = ({
const GenerateStepActions: React.FC<GenerateStepActionsProps> = ({
currentStep,
onPrev,
onNext,
@@ -27,9 +29,10 @@ export const GenerateStepActions: React.FC<GenerateStepActionsProps> = ({
generating,
generated,
generateError,
selectedCount = 1,
}) => {
const renderPrimaryButton = () => {
/* 步骤 1~3:上一步 / 下一步(必填校验由 useStepNavigation.goNext 统一处理) */
/* 步骤 1~3:上一步 / 下一步 */
if (currentStep < 4) {
return (
<button className="xx-btn xx-btn-primary" onClick={onNext}>
@@ -38,50 +41,46 @@ export const GenerateStepActions: React.FC<GenerateStepActionsProps> = ({
)
}
/* 步骤 4:确认生成视频(触发按钮在标题页) */
/* 步骤 4选择标题 — 确认生成 */
if (currentStep === 4) {
if (generating) {
return (
<button className="xx-btn xx-btn-primary" disabled>
</button>
)
}
if (generateError) {
return (
<button className="xx-btn xx-btn-primary" onClick={onConfirmGenerate}>
🔄
</button>
)
}
if (generated) {
return (
<button className="xx-btn xx-btn-primary" onClick={onNext}>
🔄
</button>
)
}
return (
<button className="xx-btn xx-btn-primary" onClick={onConfirmGenerate}>
{selectedCount > 1 ? `✨ 确认生成 ${selectedCount} 个视频` : "✨ 确认生成视频"}
</button>
)
}
/* 步骤 5渲染中禁用,完成后下一步进封面 */
/* 步骤 5确认生成进度页 — 全部完成后下一步进封面 */
if (currentStep === 5) {
if (generated) {
return (
<button className="xx-btn xx-btn-primary" onClick={onNext}>
</button>
)
}
return (
<button
className="xx-btn xx-btn-primary"
onClick={onNext}
disabled={generating || !generated}
>
{generating ? "视频生成中…" : "下一步 →"}
<button className="xx-btn xx-btn-primary" disabled>
</button>
)
}
/* 步骤 6(最后一步):无主按钮 */
/* 步骤 6封面,最后一步):无主按钮 */
return null
}
@@ -1,20 +1,20 @@
/**
* GeneratePage 步骤内容渲染
* 步骤顺序(6步):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6)
* 步骤顺序(6步Issue #1677 修正):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6)
* 步骤4预览(Canvas 网格)与步骤5进度(批量渲染网格)由 GeneratePage 直接渲染在左侧大区域。
*/
import React from "react"
import type { EditingTemplate } from "@/api/editing-planner"
import type { EditPlanClip } from "@/api/template-editor"
import type { PresetVoiceItem } from "@/api/voices"
import type { VoiceClone } from "@/api/voice-clone"
import type { CoverConfig } from "../types/cover"
import type { TitleSettings } from "../types"
import Step1TemplateSelect from "../components/Step1TemplateSelect"
import Step2MaterialSelect from "../components/Step2MaterialSelect"
import Step3VoiceSelect from "../components/Step5VoiceSelect"
import Step3VoiceWithMode from "./Step3VoiceWithMode"
import Step4TitleSettings from "../components/Step4TitleSettings"
import Step5ConfirmGenerate from "../components/Step7ConfirmGenerate"
import Step6CoverSettings from "../components/Step6CoverSettings"
import BatchGenerationGrid from "./BatchGenerationGrid"
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
import type { GeneratedVideo } from "@/api/template-editor"
export interface GenerateStepContentProps {
@@ -50,15 +50,6 @@ export interface GenerateStepContentProps {
selectedVoice: string
onSelectedVoiceChange: (id: string) => void
onServerClipsChange: (clips: EditPlanClip[]) => void
voiceMode: "preset" | "custom" | "clone"
onVoiceModeChange: (mode: "preset" | "custom" | "clone") => void
selectedClonedVoice: string
onSelectedClonedVoiceChange: (id: string) => void
clonedVoices: VoiceClone[]
addClone: (voice: VoiceClone) => void
hasProcessing: boolean
cloneModalOpen: boolean
onCloneModalOpenChange: (open: boolean) => void
/* 生成 */
generating: boolean
generated: boolean
@@ -66,13 +57,26 @@ export interface GenerateStepContentProps {
progress: number
generatedVideos: GeneratedVideo[]
onRetry: () => void
onRetryBatchTask: (taskId: string) => void
onDismissError: () => void
/* 其他 */
presetVoices: PresetVoiceItem[]
/** 批量:每个正式生成任务的独立状态(步骤5进度网格) */
batchTasks: BatchTaskState[]
/** BGM 开关 */
bgm: boolean
/** BGM 配置(来自模板) */
bgmConfig?: { enabled: boolean; music_id?: string }
/* ── 批量生成(#1677)── */
previewCount: number
previewTitles: string[]
onPreviewTitlesChange: (titles: string[]) => void
voiceModePerVideo: boolean
onVoiceModePerVideoChange: (v: boolean) => void
voiceLibraryIds: string[]
onVoiceLibraryIdsChange: (ids: string[]) => void
previewCovers: string[]
onPreviewCoversChange: (urls: string[]) => void
/** 批量模式勾选的变体索引(封面卡片按勾选顺序展示) */
selectedVariantIds?: number[]
}
export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) => {
@@ -104,17 +108,24 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
selectedVoice,
onSelectedVoiceChange,
onServerClipsChange,
voiceMode,
selectedClonedVoice,
clonedVoices,
generating,
generated,
generateError,
progress,
generatedVideos,
onRetry,
onDismissError,
presetVoices,
generatedVideos,
batchTasks,
onRetryBatchTask,
previewCount,
previewTitles,
onPreviewTitlesChange,
voiceModePerVideo,
onVoiceModePerVideoChange,
voiceLibraryIds,
onVoiceLibraryIdsChange,
previewCovers,
onPreviewCoversChange,
selectedVariantIds,
} = props
/* 当前模板的 segments,传给 Step2 构建 clips */
@@ -146,9 +157,14 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
)
case 3:
return (
<Step3VoiceSelect
<Step3VoiceWithMode
previewCount={previewCount}
selectedVoice={selectedVoice}
onSelectedVoiceChange={onSelectedVoiceChange}
voiceModePerVideo={voiceModePerVideo}
onVoiceModePerVideoChange={onVoiceModePerVideoChange}
voiceLibraryIds={voiceLibraryIds}
onVoiceLibraryIdsChange={onVoiceLibraryIdsChange}
/>
)
case 4:
@@ -167,31 +183,66 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
onApplyPreset={onApplyPreset}
activePreset={activePreset}
titlePresets={titlePresets}
previewCount={previewCount}
previewTitles={previewTitles}
onPreviewTitlesChange={onPreviewTitlesChange}
/>
)
case 5:
/* 确认生成页:批量=逐任务进度网格;单视频=进度状态卡(成片播放器在左侧大区域) */
if (previewCount > 1) {
return (
<BatchGenerationGrid
tasks={batchTasks}
titles={previewTitles}
onRetryTask={onRetryBatchTask}
/>
)
}
/* 单视频:渲染进度 / 失败重试 / 完成提示(成片播放器在右侧栏) */
return (
<Step5ConfirmGenerate
templates={userTemplates}
selectedTemplate={selectedTemplate}
materialMode={materialMode}
selectedMaterials={selectedMaterials}
smartSelectedIds={smartSelectedIds}
title={titleSettings.title}
voiceMode={voiceMode}
selectedVoice={selectedVoice}
selectedClonedVoice={selectedClonedVoice}
presetVoices={presetVoices}
clonedVoices={clonedVoices}
coverSettings={coverSettings}
generating={generating}
generated={generated}
generateError={generateError}
progress={progress}
generatedVideos={generatedVideos}
onRetry={onRetry}
onDismissError={onDismissError}
/>
<div className="xx-form-section">
<h3>🎬 </h3>
{generating && (
<div className="xx-gen-progress-card">
<div className="xx-gen-progress-header">
<div className="xx-gen-progress-info">
<div className="xx-gen-progress-phase">
{Math.round(progress)}%
</div>
<div className="xx-gen-progress-sub">
</div>
</div>
</div>
<div className="xx-gen-progress-bar">
<div
className="xx-gen-progress-bar-fill"
style={{ width: `${Math.min(Math.round(progress), 100)}%` }}
/>
</div>
</div>
)}
{generateError && !generating && (
<div className="xx-gen-error-card">
<div className="xx-gen-error-info">
<div className="xx-gen-error-title"></div>
<div className="xx-gen-error-msg">{generateError}</div>
</div>
<button type="button" className="xx-btn xx-btn-primary xx-btn-sm" onClick={onRetry}>
🔄
</button>
</div>
)}
{generated && !generating && (
<div className="xx-gen-success-card">
<div className="xx-gen-success-info">
<div className="xx-gen-success-title"> </div>
<div className="xx-gen-success-sub"></div>
</div>
</div>
)}
</div>
)
case 6:
return (
@@ -201,6 +252,11 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
selectedTemplate={selectedTemplate}
titleSettings={titleSettings}
generatedVideos={generatedVideos}
previewCount={previewCount}
previewTitles={previewTitles}
previewCovers={previewCovers}
onPreviewCoversChange={onPreviewCoversChange}
selectedVariantIndexes={selectedVariantIds}
/>
)
default:
@@ -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<PreviewCountModalProps> = ({
open,
defaultCount = 1,
onConfirm,
onCancel,
}) => {
const [count, setCount] = useState(defaultCount)
const inputRef = useRef<HTMLInputElement>(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 (
<div className="xx-modal-mask" onClick={onCancel}>
<div className="xx-modal-box xx-count-modal" onClick={(e) => e.stopPropagation()}>
<h3 style={{ margin: "0 0 8px", fontSize: 18 }}></h3>
<p style={{ margin: "0 0 20px", fontSize: 13, color: "var(--text-secondary, #666)" }}>
AI
</p>
<div className="xx-count-selector">
<button
type="button"
className="xx-count-btn"
onClick={() => setCount((c) => clamp(c - 1))}
disabled={count <= 1}
aria-label="减少"
>
</button>
<input
ref={inputRef}
type="number"
min={1}
max={MAX_PREVIEW_COUNT}
value={count}
onChange={(e) => setCount(clamp(parseInt(e.target.value, 10) || 1))}
onKeyDown={handleKeyDown}
className="xx-count-input"
/>
<button
type="button"
className="xx-count-btn"
onClick={() => setCount((c) => clamp(c + 1))}
disabled={count >= MAX_PREVIEW_COUNT}
aria-label="增加"
>
+
</button>
</div>
<div className="xx-count-quick">
{[1, 3, 5, 10].map((n) => (
<button
key={n}
type="button"
className={`xx-count-chip ${count === n ? "active" : ""}`}
onClick={() => setCount(n)}
>
{n}
</button>
))}
</div>
<div className="xx-count-actions">
<button type="button" className="xx-btn xx-btn-ghost" onClick={onCancel}>
</button>
<button type="button" className="xx-btn xx-btn-primary" onClick={handleConfirm}>
{count === 1 ? "生成 1 个视频" : `生成 ${count} 个视频`}
</button>
</div>
<p
style={{
margin: "12px 0 0",
fontSize: 12,
color: "var(--text-tertiary, #999)",
textAlign: "center",
}}
>
= 1
</p>
</div>
</div>
)
}
export default PreviewCountModal
@@ -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<Step3VoiceWithModeProps> = ({
previewCount,
selectedVoice,
onSelectedVoiceChange,
voiceModePerVideo,
onVoiceModePerVideoChange,
voiceLibraryIds,
onVoiceLibraryIdsChange,
}) => {
const isBatch = previewCount > 1
if (!isBatch) {
return (
<Step3VoiceSelect
selectedVoice={selectedVoice}
onSelectedVoiceChange={onSelectedVoiceChange}
/>
)
}
return (
<div className="xx-form-section">
{/* 共用/独立切换 */}
<div className="xx-title-ai-toggle" style={{ marginBottom: 16 }}>
<div>
<div style={{ fontWeight: 600, fontSize: 15 }}>
🎙 {voiceModePerVideo ? "每个视频独立配音" : "所有视频共用配音"}
</div>
<div style={{ fontSize: 12, color: "var(--text-tertiary, #999)", marginTop: 2 }}>
{voiceModePerVideo
? `${previewCount} 个视频分别选择不同配音`
: "所有视频使用同一个配音(默认)"}
</div>
</div>
<div
className={`xx-switch ${voiceModePerVideo ? "active" : ""}`}
onClick={() => onVoiceModePerVideoChange(!voiceModePerVideo)}
role="switch"
aria-checked={voiceModePerVideo}
tabIndex={0}
onKeyDown={(e) => {
if (e.key === "Enter" || e.key === " ") {
e.preventDefault()
onVoiceModePerVideoChange(!voiceModePerVideo)
}
}}
>
<div className="xx-switch-knob" />
</div>
</div>
{!voiceModePerVideo ? (
<Step3VoiceSelect
heading="🎙️ 共用配音"
description={`所有 ${previewCount} 个视频使用同一个配音,点击卡片可预览播放`}
selectedVoice={selectedVoice}
onSelectedVoiceChange={onSelectedVoiceChange}
/>
) : (
<div className="xx-per-voice-list">
{Array.from({ length: previewCount }, (_, i) => (
<Step3VoiceSelect
key={i}
heading={`🎙️ 视频 ${i + 1} 的配音`}
description="为这个视频单独选择配音"
compact
selectedVoice={voiceLibraryIds[i] || ""}
onSelectedVoiceChange={(id) => {
const next = [...voiceLibraryIds]
next[i] = id
onVoiceLibraryIdsChange(next)
}}
/>
))}
</div>
)}
</div>
)
}
export default Step3VoiceWithMode
@@ -1,17 +1,22 @@
/**
* Step 4 选择标题(合并原 Step4 标题输入 + Step5 标题样式面板
* Step 4 选择标题(Issue #1677 批量生成
*
* 左侧:标题文字输入 + AI生成标题 + 样式设置(位置/字号/字体/颜色/样式/预设)
* 右侧FrontendPreviewPlayer 实时预览(由 GeneratePage 统一渲染)
* 布局(由 GeneratePage 编排):左侧大区域实时预览(单=大播放器,批量=Canvas 网格),
* 右侧边栏标题设置。本组件渲染在右侧边栏:
* - 单视频:AI 标题生成器 + AutoComplete 标题库(与旧版完全一致,零回归)
* - 批量:N 个独立标题输入框(AutoComplete 支持标题库选择)+ 批量 AI 生成
* (一次生成 N 个标题,分别填入各变体,可单独换一个)
* - 标题样式(字体/颜色/位置/大小/粗斜描边/预设):全局统一
*/
import React from "react"
import { AutoComplete } from "antd"
import { PlayCircleOutlined } from "@ant-design/icons"
import React, { useMemo, useState } from "react"
import { 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
@@ -29,6 +34,42 @@ interface Step4TitleSettingsProps {
onApplyPreset: (presetKey: string) => void
activePreset: string | null
titlePresets: { key: string; label: string; previewStyle: React.CSSProperties }[]
/* ── 批量生成(#1677)── */
/** 生成数量 */
previewCount?: number
/** 每个变体的标题文字(长度=previewCount */
previewTitles?: string[]
onPreviewTitlesChange?: (titles: string[]) => void
}
/** 从本地 AI 标题模板池按主题词生成 N 个不同标题(与单视频 AI 生成同源) */
function buildBatchAiTitles(topic: string, count: number): string[] {
const styles: Array<"catchy" | "emotional" | "informative"> = [
"catchy",
"emotional",
"informative",
]
const pool: string[] = []
styles.forEach((style) => {
const templates = AI_TITLE_TEMPLATES[style] || []
templates.forEach((tpl) => pool.push(tpl.replace(/\{topic\}/g, topic)))
})
// 洗牌后取前 count 个;不足则轮转补齐
const shuffled = [...pool].sort(() => Math.random() - 0.5)
const out: string[] = []
for (let i = 0; i < count; i++) {
out.push(shuffled[i % shuffled.length] || "")
}
return out
}
function extractTopic(text: string): string {
const keywords = text
.replace(/[,。!?、,.!?]/g, " ")
.split(/\s+/)
.filter(Boolean)
if (keywords.length === 0) return "这个话题"
return keywords.slice(0, 3).join("")
}
const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
@@ -44,122 +85,208 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
onApplyPreset,
activePreset,
titlePresets,
previewCount = 1,
previewTitles,
onPreviewTitlesChange,
} = props
const isBatch = previewCount > 1
const [batchAiLoading, setBatchAiLoading] = useState(false)
const [batchAiTopic, setBatchAiTopic] = useState("")
/** 更新单个变体标题;变体0同步写回 titleSettings.title(全局样式面板/草稿/TTS 链路依赖) */
const updateVariantTitle = (index: number, val: string) => {
if (!previewTitles || !onPreviewTitlesChange) return
const next = [...previewTitles]
next[index] = val
onPreviewTitlesChange(next)
if (index === 0) {
t.updateTitle(val)
}
}
/** 批量 AI 生成:按主题词生成标题,分别填入 N 个变体 */
const handleBatchAiGenerate = async (onlyEmpty = false) => {
if (!onPreviewTitlesChange || !previewTitles) return
const topic = (batchAiTopic || t.aiTitleInput || "").trim()
if (!topic) {
message.warning("请先输入主题词,例如:萌宠日常、旅行vlog")
return
}
setBatchAiLoading(true)
try {
// 与单视频一致:本地模板模拟 AI 生成(1200ms 体验延迟)
await new Promise((resolve) => setTimeout(resolve, 800))
const picked = buildBatchAiTitles(extractTopic(topic), previewCount)
const next = [...previewTitles]
for (let i = 0; i < previewCount; i++) {
if (onlyEmpty && next[i]?.trim()) continue
if (picked[i]) next[i] = picked[i]
}
onPreviewTitlesChange(next)
if (next[0]) t.updateTitle(next[0])
message.success(`已为 ${previewCount} 个视频生成标题,可单独修改`)
} finally {
setBatchAiLoading(false)
}
}
const titleOptions = useMemo(
() => t.userTitles.map((ut) => ({ label: ut.content, value: ut.content })),
[t.userTitles],
)
return (
<div className="xx-form-section">
<div className="xx-form-section xx-title-sidebar">
<h3>📝 </h3>
{/* AI 自动选择模式 */}
{t.titleSettings.aiAutoSelect && (
{!isBatch ? (
/* ── 单视频:原有 AI 标题 + 输入框(保持不变,零回归) ── */
<>
<div className="xx-title-ai-toggle">
<span className="xx-toggle-label">AI </span>
<div className="xx-switch active" onClick={t.toggleAiAutoSelect}>
<div className="xx-switch-knob" />
</div>
</div>
{/* 显示当前 AI 选中的标题(只读)+ 换一个按钮 */}
<div className="xx-form-field">
<label> AI </label>
<div
style={{
display: "flex",
alignItems: "center",
gap: 10,
padding: "8px 12px",
background: "var(--bg-secondary, rgba(0,0,0,0.04))",
borderRadius: 8,
fontSize: 14,
color: "var(--text-primary, #333)",
}}
>
<span style={{ flex: 1 }}>{t.titleSettings.title || "AI 将自动为你选择标题"}</span>
<button
type="button"
className="xx-btn xx-btn-primary"
style={{ flexShrink: 0, fontSize: 13, padding: "4px 12px" }}
onClick={t.autoGenerateTitle}
>
🔄
</button>
</div>
</div>
</>
)}
{/* 手动选择模式 */}
{!t.titleSettings.aiAutoSelect && (
<>
<AiTitleGenerator
inputValue={t.aiTitleInput}
onInputChange={t.setAiTitleInput}
generating={t.aiTitleGenerating}
onGenerate={t.handleGenerateAiTitles}
results={t.aiTitleResults}
hasGenerated={t.hasGeneratedTitles}
onSelect={t.handleSelectAiTitle}
selectedTitle={t.titleSettings.title}
onRefresh={t.handleRefreshAiTitles}
/>
<div className="xx-title-ai-toggle">
<span className="xx-toggle-label">AI </span>
<div className="xx-switch" onClick={t.toggleAiAutoSelect}>
<div className="xx-switch-knob" />
</div>
</div>
<div className="xx-form-field">
<label></label>
<AutoComplete
placeholder="输入或从标题库选择…"
allowClear
maxLength={50}
style={{ width: "100%" }}
value={t.titleSettings.title || undefined}
onChange={(val) => 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 ? (
<span style={{ color: "var(--text-tertiary)", fontSize: 13 }}>
{t.titleSettings.aiAutoSelect ? (
<>
<div className="xx-title-ai-toggle">
<span className="xx-toggle-label">AI </span>
<div className="xx-switch active" onClick={t.toggleAiAutoSelect}>
<div className="xx-switch-knob" />
</div>
</div>
<div className="xx-form-field">
<label> AI </label>
<div
style={{
display: "flex",
alignItems: "center",
gap: 10,
padding: "8px 12px",
background: "var(--bg-secondary, rgba(0,0,0,0.04))",
borderRadius: 8,
fontSize: 14,
color: "var(--text-primary, #333)",
}}
>
<span style={{ flex: 1 }}>
{(previewTitles?.[0] ?? t.titleSettings.title) || "AI 将自动为你选择标题"}
</span>
) : null
}
/>
</div>
<button
type="button"
className="xx-btn xx-btn-primary"
style={{ flexShrink: 0, fontSize: 13, padding: "4px 12px" }}
onClick={t.autoGenerateTitle}
>
🔄
</button>
</div>
</div>
</>
) : (
<>
<AiTitleGenerator
inputValue={t.aiTitleInput}
onInputChange={t.setAiTitleInput}
generating={t.aiTitleGenerating}
onGenerate={t.handleGenerateAiTitles}
results={t.aiTitleResults}
hasGenerated={t.hasGeneratedTitles}
onSelect={t.handleSelectAiTitle}
selectedTitle={t.titleSettings.title}
onRefresh={t.handleRefreshAiTitles}
/>
<div className="xx-title-ai-toggle">
<span className="xx-toggle-label">AI </span>
<div className="xx-switch" onClick={t.toggleAiAutoSelect}>
<div className="xx-switch-knob" />
</div>
</div>
<div className="xx-form-field">
<label></label>
<AutoComplete
placeholder="输入标题文字…"
allowClear
maxLength={50}
style={{ width: "100%" }}
value={(previewTitles?.[0] ?? t.titleSettings.title) || undefined}
onChange={(val) => {
t.updateTitle(val || "")
onPreviewTitlesChange?.([val || ""])
}}
options={titleOptions}
filterOption={(inputValue, option) => {
const title = (option?.label || option?.value || "") as string
return title.toLowerCase().includes((inputValue || "").toLowerCase())
}}
/>
</div>
</>
)}
</>
) : (
/* ── 批量:AI 批量生成 + N 个独立标题输入框(AutoComplete 支持标题库) ── */
<div className="xx-batch-titles">
<div
style={{
fontSize: 12,
color: "var(--text-secondary, #666)",
marginBottom: 10,
lineHeight: 1.6,
}}
>
//
</div>
{/* 批量 AI 标题 */}
<div className="xx-batch-ai-row">
<Input
placeholder="主题词,如:萌宠日常、旅行vlog"
value={batchAiTopic || t.aiTitleInput}
onChange={(e) => {
setBatchAiTopic(e.target.value)
t.setAiTitleInput(e.target.value)
}}
maxLength={30}
size="small"
style={{ flex: 1 }}
/>
<button
type="button"
className="xx-btn xx-btn-primary xx-btn-sm"
disabled={batchAiLoading}
onClick={() => handleBatchAiGenerate(false)}
>
{batchAiLoading ? <LoadingOutlined /> : "✨"} {previewCount}
</button>
<button
type="button"
className="xx-btn xx-btn-ghost xx-btn-sm"
disabled={batchAiLoading}
onClick={() => handleBatchAiGenerate(true)}
>
</button>
</div>
{Array.from({ length: previewCount }, (_, i) => (
<div className="xx-form-field" key={i}>
<label> {i + 1} </label>
<AutoComplete
placeholder={`视频 ${i + 1} 的标题…`}
maxLength={50}
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())
}}
/>
</div>
))}
</div>
)}
{/* 标题样式面板(原 Step5 */}
<div
style={{
display: "flex",
alignItems: "center",
gap: 8,
padding: "10px 14px",
background: "rgba(59, 130, 246, 0.08)",
borderRadius: 8,
marginTop: 16,
marginBottom: 12,
border: "1px solid rgba(59, 130, 246, 0.15)",
}}
>
<PlayCircleOutlined style={{ fontSize: 16, color: "#3b82f6" }} />
<span style={{ fontSize: 12, color: "var(--text-secondary, #666)" }}>
</span>
</div>
{/* 标题样式面板(全局共用 */}
<TitleStylePanel
settings={t.titleSettings}
onUpdatePosition={onUpdatePosition}
@@ -12,6 +12,12 @@ import type { AssetItem } from "@/api/assets"
interface Step5VoiceSelectProps {
selectedVoice: string
onSelectedVoiceChange: (id: string) => void
/** 卡片标题(独立配音模式下显示"视频 N 的配音"),默认"选择配音" */
heading?: string
/** 描述文案 */
description?: string
/** 是否使用紧凑卡片样式(独立配音模式下 N 个并排) */
compact?: boolean
}
/** 获取素材实际时长(优先顶层 durationfallback 到 metadata.duration */
@@ -43,6 +49,9 @@ const formatFileSize = (bytes?: number): string => {
const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
selectedVoice,
onSelectedVoiceChange,
heading = "🎙️ 选择配音",
description = "从配音库中选择已上传的素材,点击卡片可预览播放",
compact = false,
}) => {
const navigate = useNavigate()
const [playingId, setPlayingId] = useState<string | null>(null)
@@ -148,14 +157,14 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
return (
<div className="xx-form-section">
<h3>🎙 </h3>
<p style={{ color: "#666", marginBottom: 16, fontSize: 14 }}>
</p>
<h3>{heading}</h3>
<p style={{ color: "#666", marginBottom: 16, fontSize: 14 }}>{description}</p>
<div
style={{
display: "grid",
gridTemplateColumns: "repeat(auto-fill, minmax(220px, 1fr))",
gridTemplateColumns: compact
? "repeat(auto-fill, minmax(160px, 1fr))"
: "repeat(auto-fill, minmax(220px, 1fr))",
gap: 12,
}}
>
@@ -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<Step6CoverSettingsProps> = (props) => {
@@ -46,13 +62,173 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (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<HTMLInputElement>(null)
const uploadTargetRef = useRef<number>(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<HTMLInputElement>) => {
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 (
<div className="xx-form-section">
<h3>🖼 </h3>
<div
style={{
padding: "10px 14px",
background: "rgba(16, 185, 129, 0.08)",
borderRadius: 8,
marginBottom: 16,
border: "1px solid rgba(16, 185, 129, 0.15)",
fontSize: 13,
color: "var(--text-secondary, #666)",
}}
>
🎬 {completedVideos.length}
</div>
<div style={{ display: "flex", gap: 8, marginBottom: 16 }}>
<Button
buttonType="primary"
onClick={() => void batchCovers.generateAll()}
disabled={completedVideos.length === 0 || batchCovers.loadingIndex !== null}
>
</Button>
</div>
<div className="xx-cover-grid">
{cardIndexes.map((variantIndex, cardPos) => {
const url = batchCoversList[cardPos]
const isLoading = batchCovers.loadingIndex === cardPos
const isUploading = batchCovers.uploadingIndex === cardPos
const title = batchTitles[cardPos]
return (
<div className="xx-cover-card" key={variantIndex}>
<div className="xx-cover-card-title"> {variantIndex + 1}</div>
<div className="xx-cover-card-box">
{isLoading || isUploading ? (
<div className="xx-cover-card-loading">
<Spin indicator={<LoadingOutlined style={{ fontSize: 24 }} spin />} />
<span>{isLoading ? "AI 选帧中…" : "上传中…"}</span>
</div>
) : url ? (
<img
src={url}
alt={`视频${variantIndex + 1}封面`}
className="xx-cover-card-img"
/>
) : (
<div className="xx-cover-card-placeholder">
<span style={{ fontSize: 26 }}>🖼</span>
<span style={{ fontSize: 12 }}></span>
</div>
)}
<div className="xx-cover-card-ratio">9:16</div>
</div>
{title && (
<div
style={{
fontSize: 12,
color: "var(--text-secondary, #666)",
marginTop: 6,
overflow: "hidden",
textOverflow: "ellipsis",
whiteSpace: "nowrap",
}}
title={title}
>
{title}
</div>
)}
<div style={{ display: "flex", gap: 6, marginTop: 8 }}>
<button
type="button"
className="xx-btn xx-btn-primary xx-btn-sm"
style={{ flex: 1, fontSize: 12, padding: "4px 8px" }}
onClick={() => void batchCovers.generateOne(cardPos)}
disabled={isLoading || isUploading}
>
{url ? "🔄 重新生成" : "✨ 自动生成"}
</button>
<button
type="button"
className="xx-btn xx-btn-ghost xx-btn-sm"
style={{ flex: 1, fontSize: 12, padding: "4px 8px" }}
onClick={() => handleUploadClick(variantIndex)}
disabled={isLoading || isUploading}
>
📤
</button>
</div>
</div>
)
})}
</div>
<input
ref={uploadInputRef}
type="file"
accept="image/*"
style={{ display: "none" }}
onChange={handleFileChange}
/>
</div>
)
}
/* ── 单视频:原有流程保持不变 ── */
return (
<div className="xx-form-section">
<h3>🖼 </h3>
@@ -75,7 +251,7 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
)}
<div className="xx-cover-actions">
<Button buttonType="primary" onClick={handleAutoGenerate} disabled={!finalVideo}>
<Button buttonType="primary" onClick={generateAutoCover} disabled={!finalVideo}>
</Button>
<Button buttonType="ghost" onClick={() => setShowCoverSettings(true)}>
@@ -1,79 +0,0 @@
/**
* Step 7 确认生成组件
*/
import React from "react"
import type { EditingTemplate } from "@/api/editing-planner"
import type { GeneratedVideo } from "@/api/template-editor"
import type { CoverConfig } from "../types/cover"
import type { VoiceClone } from "@/api/voice-clone"
import type { PresetVoiceItem } from "@/api/voices"
import { useStep7Generate } from "../hooks/useStep7Generate"
import SummaryCard from "./step7-confirm/SummaryCard"
import GenerationStatus from "./step7-confirm/GenerationStatus"
interface Step7ConfirmGenerateProps {
templates: EditingTemplate[]
selectedTemplate: string
materialMode: "manual" | "auto"
selectedMaterials: string[]
smartSelectedIds: string[]
title: string
voiceMode: "preset" | "custom" | "clone"
selectedVoice: string
selectedClonedVoice: string
presetVoices: PresetVoiceItem[]
clonedVoices: VoiceClone[]
coverSettings: CoverConfig
generating: boolean
generated: boolean
generateError: string | null
progress: number
generatedVideos: GeneratedVideo[]
onRetry: () => void
onDismissError: () => void
}
const Step7ConfirmGenerate: React.FC<Step7ConfirmGenerateProps> = (props) => {
const {
templateName,
materialSummary,
title,
voiceName,
coverSummary,
generating,
generated,
generateError,
progress,
generatedVideos,
getGenerationPhase,
handleScrollToPreview,
} = useStep7Generate(props)
const { onRetry, onDismissError } = props
return (
<div className="xx-form-section">
<h3> </h3>
<SummaryCard
templateName={templateName}
materialSummary={materialSummary}
title={title}
voiceName={voiceName}
coverSummary={coverSummary}
/>
<GenerationStatus
generating={generating}
generated={generated}
generateError={generateError}
progress={progress}
generatedVideos={generatedVideos}
getGenerationPhase={getGenerationPhase}
onScrollToPreview={handleScrollToPreview}
onRetry={onRetry}
onDismissError={onDismissError}
/>
</div>
)
}
export default Step7ConfirmGenerate
@@ -1,44 +0,0 @@
import React from "react"
interface SummaryCardProps {
templateName: string
materialSummary: string
title: string
voiceName: string
coverSummary: string
}
const SummaryCard: React.FC<SummaryCardProps> = ({
templateName,
materialSummary,
title,
voiceName,
coverSummary,
}) => {
return (
<div className="xx-summary-card">
<div className="xx-summary-row">
<span className="xx-summary-label"></span>
<span className="xx-summary-value">{templateName}</span>
</div>
<div className="xx-summary-row">
<span className="xx-summary-label"></span>
<span className="xx-summary-value">{materialSummary}</span>
</div>
<div className="xx-summary-row">
<span className="xx-summary-label"></span>
<span className="xx-summary-value">{title || "未选择"}</span>
</div>
<div className="xx-summary-row">
<span className="xx-summary-label"></span>
<span className="xx-summary-value">{voiceName}</span>
</div>
<div className="xx-summary-row">
<span className="xx-summary-label"></span>
<span className="xx-summary-value">{coverSummary}</span>
</div>
</div>
)
}
export default SummaryCard
+4
View File
@@ -37,6 +37,10 @@ export const STEPS = [
{ key: 6, label: "选择封面" },
]
/* ── 批量生成限制 ── */
export const MAX_PREVIEW_COUNT = 10
export const MIN_PREVIEW_COUNT = 1
/* ── 标题位置选项 ── */
export const POSITION_OPTIONS = [
{ value: "top", label: "顶部" },
+542
View File
@@ -2900,3 +2900,545 @@
color: rgba(255, 255, 255, 0.85);
white-space: nowrap;
}
/* ================================================================
Issue #1677 多视频批量生成
================================================================ */
/* ── Step4 布局对调:左侧预览大区域,右侧标题边栏 ── */
.xx-generate-layout.step4-layout {
grid-template-columns: 1fr 380px;
align-items: start;
}
.xx-generate-preview-col {
min-width: 0;
position: sticky;
top: 16px;
}
.xx-generate-preview-col .xx-form-section {
margin: 0;
}
.xx-title-sidebar {
max-height: calc(100vh - 140px);
overflow-y: auto;
}
/* ── 数量选择弹窗 ── */
.xx-modal-mask {
position: fixed;
inset: 0;
background: rgba(0, 0, 0, 0.45);
display: flex;
align-items: center;
justify-content: center;
z-index: 1000;
}
.xx-modal-box {
background: var(--bg-primary, #fff);
border-radius: 16px;
padding: 28px;
width: 420px;
max-width: calc(100vw - 32px);
box-shadow: 0 12px 48px rgba(0, 0, 0, 0.18);
}
.xx-count-selector {
display: flex;
align-items: center;
justify-content: center;
gap: 16px;
margin: 8px 0 16px;
}
.xx-count-btn {
width: 44px;
height: 44px;
border-radius: 50%;
border: 1px solid var(--border-primary, #d9d9d9);
background: var(--bg-secondary, #f5f5f5);
font-size: 22px;
line-height: 1;
cursor: pointer;
color: var(--text-primary, #333);
transition: all 0.15s;
}
.xx-count-btn:hover:not(:disabled) {
border-color: #1677ff;
color: #1677ff;
}
.xx-count-btn:disabled {
opacity: 0.4;
cursor: not-allowed;
}
.xx-count-input {
width: 88px;
height: 52px;
text-align: center;
font-size: 26px;
font-weight: 700;
border: 2px solid var(--border-primary, #d9d9d9);
border-radius: 12px;
color: var(--text-primary, #333);
background: var(--bg-primary, #fff);
}
.xx-count-input:focus {
outline: none;
border-color: #1677ff;
}
/* 隐藏 number input 上下箭头 */
.xx-count-input::-webkit-outer-spin-button,
.xx-count-input::-webkit-inner-spin-button {
-webkit-appearance: none;
margin: 0;
}
.xx-count-input {
-moz-appearance: textfield;
appearance: textfield;
}
.xx-count-quick {
display: flex;
gap: 8px;
justify-content: center;
margin-bottom: 20px;
}
.xx-count-chip {
padding: 6px 16px;
border-radius: 999px;
border: 1px solid var(--border-primary, #d9d9d9);
background: var(--bg-primary, #fff);
font-size: 13px;
cursor: pointer;
color: var(--text-secondary, #666);
transition: all 0.15s;
}
.xx-count-chip:hover {
border-color: #1677ff;
color: #1677ff;
}
.xx-count-chip.active {
background: #1677ff;
border-color: #1677ff;
color: #fff;
}
.xx-count-actions {
display: flex;
gap: 12px;
}
.xx-count-actions .xx-btn {
flex: 1;
}
/* ── 批量预览网格 ── */
.xx-variant-grid {
display: grid;
grid-template-columns: repeat(auto-fill, minmax(240px, 1fr));
gap: 16px;
}
.xx-variant-card {
position: relative;
border: 2px solid var(--border-primary, #e8e8e8);
border-radius: 12px;
padding: 10px;
background: var(--bg-primary, #fff);
cursor: pointer;
transition: all 0.18s;
}
.xx-variant-card:hover {
border-color: #91caff;
}
.xx-variant-card.selected {
border-color: #1677ff;
box-shadow: 0 0 0 3px rgba(22, 119, 255, 0.12);
}
.xx-variant-card.failed {
border-color: #ffccc7;
cursor: default;
}
.xx-variant-check {
position: absolute;
top: 14px;
left: 14px;
z-index: 3;
width: 26px;
height: 26px;
border-radius: 50%;
border: 2px solid #fff;
background: rgba(0, 0, 0, 0.35);
color: #fff;
display: flex;
align-items: center;
justify-content: center;
font-size: 14px;
font-weight: 700;
}
.xx-variant-check.checked {
background: #1677ff;
border-color: #1677ff;
}
.xx-variant-index {
font-size: 13px;
font-weight: 600;
color: var(--text-secondary, #666);
margin-bottom: 8px;
}
.xx-variant-video-wrap {
position: relative;
border-radius: 8px;
overflow: hidden;
background: #000;
aspect-ratio: 9 / 16;
max-height: 420px;
display: flex;
align-items: center;
justify-content: center;
}
.xx-variant-loading,
.xx-variant-failed {
display: flex;
flex-direction: column;
align-items: center;
gap: 10px;
color: var(--text-secondary, #999);
font-size: 12px;
padding: 16px;
text-align: center;
}
.xx-variant-progress {
width: 120px;
height: 4px;
border-radius: 2px;
background: rgba(255, 255, 255, 0.25);
overflow: hidden;
}
.xx-variant-progress-bar {
height: 100%;
background: #1677ff;
border-radius: 2px;
transition: width 0.4s;
}
.xx-variant-progress-text {
color: rgba(255, 255, 255, 0.85);
font-size: 12px;
}
.xx-variant-title-overlay {
position: absolute;
left: 8%;
right: 8%;
text-align: center;
font-weight: 700;
line-height: 1.3;
pointer-events: none;
text-shadow: 0 1px 3px rgba(0, 0, 0, 0.7);
word-break: break-all;
}
.xx-variant-title-overlay.pos-top {
top: 8%;
}
.xx-variant-title-overlay.pos-center {
top: 50%;
transform: translateY(-50%);
}
.xx-variant-title-overlay.pos-bottom,
.xx-variant-title-overlay.pos-custom {
bottom: 10%;
}
.xx-variant-footer {
min-height: 22px;
margin-top: 8px;
font-size: 12px;
}
.xx-variant-ready-tag {
color: #52c41a;
display: inline-flex;
align-items: center;
gap: 4px;
}
.xx-variant-skip-tag {
color: var(--text-tertiary, #999);
}
/* ── 批量配音列表 ── */
.xx-per-voice-list {
display: flex;
flex-direction: column;
gap: 16px;
}
.xx-per-voice-list .xx-form-section {
margin: 0;
}
/* ── 批量封面网格 ── */
.xx-cover-grid {
display: grid;
grid-template-columns: repeat(auto-fill, minmax(180px, 1fr));
gap: 16px;
}
.xx-cover-card {
border: 1px solid var(--border-primary, #e8e8e8);
border-radius: 12px;
padding: 10px;
background: var(--bg-primary, #fff);
}
.xx-cover-card-title {
font-size: 13px;
font-weight: 600;
color: var(--text-secondary, #666);
margin-bottom: 8px;
}
.xx-cover-card-box {
position: relative;
border-radius: 8px;
overflow: hidden;
background: #000;
aspect-ratio: 9 / 16;
max-height: 300px;
}
.xx-cover-card-img {
width: 100%;
height: 100%;
object-fit: cover;
display: block;
}
.xx-cover-card-placeholder,
.xx-cover-card-loading {
position: absolute;
inset: 0;
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
gap: 8px;
color: var(--text-tertiary, #999);
background: var(--bg-secondary, #f7f7f7);
font-size: 13px;
}
.xx-cover-card-ratio {
position: absolute;
right: 6px;
bottom: 6px;
background: rgba(0, 0, 0, 0.55);
color: #fff;
font-size: 10px;
padding: 1px 6px;
border-radius: 4px;
}
/* ── 响应式:窄屏 Step4 回退单列 ── */
@media (max-width: 960px) {
.xx-generate-layout.step4-layout {
grid-template-columns: 1fr;
}
.xx-generate-preview-col {
position: static;
}
.xx-title-sidebar {
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;
}
}
@@ -29,6 +29,19 @@ export interface UseGenerateVideoProps {
}
/** 生成成功后的回调(用于清除持久化的 previewTaskId 等状态) */
onGenerationSuccess?: () => void
/* ── 批量生成(#1677)── */
/** 生成数量(1=单条旧逻辑,>1=批量) */
previewCount?: number
/** 每个变体的标题文字(长度=count 时各自独立) */
variantTitles?: string[]
/** 独立配音模式下每个变体的配音ID(空数组=共用 selectedVoice */
variantVoiceLibraryIds?: string[]
/** 是否独立配音 */
voiceModePerVideo?: boolean
/** 每个变体的封面URL(空数组=回退 coverSettings */
variantCoverUrls?: string[]
/** 勾选要生成的变体索引(批量模式) */
selectedVariantIndexes?: number[]
}
/** 生成阶段 */
@@ -44,7 +57,7 @@ export interface UseGenerateVideoResult {
generated: boolean
generateError: string | null
generatedVideos: GeneratedVideo[]
generate: () => Promise<void>
generate: () => Promise<boolean>
retry: () => void
dismissError: () => void
download: () => Promise<void>
@@ -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<BatchTaskState>) => void
}
/** 最大连续错误次数(仅对可重试错误),超过后终止轮询 */
@@ -17,31 +31,30 @@ const MAX_RETRYABLE_ERRORS = 10
const MAX_RESULTS_RETRIES = 3
/**
* 生成状态轮询 Hookv2改用 /generation/tasks/{task_id}
* 生成状态轮询 Hookv4批量任务独立状态 + 单任务重试
*
* 旧版轮询 GET /templates/{id}/editor/generation-status 依赖 plan 维度状态,
* 在编辑流程数据链路断裂时拿不到 task_id。新版直接使用 POST /generation/tasks
* 返回的 task_id 轮询任务详情,不再依赖 plan。
*
* 错误处理:
* - 4xx(尤其 404)视为不可恢复,立即 onFailed,不再重试
* - 5xx / 网络错误重试,最多连续 MAX_RETRYABLE_ERRORS 次
* - 任务完成后获取结果失败会重试 MAX_RESULTS_RETRIES 次,仍失败则 onFailed
* startPolling(taskId) 轮询单个任务;
* startPollingBatch(tasks) 并行轮询 N 个任务:
* - 每个任务独立进度/状态/失败,通过 onBatchTaskUpdate 实时回传
* - 全部成功才 onComplete(聚合视频按变体顺序);任一失败不影响其他任务继续
* - retryTask(taskId) 单独重试失败任务(重新轮询,后端任务仍在跑则直接接续)
*/
export const useGenerationPolling = ({
export function useGenerationPolling({
onProgress,
onComplete,
onFailed,
}: UseGenerationPollingOptions) => {
const progressTimer = useRef<ReturnType<typeof setTimeout>>()
onBatchTaskUpdate,
}: UseGenerationPollingOptions) {
const progressTimer = useRef<ReturnType<typeof setTimeout>[]>([])
const cancelledRef = useRef(false)
/** 批量任务上下文:taskId → 变体序号 */
const batchContextRef = useRef<Map<string, number>>(new Map())
const [, forceTick] = useState(0)
const clearTimer = useCallback(() => {
cancelledRef.current = true
if (progressTimer.current) {
clearTimeout(progressTimer.current)
progressTimer.current = undefined
}
progressTimer.current.forEach((t) => clearTimeout(t))
progressTimer.current = []
}, [])
/** 任务完成后拉取结果列表,带重试 */
@@ -62,85 +75,237 @@ export const useGenerationPolling = ({
[],
)
const extractErrorMessage = (pollErr: unknown, status: number): string => {
const msg =
(axios.isAxiosError(pollErr) &&
(pollErr.response?.data as { detail?: string; message?: string } | undefined)?.detail) ||
(axios.isAxiosError(pollErr) &&
(pollErr.response?.data as { detail?: string; message?: string } | undefined)?.message) ||
`查询任务失败 (${status})`
return safeExtractError(msg)
}
/**
* 轮询单个任务。
* - isBatch=true:状态变化通过 onBatchTaskUpdate 回传,不触发整体 onProgress/onComplete
* - resolve(videos) 成功;reject(Error) 失败
*/
const pollSingleTask = useCallback(
(
taskId: string,
runId: number,
callbacks?: {
onTaskProgress?: (pct: number) => void
onTaskCompleted?: (videos: unknown[]) => void
onTaskFailed?: (msg: string) => void
},
): Promise<unknown[]> => {
return new Promise((resolve, reject) => {
let consecutiveErrors = 0
let done = false
const poll = async () => {
if (cancelledRef.current || done) return
try {
const task = await getGenerationTask(taskId)
if (cancelledRef.current || done) return
consecutiveErrors = 0
if (task.status === "completed") {
done = true
const videos = await fetchResultsWithRetry(taskId)
if (cancelledRef.current) return
if (videos === null) {
const msg = "视频已生成,但获取结果列表失败,请稍后在任务列表查看"
callbacks?.onTaskFailed?.(msg)
reject(new Error(msg))
return
}
callbacks?.onTaskCompleted?.(videos)
resolve(videos)
return
}
if (task.status === "failed" || task.status === "cancelled") {
done = true
const rawMsg =
task.error_info?.error_message ||
task.error_message ||
(task.status === "cancelled" ? "任务已取消" : "视频生成失败,请联系管理员或重试")
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)))
callbacks?.onTaskProgress?.(pct)
if (!callbacks && runId === 0) {
onProgress(pct)
}
const timer = setTimeout(poll, 2000)
progressTimer.current.push(timer)
} catch (pollErr) {
if (cancelledRef.current || done) return
console.error("[轮询出错] taskId:", taskId, pollErr)
const status = axios.isAxiosError(pollErr) ? pollErr.response?.status : undefined
if (status && status >= 400 && status < 500) {
done = true
const msg = extractErrorMessage(pollErr, status)
callbacks?.onTaskFailed?.(msg)
reject(new Error(msg))
return
}
consecutiveErrors += 1
if (consecutiveErrors >= MAX_RETRYABLE_ERRORS) {
done = true
const msg = "任务状态查询连续失败,请稍后在任务列表查看结果"
callbacks?.onTaskFailed?.(msg)
reject(new Error(msg))
return
}
const timer = setTimeout(poll, 3000)
progressTimer.current.push(timer)
}
}
const timer = setTimeout(poll, 1500)
progressTimer.current.push(timer)
})
},
[onProgress, fetchResultsWithRetry],
)
/** 单任务轮询(单视频,兼容旧调用) */
const startPolling = useCallback(
(taskId: string) => {
cancelledRef.current = false
let consecutiveErrors = 0
const poll = async () => {
if (cancelledRef.current) return
try {
const task = await getGenerationTask(taskId)
consecutiveErrors = 0
if (task.status === "completed") {
onProgress(100)
const videos = await fetchResultsWithRetry(taskId)
if (cancelledRef.current) return
if (videos === null) {
const errorMsg = "视频已生成,但获取结果列表失败,请稍后在任务列表查看"
console.error("[生成结果获取失败] taskId:", taskId)
onFailed(errorMsg)
message.error(errorMsg)
return
}
onComplete(videos)
message.success("视频生成完成!")
return
}
if (task.status === "failed" || task.status === "cancelled") {
const rawMsg =
task.error_info?.error_message ||
task.error_message ||
(task.status === "cancelled" ? "任务已取消" : "视频生成失败,请联系管理员或重试")
const errorMsg = safeExtractError(rawMsg)
console.error("[生成失败] taskId:", taskId, "响应:", task)
onFailed(errorMsg)
message.error(errorMsg)
return
}
// pending / waiting / running — 继续轮询
const pct = Math.max(0, Math.min(99, Math.round(Number(task.progress) || 0)))
onProgress(pct)
progressTimer.current = setTimeout(poll, 2000)
} catch (pollErr) {
batchContextRef.current.clear()
pollSingleTask(taskId, 0)
.then((videos) => {
if (cancelledRef.current) return
console.error("[轮询出错] taskId:", taskId, pollErr)
onProgress(100)
onComplete(videos)
message.success("视频生成完成!")
})
.catch((err: Error) => {
if (cancelledRef.current) return
console.error("[生成失败] taskId:", taskId, err.message)
onFailed(err.message)
message.error(err.message)
})
},
[pollSingleTask, onProgress, onComplete, onFailed],
)
// 4xx 不可恢复,立即失败
const status = axios.isAxiosError(pollErr) ? pollErr.response?.status : undefined
if (status && status >= 400 && status < 500) {
const msg =
(axios.isAxiosError(pollErr) &&
(pollErr.response?.data as { detail?: string; message?: string } | undefined)
?.detail) ||
(axios.isAxiosError(pollErr) &&
(pollErr.response?.data as { detail?: string; message?: string } | undefined)
?.message) ||
`查询任务失败 (${status})`
const errorMsg = safeExtractError(msg)
onFailed(errorMsg)
message.error(errorMsg)
return
}
/**
* 批量多任务轮询:
* - 每个任务独立进度/状态回传 onBatchTaskUpdate
* * 全部完成后按变体顺序聚合视频 onComplete
* - 部分失败:整体不 onFailed(第5步逐卡片展示失败+重试按钮);全部失败才 onFailed
*/
const startPollingBatch = useCallback(
(tasks: { taskId: string; variantIndex: number }[]) => {
cancelledRef.current = false
const runId = Date.now()
const progressMap = new Map<string, number>()
const resultMap = new Map<string, unknown[]>()
const failureMap = new Map<string, string>()
batchContextRef.current = new Map(tasks.map((t) => [t.taskId, t.variantIndex]))
consecutiveErrors += 1
if (consecutiveErrors >= MAX_RETRYABLE_ERRORS) {
const errorMsg = "任务状态查询连续失败,请稍后在任务列表查看结果"
onFailed(errorMsg)
message.error(errorMsg)
return
}
progressTimer.current = setTimeout(poll, 3000)
const reportAggregateProgress = () => {
if (cancelledRef.current) return
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 checkAllSettled = () => {
if (resultMap.size + failureMap.size < tasks.length) return
if (resultMap.size === tasks.length) {
onProgress(100)
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)
}
}
progressTimer.current = setTimeout(poll, 1500)
tasks.forEach(({ taskId, variantIndex }) => {
onBatchTaskUpdate?.(taskId, {
taskId,
variantIndex,
status: "running",
progress: 0,
error: null,
videos: [],
})
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
})
})
},
[onProgress, onComplete, onFailed, fetchResultsWithRetry],
[pollSingleTask, onProgress, onComplete, onFailed, onBatchTaskUpdate],
)
return { startPolling, 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 }
}
@@ -0,0 +1,150 @@
/**
* 批量封面 HookIssue #1677
* N 个视频时:逐个自动生成封面(从对应成片抽帧 + 叠加对应标题)或上传自定义封面
*/
import { useCallback, useState } from "react"
import { message } from "antd"
import { generateCover } from "@/api/generation"
import { uploadAssetDirect, getAssetLibraries } from "@/api/assets"
import type { GeneratedVideo } from "@/api/template-editor"
interface UseBatchCoversOptions {
selectedTemplate: string
generatedVideos: GeneratedVideo[]
/** 每个变体的标题文字 */
titles: string[]
/** 标题样式(全局共用) */
titleStyle: {
font: string
size: number
color: string
position: string
bold: boolean
stroke: boolean
shadow: boolean
}
covers: string[]
onCoversChange: (urls: string[]) => void
}
export function useBatchCovers({
selectedTemplate,
generatedVideos,
titles,
titleStyle,
covers,
onCoversChange,
}: UseBatchCoversOptions) {
const [loadingIndex, setLoadingIndex] = useState<number | null>(null)
const [uploadingIndex, setUploadingIndex] = useState<number | null>(null)
const patchCover = useCallback(
(index: number, url: string) => {
const next = [...covers]
next[index] = url
onCoversChange(next)
},
[covers, onCoversChange],
)
/** 为第 index 个视频自动生成封面 */
const generateOne = useCallback(
async (index: number) => {
const finalVideos = generatedVideos.filter((v) => v.status === "completed")
const target = finalVideos[index] || generatedVideos[index]
if (!target) {
message.warning("该视频尚未生成完成")
return
}
setLoadingIndex(index)
try {
const titleText = titles[index] || ""
const response = await generateCover(selectedTemplate, {
generated_video_id: target.id,
video_url: target.file_url || target.download_url || "",
cover_type: "ai_frame",
...(titleText
? {
title_config: {
text: titleText,
font: titleStyle.font,
font_size: titleStyle.size,
font_color: titleStyle.color,
position: titleStyle.position,
bold: titleStyle.bold,
stroke: titleStyle.stroke,
shadow: titleStyle.shadow,
},
}
: {}),
})
const url = response.cover?.image_url || response.cover?.thumbnail_url || ""
if (url) {
patchCover(index, url)
message.success(`视频 ${index + 1} 封面生成成功`)
} else {
message.warning(`视频 ${index + 1} 封面生成未返回图片,请重试`)
}
} catch (err) {
console.error(`[封面] 视频 ${index + 1} 生成失败:`, err)
message.error(`视频 ${index + 1} 封面生成失败,请重试`)
} finally {
setLoadingIndex(null)
}
},
[generatedVideos, titles, titleStyle, selectedTemplate, patchCover],
)
/** 为第 index 个视频上传自定义封面 */
const uploadOne = useCallback(
async (index: number, file: File) => {
setUploadingIndex(index)
try {
const libs = await getAssetLibraries()
const imageLib = libs.find((l) => l.kind === "image") || libs[0]
if (!imageLib) {
message.error("未找到素材库,请先创建")
return
}
const result = await uploadAssetDirect({
file,
library_id: imageLib.id,
})
const url = result?.url || ""
if (url) {
patchCover(index, url)
message.success(`视频 ${index + 1} 封面已上传`)
} else {
message.warning("上传完成但未获取到图片URL,请重试")
}
} catch (err) {
console.error(`[封面] 视频 ${index + 1} 上传失败:`, err)
message.error("封面上传失败,请重试")
} finally {
setUploadingIndex(null)
}
},
[patchCover],
)
/** 一键全部自动生成(串行,避免队列限流) */
const generateAll = useCallback(async () => {
const finalVideos = generatedVideos.filter((v) => v.status === "completed")
for (let i = 0; i < finalVideos.length; i++) {
if (covers[i]) continue // 已有封面跳过
// eslint-disable-next-line no-await-in-loop
await generateOne(i)
}
message.success("全部封面已生成")
}, [generatedVideos, covers, generateOne])
return {
loadingIndex,
uploadingIndex,
generateOne,
uploadOne,
generateAll,
}
}
export default useBatchCovers
@@ -1,6 +1,6 @@
/**
* GeneratePage 表单状态管理
* 集中管理 7 步向导的所有共享状态、API 加载、URL 参数解析
* 集中管理 5 步向导的所有共享状态、API 加载、URL 参数解析
*/
import { useState } from "react"
import { useSearchParams } from "react-router-dom"
@@ -100,6 +100,26 @@ export interface GenerateFormState {
/** 从预览响应中提取的 source_edit_plan_id(供 fallback 路径使用) */
storedSourceEditPlanId: string | null
setStoredSourceEditPlanId: (planId: string | null) => void
/* ── 批量生成(Issue #1677)── */
/** 生成数量(1~10),1=单条旧逻辑 */
previewCount: number
setPreviewCount: (n: number) => void
/** 每个变体的标题文字,长度=previewCount[0] 与 titleSettings.title 保持同步 */
previewTitles: string[]
setPreviewTitles: (titles: string[] | ((prev: string[]) => string[])) => void
/** false=所有视频共用一个配音;true=每个视频独立配音 */
voiceModePerVideo: boolean
setVoiceModePerVideo: (v: boolean) => void
/** 独立配音模式下每个变体的配音素材ID,长度=previewCount */
voiceLibraryIds: string[]
setVoiceLibraryIds: (ids: string[] | ((prev: string[]) => string[])) => void
/** 每个变体的封面URL(自动生成或上传),长度=previewCount,空串=未设置 */
previewCovers: string[]
setPreviewCovers: (urls: string[] | ((prev: string[]) => string[])) => void
/** 确认生成时勾选的变体索引 */
selectedVariantIds: number[]
setSelectedVariantIds: (ids: number[] | ((prev: number[]) => number[])) => void
}
export const useGenerateFormState = (): GenerateFormState => {
@@ -185,6 +205,14 @@ export const useGenerateFormState = (): GenerateFormState => {
null,
)
/* ── 批量生成状态(Issue #1677)── */
const [previewCount, setPreviewCount] = useState(1)
const [previewTitles, setPreviewTitles] = useState<string[]>([""])
const [voiceModePerVideo, setVoiceModePerVideo] = useState(false)
const [voiceLibraryIds, setVoiceLibraryIds] = useState<string[]>([""])
const [previewCovers, setPreviewCovers] = useState<string[]>([""])
const [selectedVariantIds, setSelectedVariantIds] = useState<number[]>([0])
/* ── 从 URL / 编辑计划加载配置 ── */
usePlanConfigLoader({
editPlanId,
@@ -233,5 +261,17 @@ export const useGenerateFormState = (): GenerateFormState => {
setPreviewTaskId,
storedSourceEditPlanId,
setStoredSourceEditPlanId,
previewCount,
setPreviewCount,
previewTitles,
setPreviewTitles,
voiceModePerVideo,
setVoiceModePerVideo,
voiceLibraryIds,
setVoiceLibraryIds,
previewCovers,
setPreviewCovers,
selectedVariantIds,
setSelectedVariantIds,
}
}
@@ -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<string | null>(null)
const [generatedVideos, setGeneratedVideos] = useState<GeneratedVideo[]>([])
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
const [batchTasks, setBatchTasks] = useState<BatchTaskState[]>([])
const handleBatchTaskUpdate = useCallback((taskId: string, patch: Partial<BatchTaskState>) => {
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, clearTimer } = useGenerationPolling({
/* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入),
按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */
useEffect(() => {
if (batchTasks.length === 0) return
const byVariant = new Map<number, GeneratedVideo>()
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 {
@@ -83,7 +143,11 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
}
}
const hide = message.loading("正在生成预览视频...", 0)
const isBatch = (props.previewCount || 1) > 1
const hide = message.loading(
isBatch ? `正在生成 ${props.previewCount} 个视频...` : "正在生成预览视频...",
0,
)
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
@@ -92,6 +156,29 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
? props.selectedClonedVoice || props.selectedVoice || ""
: props.selectedVoice || ""
/* ── 批量变体数组(长度1=共用,长度=count=独立,空=回退单值) ── */
const indexes =
isBatch && props.selectedVariantIndexes?.length
? props.selectedVariantIndexes
: Array.from({ length: props.previewCount || 1 }, (_, i) => i)
const batchCount = isBatch ? indexes.length : 1
// 标题文字数组:批量时按勾选顺序
const titlesArr =
isBatch && (props.variantTitles?.length || 0) >= batchCount
? indexes.map((i) => props.variantTitles![i] || props.titleSettings?.title || "")
: []
// 配音数组:独立配音模式按勾选顺序;否则不传(回退共用 voice_library_id
const voiceArr =
isBatch && props.voiceModePerVideo && props.variantVoiceLibraryIds?.length
? indexes.map((i) => props.variantVoiceLibraryIds![i] || voiceLibraryId)
: []
// 封面数组:批量时按勾选顺序(未设置封面的变体传空串,后端回退智能封面)
const coversArr =
isBatch && props.variantCoverUrls?.length
? indexes.map((i) => props.variantCoverUrls![i] || "")
: []
try {
const taskResp = await createGenerationTask({
template_id: selectedTemplate,
@@ -109,6 +196,10 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
...(props.bgmConfig?.music_id ? { preset_id: props.bgmConfig.music_id } : {}),
},
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
...(isBatch ? { count: batchCount } : {}),
...(titlesArr.length ? { titles: titlesArr } : {}),
...(voiceArr.length ? { voice_library_ids: voiceArr } : {}),
...(coversArr.length ? { cover_urls: coversArr } : {}),
...(props.titleSettings?.title
? {
title_config: {
@@ -133,12 +224,17 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
: {}),
})
hide()
const taskId = taskResp.items?.[0]?.id
const taskIds = (taskResp.items || []).map((it) => it.id).filter(Boolean)
if (!taskId) {
if (taskIds.length === 0) {
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
}
startPolling(taskId)
if (taskIds.length > 1) {
// 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建)
startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i })))
} else {
startPolling(taskIds[0])
}
} catch (err) {
hide()
throw err
@@ -154,13 +250,21 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
return false
}
return true
}, [props, clearTimer, startPolling, selectedTemplate])
}, [props, clearTimer, startPolling, startPollingBatch, selectedTemplate])
const retry = useCallback(() => {
setGenerateError(null)
generate()
}, [generate])
/** 第5步:单独重试某个失败任务 */
const retryBatchTask = useCallback(
(taskId: string) => {
retryTask(taskId)
},
[retryTask],
)
const dismissError = useCallback(() => {
setGenerateError(null)
}, [])
@@ -205,6 +309,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
generatedVideos,
generate,
retry,
retryBatchTask,
batchTasks,
dismissError,
download,
share,
@@ -114,8 +114,11 @@ export function useServerPreview({
const resp = await createPreview(request)
if (seq !== requestSeqRef.current || !mountedRef.current) return
setTaskId(resp.task_id)
onCreatedRef.current?.(resp.task_id, resp.source_edit_plan_id)
// 兼容批量响应 {items, total}:取第一个变体
const firstTask = resp.items?.[0]
const taskId = firstTask?.task_id || ""
setTaskId(taskId)
onCreatedRef.current?.(taskId, resp.source_edit_plan_id)
let completed = false
@@ -132,7 +135,7 @@ export function useServerPreview({
if (completed || seq !== requestSeqRef.current || !mountedRef.current) return
try {
const st = await getPreviewStatus(resp.task_id)
const st = await getPreviewStatus(taskId)
if (completed || seq !== requestSeqRef.current || !mountedRef.current) return
if (st.status === "completed" && st.video_url) {
@@ -1,112 +0,0 @@
/**
* Step 7 确认生成 Hook
* 封装生成确认页的展示逻辑
*/
import { useMemo } from "react"
import { useQuery } from "@tanstack/react-query"
import type { EditingTemplate } from "@/api/editing-planner"
import type { GeneratedVideo } from "@/api/template-editor"
import type { CoverConfig } from "../types/cover"
import type { VoiceClone } from "@/api/voice-clone"
import type { PresetVoiceItem } from "@/api/voices"
import { getAssetsByKind } from "@/api/assets"
import { COVER_MODE_LABELS } from "../constants"
interface UseStep7GenerateProps {
templates: EditingTemplate[]
selectedTemplate: string
materialMode: "manual" | "auto"
selectedMaterials: string[]
smartSelectedIds: string[]
title: string
voiceMode: "preset" | "custom" | "clone"
selectedVoice: string
selectedClonedVoice: string
presetVoices: PresetVoiceItem[]
clonedVoices: VoiceClone[]
coverSettings: CoverConfig
generating: boolean
generated: boolean
generateError: string | null
progress: number
generatedVideos: GeneratedVideo[]
}
export function useStep7Generate({
templates,
selectedTemplate,
materialMode,
selectedMaterials,
smartSelectedIds,
title,
voiceMode: _voiceMode,
selectedVoice,
selectedClonedVoice: _selectedClonedVoice,
presetVoices: _presetVoices,
clonedVoices: _clonedVoices,
coverSettings,
generating,
generated,
generateError,
progress,
generatedVideos,
}: UseStep7GenerateProps) {
const templateName = useMemo(
() => templates.find((t) => t.id === selectedTemplate)?.name ?? "未选择",
[templates, selectedTemplate],
)
const materialSummary = useMemo(() => {
if (materialMode === "auto") {
return `${smartSelectedIds.length} 个素材(智能匹配)`
}
return `${selectedMaterials.length} 个素材`
}, [materialMode, selectedMaterials.length, smartSelectedIds.length])
// 从配音素材库中查找 voiceName
const { data: voiceMaterials = [] } = useQuery({
queryKey: ["assets", "voice"],
queryFn: () => getAssetsByKind("voice", { limit: 50 }),
})
const voiceName = useMemo(() => {
const asset = voiceMaterials.find((v) => v.id === selectedVoice)
return asset ? asset.name : "未选择"
}, [voiceMaterials, selectedVoice])
const coverSummary = useMemo(() => {
if (!coverSettings.enabled) return "不使用"
return COVER_MODE_LABELS[coverSettings.mode] || "智能封面"
}, [coverSettings])
const getGenerationPhase = (p: number) => {
if (p < 20) return { label: "分析素材与配置", icon: "🔍" }
if (p < 50) return { label: "智能剪辑合成", icon: "🎬" }
if (p < 80) return { label: "渲染视频中", icon: "⚡" }
return { label: "即将完成", icon: "✨" }
}
const handleScrollToPreview = () => {
const el =
document.querySelector(".xx-inline-video-player") ||
document.querySelector(".xx-preview-section")
el?.scrollIntoView({ behavior: "smooth", block: "start" })
}
return {
templateName,
materialSummary,
title,
voiceName,
coverSummary,
generating,
generated,
generateError,
progress,
generatedVideos,
getGenerationPhase,
handleScrollToPreview,
}
}
export default useStep7Generate
@@ -1,6 +1,11 @@
/**
* GeneratePage 步骤导航
* 步骤顺序(6步):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6)
* 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,10 +18,10 @@ export interface UseStepNavigationOptions {
selectedMaterials: string[]
smartSelectedIds: string[]
titleSettings: TitleSettings
/** 预览是否已就绪(素材已加载,可播放 */
previewReady: boolean
/** 是否已完成视频生成(步骤5确认生成后才能进入封面) */
/** 是否已完成视频生成(步骤5全部渲染完成后才能进入封面 */
generated: boolean
/** Step1 点下一步时弹出数量选择弹窗 */
onOpenCountModal: () => void
}
export interface UseStepNavigationReturn {
@@ -32,14 +37,18 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav
materialMode,
selectedMaterials,
smartSelectedIds,
titleSettings,
previewReady,
generated,
onOpenCountModal,
} = options
const goNext = () => {
if (currentStep === 1 && !selectedTemplate) {
message.warning("请先选择一个模板")
if (currentStep === 1) {
if (!selectedTemplate) {
message.warning("请先选择一个模板")
return
}
// 选完模板弹数量选择弹窗(每次都弹,不记忆)
onOpenCountModal()
return
}
if (currentStep === 2 && materialMode === "manual" && selectedMaterials.length === 0) {
@@ -50,21 +59,12 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav
message.warning("请先进行智能匹配并选择素材")
return
}
// Step4(标题+预览):标题必填 + 预览必须已加载
if (currentStep === 4) {
if (!titleSettings.title.trim()) {
message.warning("请选择或输入标题")
// 步骤5(确认生成):全部渲染完成后才能下一步进封面
if (currentStep === 5) {
if (!generated) {
message.warning("视频还在渲染中,请等待生成完成")
return
}
if (!previewReady) {
message.warning("预览视频正在加载,请稍候")
return
}
}
// Step5(确认生成):必须已完成生成才能进入封面
if (currentStep === 5 && !generated) {
message.warning("请先生成视频")
return
}
if (currentStep < 6) {
setCurrentStep((s) => s + 1)
@@ -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);
}
+154 -25
View File
@@ -1,39 +1,110 @@
/**
* 个人设置页面
* 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, 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 = () => {
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 [wechatBindOpen, setWechatBindOpen] = useState(false)
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 handleBindSuccess = () => {
setWechatBindOpen(false)
queryClient.invalidateQueries({ queryKey: ["currentUser"] })
message.success("微信绑定成功")
}
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 (
<div className="xx-settings-page">
<PageHead title="个人设置" description="管理您的账户信息" />
<div className="xx-settings-card">
<h3></h3>
<div className="xx-settings-notice">
<span className="xx-settings-notice-icon"></span>
<div>
<strong></strong>
<p></p>
</div>
</div>
<div className="xx-settings-form">
<div className="xx-settings-field">
<label className="xx-settings-label"></label>
@@ -42,25 +113,83 @@ const Settings: React.FC = () => {
<div className="xx-settings-field">
<label className="xx-settings-label"></label>
<Input value={user?.email || ""} disabled placeholder="邮箱" />
</div>
<div className="xx-settings-field">
<label className="xx-settings-label"></label>
<Input
value={displayName}
onChange={(e) => setDisplayName(e.target.value)}
placeholder="请输入显示名称"
value={user?.email && !user.email.endsWith("@wechat.local") ? user.email : ""}
disabled
placeholder={user?.email?.endsWith("@wechat.local") ? "微信账号暂未绑定邮箱" : "邮箱"}
/>
</div>
<div className="xx-settings-field">
<Button buttonType="primary" buttonSize="md" onClick={handleSave} disabled>
<label className="xx-settings-label"></label>
<Input
value={displayName}
onChange={(e) => setDisplayName(e.target.value)}
placeholder="请输入昵称"
maxLength={20}
/>
</div>
<div className="xx-settings-field">
<Button
buttonType="primary"
buttonSize="md"
onClick={() => saveProfileMutation.mutate()}
loading={saveProfileMutation.isPending}
disabled={!displayName.trim() || !displayNameDirty}
>
</Button>
</div>
</div>
</div>
<div className="xx-settings-card">
<h3></h3>
<div className="xx-settings-wechat">
<div className="xx-settings-wechat-info">
<span className="xx-wechat-icon">💬</span>
<div>
{wechatBound ? (
<>
<strong>
{user?.wechat_nickname ? `${user.wechat_nickname}` : ""}
</strong>
<p>使</p>
</>
) : (
<>
<strong></strong>
<p>使</p>
</>
)}
</div>
</div>
<div className="xx-settings-wechat-actions">
{wechatBound ? (
<Button
buttonType="ghost"
buttonSize="md"
onClick={handleUnbind}
loading={unbindMutation.isPending}
>
</Button>
) : (
<Button buttonType="primary" buttonSize="md" onClick={() => setWechatBindOpen(true)}>
</Button>
)}
</div>
</div>
</div>
<WechatQrModal
open={wechatBindOpen}
scene="bind"
onClose={() => setWechatBindOpen(false)}
onBindSuccess={handleBindSuccess}
/>
</div>
)
}
+6
View File
@@ -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 <Navigate to="/login" replace />
}
// 微信新用户未完成昵称引导时,禁止进入主界面
if (!profileCompleted) {
return <Navigate to="/welcome/wechat" replace />
}
return <>{children}</>
}
+10
View File
@@ -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: <WechatCallback />,
},
{
path: "/auth/wechat/bind/callback",
element: <WechatBindCallback />,
},
{
path: "/welcome/wechat",
element: <WechatOnboarding />,
},
]
+6
View File
@@ -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 {
+106
View File
@@ -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 只记录 preparecomplete 若调用会抛 "不应被调用"
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 XMLHttpRequestsend 之后下一 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()
+130
View File
@@ -0,0 +1,130 @@
/**
* 上传去重/幂等工具单测(Issue #1714
*/
import { describe, it, expect, vi } from "vitest"
import {
computeFileHash,
findDuplicateInQueue,
HASH_FULL_READ_LIMIT,
HASH_SAMPLE_CHUNK,
makeClientUploadId,
makeFileFingerprint,
} from "@/api/assets/uploadDedup"
const makeFile = (name: string, size = 100, lastModified = 1_700_000_000_000) =>
new File([new Uint8Array(size)], name, { type: "video/mp4", lastModified })
describe("makeFileFingerprint", () => {
it("同一文件(name+size+lastModified 相同)指纹一致", () => {
const a = makeFile("a.mp4", 1000, 12345)
const b = makeFile("a.mp4", 1000, 12345)
expect(makeFileFingerprint(a)).toBe(makeFileFingerprint(b))
})
it("文件名/大小/修改时间任一不同指纹即不同", () => {
const base = makeFile("a.mp4", 1000, 100)
expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("b.mp4", 1000, 100)))
expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("a.mp4", 1001, 100)))
expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("a.mp4", 1000, 101)))
})
})
describe("findDuplicateInQueue", () => {
const queue = [
{ fileKey: "k1", status: "preparing" },
{ fileKey: "k2", status: "uploading" },
{ fileKey: "k3", status: "ingesting" },
{ fileKey: "k4", status: "done" },
{ fileKey: "k5", status: "error" },
]
it("在途状态(preparing/uploading/ingesting/done)命中重复", () => {
expect(findDuplicateInQueue(queue, "k1")?.status).toBe("preparing")
expect(findDuplicateInQueue(queue, "k2")?.status).toBe("uploading")
expect(findDuplicateInQueue(queue, "k3")?.status).toBe("ingesting")
expect(findDuplicateInQueue(queue, "k4")?.status).toBe("done")
})
it("未命中返回 null", () => {
expect(findDuplicateInQueue(queue, "missing")).toBeNull()
})
it("排除 error 状态后,失败项不算重复(允许重新激活)", () => {
expect(findDuplicateInQueue(queue, "k5", ["error"])).toBeNull()
})
it("同时排除 done 后,已完成项也不算重复", () => {
expect(findDuplicateInQueue(queue, "k4", ["error", "done"])).toBeNull()
// 但在途的仍然命中
expect(findDuplicateInQueue(queue, "k1", ["error", "done"])).not.toBeNull()
})
})
describe("makeClientUploadId", () => {
it("生成带前缀且互不相同的幂等 token", () => {
const ids = new Set(Array.from({ length: 20 }, () => makeClientUploadId()))
expect(ids.size).toBe(20)
for (const id of ids) expect(id.startsWith("up_")).toBe(true)
})
})
describe("computeFileHash", () => {
it("相同内容 hash 一致、不同内容 hash 不同", async () => {
const f1 = makeFile("a.mp4", 4096)
const f2 = makeFile("b.mp4", 4096)
// 两个文件都是 0 填充,内容相同 → hash 一致
expect(await computeFileHash(f1)).toBe(await computeFileHash(f2))
const f3 = new File([new Uint8Array(4096).fill(7)], "c.mp4", { type: "video/mp4" })
expect(await computeFileHash(f1)).not.toBe(await computeFileHash(f3))
})
it("返回 64 位十六进制(SHA-256,与后端 file_hash 长度一致)", async () => {
const hash = await computeFileHash(makeFile("a.mp4", 1024))
expect(hash).toMatch(/^[0-9a-f]{64}$/)
})
})
describe("computeFileHash 大文件抽样(>64MB", () => {
it("抽样路径正常返回 64 位 hex,且大小不同则 hash 不同", async () => {
// mock 一个「声称」300MB 的 Fileslice 返回小 buffer 即可,不真分配 300MB
const makeBig = (declaredSize: number, head: number) => {
const f = new File([new Uint8Array([head, 2, 3])], "big.mov", { type: "video/quicktime" })
Object.defineProperty(f, "size", { value: declaredSize, configurable: true })
// slice 仍按真实内容返回小片段(头尾片段内容由底层小 buffer 决定)
return f
}
const h1 = await computeFileHash(makeBig(300 * 1024 * 1024, 1))
const h2 = await computeFileHash(makeBig(301 * 1024 * 1024, 1))
expect(h1).toMatch(/^[0-9a-f]{64}$/)
// 声明大小不同 → 写入的 64 位 size 字段不同 → hash 必须不同(锁定 setBigUint64 路径)
expect(h1).not.toBe(h2)
})
it("≤64MB 走全量读取(slice 一次覆盖整个文件)", async () => {
const f = new File([new Uint8Array(1024).fill(9)], "full.mp4", { type: "video/mp4" })
Object.defineProperty(f, "size", { value: HASH_FULL_READ_LIMIT, configurable: true })
const sliceSpy = vi.spyOn(f, "slice")
await computeFileHash(f)
// 全量路径:唯一一次 slice 为 (0, size)
expect(sliceSpy).toHaveBeenCalledTimes(1)
expect(sliceSpy).toHaveBeenCalledWith(0, HASH_FULL_READ_LIMIT)
sliceSpy.mockRestore()
})
it(">64MB 只读取头尾各 16MB 抽样,绝不整文件读入内存", async () => {
const f = new File([new Uint8Array(1024).fill(9)], "big.mp4", { type: "video/mp4" })
Object.defineProperty(f, "size", { value: HASH_FULL_READ_LIMIT + 1, configurable: true })
const sliceSpy = vi.spyOn(f, "slice")
await computeFileHash(f)
// 抽样路径:两次 slice —— 头部 (0, 16MB) 与尾部 (size-16MB, size)
expect(sliceSpy).toHaveBeenCalledTimes(2)
expect(sliceSpy).toHaveBeenNthCalledWith(1, 0, HASH_SAMPLE_CHUNK)
expect(sliceSpy).toHaveBeenNthCalledWith(
2,
HASH_FULL_READ_LIMIT + 1 - HASH_SAMPLE_CHUNK,
HASH_FULL_READ_LIMIT + 1,
)
sliceSpy.mockRestore()
})
})
+64
View File
@@ -0,0 +1,64 @@
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
describe("wxLogin 工具", () => {
describe("parseWxAuthUrl", () => {
it("从微信授权链接解析出 appid/redirect_uri/stateredirect_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(/加载失败/)
})
})
})
@@ -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<string, unknown>) =>
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(<WechatQrModal open={false} scene="login" onClose={vi.fn()} />)
expect(screen.queryByText("微信扫码登录")).toBeNull()
})
it("登录场景:open 后请求授权链接、写入 state、用 WxLogin 渲染二维码", async () => {
render(<WechatQrModal open scene="login" onClose={vi.fn()} />)
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(<WechatQrModal open scene="bind" onClose={vi.fn()} />)
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(<WechatQrModal open scene="login" onClose={vi.fn()} />)
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(<WechatQrModal open scene="login" onClose={vi.fn()} onLoginSuccess={onSuccess} />)
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(<WechatQrModal open scene="login" onClose={vi.fn()} />)
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(<WechatQrModal open scene="bind" onClose={vi.fn()} onBindSuccess={onBindSuccess} />)
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(<WechatQrModal open scene="login" onClose={vi.fn()} onLoginSuccess={onSuccess} />)
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()
})
})
+124 -33
View File
@@ -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 }) => (
<div data-testid="page-head">
@@ -12,51 +12,142 @@ vi.mock("@/components/layout/PageHead", () => ({
),
}))
const mockSetUser = vi.fn()
const mockInvalidate = vi.fn()
let authState: Record<string, unknown> = {
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<string, unknown>)
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(
<QueryClientProvider client={queryClient}>
<MemoryRouter>
<Settings />
</MemoryRouter>
</QueryClientProvider>,
)
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(
<MemoryRouter>
<Settings />
</MemoryRouter>,
)
expect(screen.getByText("个人设置")).toBeTruthy()
setUser: mockSetUser,
}
})
it("should display user info", () => {
render(
<MemoryRouter>
<Settings />
</MemoryRouter>,
)
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(
<MemoryRouter>
<Settings />
</MemoryRouter>,
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: "新昵称" })
})
})
})
@@ -26,22 +26,32 @@ interface FakeHandle {
fields: Record<string, string>
max_size_bytes: number
asset_id: string
duplicated?: boolean
skip_transfer?: boolean
}
transfer: ReturnType<typeof vi.fn>
complete: ReturnType<typeof vi.fn>
/** 手动结束传输(transfer 被调用后挂载);finish(true) 以失败结束 */
finish: (fail?: boolean) => void
/** complete 已被调用的次数 */
completeCalls: { resolve: () => void; reject: (err: unknown) => void }[]
}
let activeTransfers = 0
let maxConcurrent = 0
/**
* 创建一个假 handletransfer 返回挂起的 promise
* finish 槽位在 transfer executor 同步执行时挂载,测试中调用 finish() 控制成败
* 创建一个假 handle
* - transfer 返回挂起的 promisefinish()/finish(true) 控制成败
* - complete 每次调用返回独立的挂起 promise,由 completeCalls 记录控制,
* 成功调 resolve(idx) / 失败调 reject(idx)(模拟超时)
*/
const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?: boolean }) => {
const h = {
const makeFakeHandle = (opts: {
id: string
duplicated?: boolean
failTransfer?: boolean
completeAuto?: boolean
/** prepare 阶段就命中去重:prepare 响应 skip_transfer/duplicated=true */
prepareDedup?: boolean
}) => {
const h: FakeHandle = {
prepared: {
upload_url: "https://oss.example.com/u",
method: "POST",
@@ -50,17 +60,43 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?:
fields: {},
max_size_bytes: 2_000_000_000,
asset_id: opts.id,
duplicated: opts.prepareDedup ? true : undefined,
skip_transfer: opts.prepareDedup ? true : undefined,
},
transfer: vi.fn(),
complete: vi.fn().mockResolvedValue({
storage_key: "uploads/x/y.mp4",
ingest_job_id: opts.duplicated ? "" : "job-1",
url: "https://oss.example.com/u",
duplicated: opts.duplicated,
asset_id: opts.id,
}),
finish: (() => {}) as (fail?: boolean) => void,
complete: vi.fn(),
finish: () => {},
completeCalls: [],
}
h.complete.mockImplementation(
() =>
new Promise<{
storage_key: string
ingest_job_id: string
url: string
duplicated: boolean
asset_id: string
}>((resolve, reject) => {
h.completeCalls.push({
resolve: () =>
resolve({
storage_key: "uploads/x/y.mp4",
ingest_job_id: opts.duplicated ? "" : `job-${opts.id}`,
url: "https://oss.example.com/u",
duplicated: !!opts.duplicated,
asset_id: opts.id,
}),
reject,
})
// 默认立即成功,保持旧用例简单
if (opts.completeAuto !== false) {
const idx = h.completeCalls.length - 1
Promise.resolve().then(() => h.completeCalls[idx]?.resolve())
}
}),
)
h.transfer.mockImplementation(
() =>
new Promise<void>((_resolve, reject) => {
@@ -78,17 +114,24 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?:
type FakeHandleLike = ReturnType<typeof makeFakeHandle>
let activeTransfers = 0
let maxConcurrent = 0
/** prepare mock:调用序号生成稳定 id,立即把 handle(含 finish 槽位)推入数组 */
const installPrepareMock = (
handles: FakeHandleLike[],
optOverrides?: (id: string) => { duplicated?: boolean; failTransfer?: boolean },
optOverrides?: (id: string) => {
duplicated?: boolean
failTransfer?: boolean
completeAuto?: boolean
},
) => {
let callNo = 0
;(prepareDirectUploadHandle as unknown as ReturnType<typeof vi.fn>).mockImplementation(
async () => {
const id = `asset-${callNo++}`
const overrides = optOverrides?.(id) ?? {}
const h = makeFakeHandle({ id, ...overrides })
const h = makeFakeHandle({ id, completeAuto: true, ...overrides })
handles.push(h)
await new Promise((r) => setTimeout(r, 10))
return h
@@ -187,6 +230,11 @@ describe("useAssetUpload", () => {
})
await waitFor(() => expect(result.current.uploadItems[0].status).toBe("error"))
// 失败卡片记录失败阶段与完整错误原因(不再只显示"上传失败")
const failed = result.current.uploadItems[0]
expect(failed.failedStage).toBe("transfer")
expect(failed.error).toContain("OSS boom")
// 重试:重新 preparehandles[1] 成功)
const tempId = result.current.uploadItems[0].tempId
await act(async () => {
@@ -223,4 +271,142 @@ describe("useAssetUpload", () => {
expect(result.current.uploadItems[0].status).toBe("done")
})
})
it("同一文件多次选择不重复入队(指纹去重)", async () => {
const handles: FakeHandleLike[] = []
installPrepareMock(handles)
const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), {
wrapper: createWrapper(),
})
// 同一文件(name+size+lastModified 完全一致)第一次入队
const sameFile = mp4("same.mp4")
await act(async () => {
result.current.enqueueUploads([sameFile])
})
await waitFor(() => expect(handles.length).toBe(1))
expect(result.current.uploadItems).toHaveLength(1)
// transfer 挂起期间,再次选择同一文件(模拟用户反复点选/拖拽)
await act(async () => {
result.current.enqueueUploads([sameFile])
})
await act(async () => {
result.current.enqueueUploads([sameFile])
})
// 队列表只有 1 项、prepare 只有 1 次
expect(result.current.uploadItems).toHaveLength(1)
expect(handles.length).toBe(1)
// 完成后再次重复选择(已 done):仍然不新增
await act(async () => {
handles[0].finish()
})
await waitFor(() => expect(result.current.uploadItems[0].status).toBe("done"))
await act(async () => {
result.current.enqueueUploads([sameFile])
})
expect(result.current.uploadItems).toHaveLength(1)
expect(handles.length).toBe(1)
// 不同文件正常入队
await act(async () => {
result.current.enqueueUploads([mp4("other.mp4")])
})
await waitFor(() => expect(handles.length).toBe(2))
expect(result.current.uploadItems).toHaveLength(2)
})
it("complete 失败(超时)后重试:只重发 complete,不重新 prepare/直传", async () => {
const handles: FakeHandleLike[] = []
installPrepareMock(handles, () => ({ completeAuto: false }))
const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), {
wrapper: createWrapper(),
})
await act(async () => {
result.current.enqueueUploads([mp4("slow.mp4")])
})
await waitFor(() => expect(handles.length).toBe(1))
await waitFor(() => expect(handles[0].transfer).toHaveBeenCalled())
await act(async () => {
handles[0].finish()
})
// complete 被调用但挂起
await waitFor(() => expect(handles[0].complete).toHaveBeenCalledTimes(1))
// 模拟 complete 超时(后端记录可能已建成)
await act(async () => {
handles[0].completeCalls[0]?.reject(new Error("complete timeout (ECONNABORTED)"))
})
const tempId = result.current.uploadItems[0].tempId
await waitFor(() => {
const it = result.current.uploadItems.find((x) => x.tempId === tempId)
expect(it?.status).toBe("error")
expect(it?.failedStage).toBe("complete")
// 卡片同时展示真实失败原因与"重试不会重新上传"提示
expect(it?.error).toContain("complete timeout")
expect(it?.error).toContain("不会重新上传文件")
})
// 点重试:pump 复用 handle,只再调一次 completetransfer/prepare 不重复)
await act(async () => {
result.current.retryUpload(tempId)
})
await waitFor(() => expect(handles[0].complete).toHaveBeenCalledTimes(2))
expect(handles.length).toBe(1) // 没有重新 prepare
expect(handles[0].transfer).toHaveBeenCalledTimes(1) // 没有重新直传
// 第二次 complete 成功
await act(async () => {
handles[0].completeCalls[1]?.resolve()
})
await waitFor(() => {
expect(result.current.uploadItems.find((x) => x.tempId === tempId)?.status).toBe("done")
})
})
it("prepare 阶段失败:标记 prepare 阶段并保留后端错误明细", async () => {
;(prepareDirectUploadHandle as unknown as ReturnType<typeof vi.fn>).mockRejectedValueOnce({
isAxiosError: true,
response: { status: 500, data: { detail: "签名服务内部错误" } },
message: "Request failed with status code 500",
})
const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), {
wrapper: createWrapper(),
})
await act(async () => {
result.current.enqueueUploads([mp4("prep-fail.mp4")])
})
await waitFor(() => expect(result.current.uploadItems[0]?.status).toBe("error"))
const it = result.current.uploadItems[0]
expect(it.failedStage).toBe("prepare")
expect(it.error).toContain("签名服务内部错误")
})
it("prepare 返回 skip_transfer=true 时立即跳过 transfer+complete,标记 done+duplicated", async () => {
const h = makeFakeHandle({ id: "a-skip", prepareDedup: true })
;(prepareDirectUploadHandle as unknown as ReturnType<typeof vi.fn>).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")
})
})
})
@@ -0,0 +1,144 @@
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
import { render, screen, waitFor, cleanup } from "@testing-library/react"
import { MemoryRouter } from "react-router-dom"
import WechatBindCallback from "@/pages/auth/WechatBindCallback"
const mockNavigate = vi.fn()
const mockSetUser = vi.fn()
const mockParams = new URLSearchParams({ code: "bind_code", state: "bind_state" })
const mockSearchParams = [mockParams] as const
const localStorageStore: Record<string, string> = {}
vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => localStorageStore[key] || null)
vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => {
localStorageStore[key] = val
})
vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => {
delete localStorageStore[key]
})
let bindError: unknown = null
const mockBindResult = { user: { id: "u1", wechat_bound: true } }
vi.mock("react-router-dom", async () => {
const actual = await vi.importActual("react-router-dom")
return {
...actual,
useNavigate: () => mockNavigate,
useSearchParams: () => mockSearchParams,
}
})
vi.mock("@/api/auth", () => ({
bindWechat: vi.fn(async () => {
if (bindError) throw bindError
return mockBindResult
}),
normalizeUser: (u: unknown) => u,
}))
vi.mock("@/store/authStore", () => ({
useAuthStore: (selector: (state: unknown) => unknown) => selector({ setUser: mockSetUser }),
}))
// iframe 场景:默认非 iframe;用例可 mockReturnValue(true)
const { mockIsInIframe, mockPostResult } = vi.hoisted(() => ({
mockIsInIframe: vi.fn(() => false),
mockPostResult: vi.fn(),
}))
vi.mock("@/components/auth/WechatQrModal/messages", () => ({
isInIframe: () => mockIsInIframe(),
postWechatQrResult: (...args: unknown[]) => mockPostResult(...args),
}))
const renderPage = () =>
render(
<MemoryRouter>
<WechatBindCallback />
</MemoryRouter>,
)
describe("WechatBindCallback Page", () => {
afterEach(() => {
cleanup()
})
beforeEach(() => {
vi.clearAllMocks()
mockIsInIframe.mockReturnValue(false)
bindError = null
Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k))
mockParams.set("code", "bind_code")
mockParams.set("state", "bind_state")
localStorageStore.wechat_bind_state = "bind_state"
})
it("绑定成功跳转设置页并携带 success 标记", async () => {
renderPage()
await waitFor(() => {
expect(mockNavigate).toHaveBeenCalledWith("/app/profile?wechat_bind=success", {
replace: true,
})
})
expect(mockSetUser).toHaveBeenCalled()
})
it("本地无 wechat_bind_state(微信内/跨浏览器)不再误杀,绑定正常完成", async () => {
delete localStorageStore.wechat_bind_state
renderPage()
await waitFor(() => {
expect(mockNavigate).toHaveBeenCalledWith("/app/profile?wechat_bind=success", {
replace: true,
})
})
})
it("后端报错(微信已被其他账号绑定)时页面透传真实原因,不静默跳走", async () => {
bindError = {
isAxiosError: true,
response: { status: 409, data: { detail: "该微信已绑定其他账号" } },
message: "Request failed with status code 409",
}
renderPage()
await waitFor(() => {
expect(screen.getByText(/该微信已绑定其他账号/)).toBeTruthy()
})
expect(mockNavigate).not.toHaveBeenCalled()
})
it("缺少 code/state 时提示无效回调", async () => {
mockParams.delete("code")
renderPage()
await waitFor(() => {
expect(screen.getByText(/无效的回调参数/)).toBeTruthy()
})
})
describe("iframe(弹窗内嵌二维码)场景", () => {
it("绑定成功时 postMessage 通知父窗口,不做 navigate", async () => {
mockIsInIframe.mockReturnValue(true)
renderPage()
await waitFor(() => {
expect(mockPostResult).toHaveBeenCalledWith("bind", true)
})
expect(mockSetUser).toHaveBeenCalled()
expect(mockNavigate).not.toHaveBeenCalled()
})
it("绑定失败时把真实原因 postMessage 给父窗口", async () => {
mockIsInIframe.mockReturnValue(true)
bindError = {
isAxiosError: true,
response: { status: 409, data: { detail: "该微信已绑定其他账号" } },
}
renderPage()
await waitFor(() => {
expect(mockPostResult).toHaveBeenCalledWith("bind", false, {
detail: expect.stringContaining("该微信已绑定其他账号"),
})
})
expect(screen.queryByText(/返回设置/)).toBeNull()
expect(mockNavigate).not.toHaveBeenCalled()
})
})
})
@@ -1,79 +1,223 @@
import { describe, expect, it, vi, beforeEach } from "vitest"
import { render, screen } from "@testing-library/react"
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
import { render, screen, waitFor, cleanup } from "@testing-library/react"
import { MemoryRouter } from "react-router-dom"
import WechatCallback from "@/pages/auth/WechatCallback"
const mockNavigate = vi.fn()
const mockSetAuth = vi.fn()
// useSearchParams 返回模块级稳定引用(数组元素同一 URLSearchParams 实例),
// 避免每次 render 返回新数组/新实例导致 useEffect 依赖变化重跑
const mockParams = new URLSearchParams({ code: "test_code", state: "test_state" })
const mockSearchParams = [mockParams] as const
const mockAuthState = { setAuth: mockSetAuth }
// 文件级 localStorage mock(避免每个用例重复 spy 导致链式污染)
const localStorageStore: Record<string, string> = {}
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<string, unknown> = {}
let mockCurrentUser: Record<string, unknown> = {}
let callbackError: unknown = null
vi.mock("react-router-dom", async () => {
const actual = await vi.importActual("react-router-dom")
return {
...actual,
useNavigate: () => vi.fn(),
useSearchParams: () => [new URLSearchParams({ code: "test_code", state: "test_state" })],
useNavigate: () => mockNavigate,
useSearchParams: () => mockSearchParams,
}
})
vi.mock("@/api/auth", () => ({
wechatCallback: vi.fn(() => new Promise(() => {})), // pending promise,保持loading
getCurrentUser: vi.fn(),
wechatCallback: vi.fn(async () => {
if (callbackError) throw callbackError
return mockCallbackResult
}),
getCurrentUser: vi.fn(async () => mockCurrentUser),
normalizeUser: (u: unknown) => u,
}))
vi.mock("@/api/auth/tokenRefresh", () => ({
scheduleProactiveRefresh: vi.fn(),
cancelProactiveRefresh: vi.fn(),
}))
vi.mock("@/store/authStore", () => ({
useAuthStore: () => ({
setAuth: vi.fn(),
}),
useAuthStore: (selector: (state: unknown) => unknown) => selector({ setAuth: mockSetAuth }),
}))
vi.mock("@/components/auth/BindContactModal", () => ({
default: ({ open }: { open: boolean }) => (
<div data-testid="bind-contact-modal" style={{ display: open ? "block" : "none" }}>
BindContactModal
</div>
),
// iframe 场景:默认非 iframe;用例可 mockReturnValue(true)
const { mockIsInIframe, mockPostResult } = vi.hoisted(() => ({
mockIsInIframe: vi.fn(() => false),
mockPostResult: vi.fn(),
}))
vi.mock("@/components/auth/WechatQrModal/messages", () => ({
isInIframe: () => mockIsInIframe(),
postWechatQrResult: (...args: unknown[]) => mockPostResult(...args),
}))
vi.mock("antd", async () => {
const actual = await vi.importActual("antd")
return {
...actual,
message: {
success: vi.fn(),
error: vi.fn(),
},
}
})
const renderPage = () =>
render(
<MemoryRouter>
<WechatCallback />
</MemoryRouter>,
)
describe("WechatCallback Page", () => {
afterEach(() => {
cleanup()
})
beforeEach(() => {
// mock localStorage,设置wechat_state匹配,让校验通过
const store: Record<string, string> = {
wechat_state: "test_state",
vi.clearAllMocks()
mockIsInIframe.mockReturnValue(false)
callbackError = null
// 默认正常回调参数;用例可改写 mockParams 模拟 error 重定向
Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k))
mockParams.set("code", "test_code")
mockParams.set("state", "test_state")
localStorageStore.wechat_state = "test_state"
mockCallbackResult = {
access_token: "at",
refresh_token: "rt",
is_new_user: false,
}
vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => store[key] || null)
vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => {
store[key] = val
mockCurrentUser = {
id: "u1",
display_name: "老用户",
profile_completed: true,
}
})
it("老用户登录成功跳转首页/来源页", async () => {
renderPage()
await waitFor(() => {
expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true })
})
vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => {
delete store[key]
expect(mockSetAuth).toHaveBeenCalled()
})
it("新用户(is_new_user)跳转昵称引导页", async () => {
mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: true }
mockCurrentUser = { id: "u2", display_name: "微信用户", profile_completed: false }
renderPage()
await waitFor(() => {
expect(mockNavigate).toHaveBeenCalledWith("/welcome/wechat", { replace: true })
})
})
it("should render without crashing", () => {
const { container } = render(
<MemoryRouter>
<WechatCallback />
</MemoryRouter>,
)
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(
<MemoryRouter>
<WechatCallback />
</MemoryRouter>,
)
// wechatCallback 返回 pending promise,所以应该显示 loading
expect(screen.getByText("正在登录...")).toBeTruthy()
it("本地无 wechat_state(微信内打开/跨浏览器场景)不再误杀,正常完成登录", async () => {
delete localStorageStore.wechat_state
renderPage()
await waitFor(() => {
expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true })
})
// state 已被清理
expect(localStorageStore.wechat_state).toBeUndefined()
})
it("后端返回 detail 错误时,页面透传真实原因(不再吞成通用提示)", async () => {
callbackError = {
isAxiosError: true,
response: { status: 400, data: { detail: "微信授权码已过期,请重新扫码" } },
message: "Request failed with status code 400",
}
renderPage()
await waitFor(() => {
expect(screen.getByText(/微信授权码已过期,请重新扫码/)).toBeTruthy()
})
expect(screen.queryByText(/^微信登录失败,请重试$/)).toBeNull()
expect(mockNavigate).not.toHaveBeenCalled()
})
it("微信重定向带 error(用户拒绝授权)时展示授权失败原因", async () => {
for (const k of Array.from(mockParams.keys())) mockParams.delete(k)
mockParams.set("error", "access_denied")
mockParams.set("error_description", "The+user+denied+the+request")
renderPage()
await waitFor(() => {
expect(screen.getByText(/微信授权失败/)).toBeTruthy()
expect(screen.getByText(/access_denied/)).toBeTruthy()
})
expect(mockNavigate).not.toHaveBeenCalled()
})
it("缺少 code/state 参数时提示无效回调", async () => {
mockParams.delete("code")
renderPage()
await waitFor(() => {
expect(screen.getByText(/无效的回调参数/)).toBeTruthy()
})
})
it("处理中显示 loading", () => {
renderPage()
expect(screen.getByText("微信登录中...")).toBeTruthy()
})
describe("iframe(弹窗内嵌二维码)场景", () => {
it("登录成功时 postMessage 通知父窗口(needOnboarding=false),不做 navigate", async () => {
mockIsInIframe.mockReturnValue(true)
renderPage()
await waitFor(() => {
expect(mockPostResult).toHaveBeenCalledWith("login", true, { needOnboarding: false })
})
expect(mockSetAuth).toHaveBeenCalled()
expect(mockNavigate).not.toHaveBeenCalled()
})
it("新用户成功时上报 needOnboarding=true", async () => {
mockIsInIframe.mockReturnValue(true)
mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: true }
renderPage()
await waitFor(() => {
expect(mockPostResult).toHaveBeenCalledWith("login", true, { needOnboarding: true })
})
expect(mockNavigate).not.toHaveBeenCalled()
})
it("后端报错时把真实原因 postMessage 给父窗口,页面不渲染错误/按钮", async () => {
mockIsInIframe.mockReturnValue(true)
callbackError = {
isAxiosError: true,
response: { status: 400, data: { detail: "state 已过期或已被使用" } },
}
renderPage()
await waitFor(() => {
expect(mockPostResult).toHaveBeenCalledWith("login", false, {
detail: expect.stringContaining("state 已过期或已被使用"),
})
})
expect(screen.queryByText(/返回登录/)).toBeNull()
expect(mockNavigate).not.toHaveBeenCalled()
})
it("微信重定向 error(拒绝授权)在 iframe 内也上报父窗口", async () => {
mockIsInIframe.mockReturnValue(true)
for (const k of Array.from(mockParams.keys())) mockParams.delete(k)
mockParams.set("error", "access_denied")
renderPage()
await waitFor(() => {
expect(mockPostResult).toHaveBeenCalledWith("login", false, {
detail: expect.stringContaining("access_denied"),
})
})
})
})
})
@@ -0,0 +1,186 @@
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
import { render, screen, fireEvent, waitFor, cleanup } from "@testing-library/react"
import { MemoryRouter } from "react-router-dom"
import { QueryClient, QueryClientProvider } from "@tanstack/react-query"
import WechatOnboarding from "@/pages/auth/WechatOnboarding"
const queryClient = new QueryClient({
defaultOptions: { queries: { retry: false }, mutations: { retry: false } },
})
const mockNavigate = vi.fn()
const mockSetUser = vi.fn()
let updateProfileMock = vi.fn()
vi.mock("react-router-dom", async () => {
const actual = await vi.importActual("react-router-dom")
return { ...actual, useNavigate: () => mockNavigate }
})
let authState: Record<string, unknown> = {}
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(
<QueryClientProvider client={queryClient}>
<MemoryRouter>
<WechatOnboarding />
</MemoryRouter>
</QueryClientProvider>,
)
describe("WechatOnboarding 昵称引导页", () => {
afterEach(() => {
cleanup()
})
beforeEach(() => {
vi.clearAllMocks()
authState = {
isAuthenticated: true,
user: { id: "u1", display_name: "", profile_completed: false },
setUser: mockSetUser,
}
localStorage.setItem("access_token", "at")
updateProfileMock = vi.fn(async (data: { display_name: string }) => ({
id: "u1",
display_name: data.display_name,
profile_completed: true,
}))
})
it("未登录时跳转登录页", () => {
authState = {
isAuthenticated: false,
user: null,
setUser: mockSetUser,
}
localStorage.removeItem("access_token")
renderPage()
expect(mockNavigate).not.toHaveBeenCalled()
// Navigate 组件渲染即生效;这里断言页面不含昵称表单
expect(screen.queryByText("进入小虾智剪")).toBeNull()
})
it("资料已完善的用户跳 dashboard", () => {
authState = {
isAuthenticated: true,
user: { id: "u1", display_name: "已起名", profile_completed: true },
setUser: mockSetUser,
}
renderPage()
expect(screen.queryByText("进入小虾智剪")).toBeNull()
})
it("昵称输入框不预填,必须用户自己输入", () => {
renderPage()
expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy()
expect((screen.getByPlaceholderText("请输入您的昵称") as HTMLInputElement).value).toBe("")
})
it("新用户可见昵称表单并能提交", async () => {
renderPage()
expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy()
fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
target: { value: "小虾用户" },
})
fireEvent.click(screen.getByText("进入小虾智剪"))
await waitFor(() => {
expect(updateProfileMock).toHaveBeenCalledWith({ display_name: "小虾用户" })
})
await waitFor(() => {
expect(mockSetUser).toHaveBeenCalled()
expect(mockNavigate).toHaveBeenCalledWith("/app/dashboard", { replace: true })
})
})
it("昵称为空时不允许提交(表单校验)", async () => {
renderPage()
fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
target: { value: " " },
})
fireEvent.click(screen.getByText("进入小虾智剪"))
// 等待表单校验
await waitFor(
() => {
expect(updateProfileMock).not.toHaveBeenCalled()
},
{ timeout: 1000 },
)
})
it("连点提交按钮只触发一次请求(防重复提交)", async () => {
// mutation 挂起不立即完成,模拟慢网络下连续双击
let resolveSubmit: (v: unknown) => void = () => {}
updateProfileMock = vi.fn(
() =>
new Promise((resolve) => {
resolveSubmit = resolve
}),
)
renderPage()
fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
target: { value: "小虾用户" },
})
const btn = screen.getByText("进入小虾智剪")
fireEvent.click(btn)
// 第一次点击后立即再点(此时重渲染/loading 可能还没生效)
fireEvent.click(btn)
fireEvent.click(btn)
await waitFor(() => {
expect(updateProfileMock).toHaveBeenCalledTimes(1)
})
// 释放挂起的 Promise,避免泄漏
resolveSubmit({ id: "u1", display_name: "小虾用户", profile_completed: true })
})
it("提交失败后守卫复位,允许再次提交", async () => {
updateProfileMock = vi
.fn()
.mockRejectedValueOnce({
isAxiosError: true,
response: { status: 500, data: { detail: "服务内部错误" } },
})
.mockResolvedValueOnce({ id: "u1", display_name: "小虾用户", profile_completed: true })
renderPage()
fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
target: { value: "小虾用户" },
})
fireEvent.click(screen.getByText("进入小虾智剪"))
await waitFor(() => {
expect(updateProfileMock).toHaveBeenCalledTimes(1)
})
// 失败后再点一次,应能重新提交
fireEvent.click(screen.getByText("进入小虾智剪"))
await waitFor(() => {
expect(updateProfileMock).toHaveBeenCalledTimes(2)
})
})
it("提交失败显示错误且不跳转", async () => {
updateProfileMock = vi.fn(async () => {
throw new Error("500")
})
renderPage()
fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
target: { value: "小虾用户" },
})
fireEvent.click(screen.getByText("进入小虾智剪"))
await waitFor(() => {
expect(mockNavigate).not.toHaveBeenCalled()
})
})
})
@@ -19,6 +19,10 @@ import "@/pages/generate/GeneratePage"
import "@/pages/generate/components/Step2MaterialSelect"
import "@/pages/generate/components/Step4TitleSettings"
import "@/pages/generate/components/Step5VoiceSelect"
import "@/pages/generate/components/Step3VoiceWithMode"
import "@/pages/generate/components/CanvasPreviewGrid"
import "@/pages/generate/components/BatchGenerationGrid"
import "@/pages/generate/components/PreviewCountModal"
import "@/pages/generate/components/PreviewVideoPanel"
import "@/pages/generate/components/GenerateResultPanel"
import "@/pages/generate/components/GenerateStepContent"
@@ -45,6 +49,7 @@ describe("GeneratePage module smoke test", () => {
})
})
import "@/pages/generate/hooks/useGenerateVideo"
import "@/pages/generate/hooks/useBatchCovers"
import "@/pages/generate/hooks/usePreviewAssets"
import "@/pages/generate/hooks/useSegmentScheduler"
import "@/pages/generate/hooks/generate-video/useGenerationPolling"
+461 -249
View File
@@ -31,25 +31,65 @@ SCENE_CHANGE_THRESHOLD = 30 # 灰度差异阈值
MIN_KEYFRAME_INTERVAL_SEC = 1.0 # 最小关键帧间隔(秒)
MAX_KEYFRAMES = 30 # 最大关键帧数
MIN_KEYFRAMES = 5 # 最小关键帧数
FINGERPRINT_SAMPLE_INTERVAL_SEC = 1.0 # 指纹采样间隔(秒):密集均匀采样,保证两视频时序可对齐
FINGERPRINT_MAX_SAMPLES = 30 # 长视频采样数上限(超过后采样间隔自动放宽)
LONG_VIDEO_SEGMENT_SEC = 30 # 长视频每段秒数
LONG_VIDEO_DURATION_THRESHOLD_SEC = 180 # 3 分钟阈值
MIN_FRAMES_PER_SEGMENT = 2 # 长视频每段最少帧数
# ── 滑动窗口匹配常量 ────────────────────────────────────────────
SEGMENT_MATCH_THRESHOLD = 8 # 帧匹配汉明距离阈值
MIN_CONSECUTIVE_MATCHES = 5 # 最少连续匹配帧数
# ── 滑动窗口匹配常量Issue #1702 二次校准) ─────────────────────
# 阈值经 staging 真实数据两轮回归校准(worker 容器内离线实验):
# 第一轮(2026-09-05):同源对 <=12 命中 4/11,异源最小距离 24 → 定 12;
# 第二轮(2026-09-05,证据视频 B->A 仍漏检):扩大样本到该用户全部
# 15 个真实成片(13 个异源候选)实测:
# - 同源成片对(A 20s / B、C 各 11.75s1s 密集采样):
# B->A 中位数距离 14<=16 命中 8/11=0.73C->A 8/11=0.73
# - 异源成片对(13 个真实视频):每帧全局最近邻最小距离 18,
# <=16 命中帧数全部为 0(最近邻 18 仅个别帧,中位数 22~28)
# 12 漏掉同源降重对(降重滤镜/字幕/画面扰动把距离从 ~8 推到 14~16);
# 16 对同源命中 0.73+ 且与异源分布(最近邻 >=18)仍有 >=2bit 安全裕度,
# 异源 <=16 命中 0 帧,无误报空间。
PHASH_THRESHOLD = 16
SEGMENT_MATCH_THRESHOLD = PHASH_THRESHOLD # 片段匹配阈值与帧匹配统一(#1702:阈值常量统一来源)
MIN_CONSECUTIVE_MATCHES = 5 # 连续匹配默认门槛;短视频自适应 min(5, max(2, 分片数//2))
MAX_GAP = 2 # 允许的最大间隙帧数
NEIGHBOR_WINDOW = 1 # 分片时序对齐:允许 ±1 邻接偏移(1s 密集采样下即 ±1s,缓解切点不一致)
# ── 融合判定常量 ────────────────────────────────────────────────
PHASH_WEIGHT = 0.7 # pHash 权重
HISTOGRAM_WEIGHT = 0.3 # 直方图权重
MATCH_RATIO_THRESHOLD = 0.7 # 至少 70% 帧匹配
MATCH_RATIO_THRESHOLD = 0.7 # 全片重复(is_duplicate至少 70% 帧匹配
PARTIAL_COVERAGE_THRESHOLD = 0.5 # 局部复用覆盖率 >=50% 也判全片重复
DUPLICATE_THRESHOLD = 0.70 # 融合后相似度阈值
# ── 降重裁剪规避常量(Issue #1702) ─────────────────────────────
# 成片强制 2-5% random_edge_crop 降重只服务外部平台;自查重指纹取中心 90%
# 区域,使两次不同裁剪的同源画面 pHash 距离回到同分布。
FINGERPRINT_CENTER_CROP_RATIO = 0.90
# ── 感知哈希 & 颜色直方图工具函数 ────────────────────────────────
def center_crop_frame(image: np.ndarray, ratio: float = FINGERPRINT_CENTER_CROP_RATIO) -> np.ndarray:
"""取画面中心 ratio 比例区域(裁除四边边缘)。
查重指纹用:random_edge_crop 降重(2-5% 四边随机裁剪)会让同源画面 pHash
位翻转 12-16,污染自查重(Issue #1702)。算 pHash/颜色直方图前先居中裁除
边缘 10%,两次不同裁剪的同源画面中心区域基本重合,指纹不再被降重污染。
降重只服务外部平台,不影响内部查重。
"""
if image is None or image.size == 0:
return image
h, w = image.shape[:2]
ch, cw = int(h * ratio), int(w * ratio)
if ch <= 0 or cw <= 0 or (ch >= h and cw >= w):
return image
y0 = (h - ch) // 2
x0 = (w - cw) // 2
return image[y0 : y0 + ch, x0 : x0 + cw]
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
"""计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。
@@ -101,11 +141,17 @@ def hamming_distance(hash1: str, hash2: str) -> int:
def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
"""Compute color histogram for an image."""
"""Compute BGR color histogram for an image.
Issue #1702: 每个通道独立做 NORM_L1 归一化(通道内 Σ=1,是概率分布),
三通道拼接存储。Bhattacharyya 系数对拼接向量直接 Σ√(a*b) 会得到
3 通道之和(范围 [0,3],实测 ~14.9 是旧 L2 归一化的错误结果),
消费方 _bhattacharyya_coefficient 按通道数平均归一到 [0,1]。
"""
hist = []
for i in range(3):
h = cv2.calcHist([image], [i], None, [bins], [0, 256])
h = cv2.normalize(h, h).flatten()
h = cv2.normalize(h, h, norm_type=cv2.NORM_L1).flatten()
hist.extend(h)
return hist
@@ -210,6 +256,30 @@ def detect_keyframe_timestamps(
return keyframe_times
def sample_fingerprint_timestamps(
duration: float,
*,
interval_sec: float = FINGERPRINT_SAMPLE_INTERVAL_SEC,
max_samples: int = FINGERPRINT_MAX_SAMPLES,
) -> list[float]:
"""指纹采样时间戳:固定间隔密集均匀采样(Issue #1702)。
动态场景检测抽帧(#1659)在两个同源视频上会各自取到不同时刻,切点/取帧
错位让对齐帧的 pHash 距离都很大(实测同源对最小距离 12 且配对时序错乱)。
改为固定 1s 间隔均匀采样后,复用片段的帧时刻天然对齐,配合 ±1 邻接窗口
即可检出同源/局部复用。长视频(>max_samples*interval)自动放宽间隔到
duration/max_samples,保证分片数有上限。
"""
if duration <= 0:
return []
step = interval_sec
n_uniform = int(duration / step)
if n_uniform > max_samples:
step = duration / max_samples
count = max(1, int(duration / step))
return [step * (i + 0.5) for i in range(count)]
# ── 数据类 ──────────────────────────────────────────────────────
@@ -297,23 +367,33 @@ def find_duplicate_segments(
target_chunks: list,
*,
match_threshold: int = SEGMENT_MATCH_THRESHOLD,
min_consecutive: int = MIN_CONSECUTIVE_MATCHES,
min_consecutive: Optional[int] = None,
max_gap: int = MAX_GAP,
neighbor_window: int = NEIGHBOR_WINDOW,
) -> list[DuplicateSegment]:
"""滑动窗口时序匹配:找出两组分片之间的重复片段。
"""滑动窗口时序匹配:找出两组分片之间的重复片段Issue #1702 重构)
算法:
1. 对每个 query chunk,找到 target 汉明距离最小的 chunk
2. 距离 <= match_threshold 视为匹配
3. 找连续匹配的 run(允许 max_gap 帧间隙)
4. 连续匹配数 >= min_consecutive 的 run 报告为重复片段
1. 构建 query×target 全量汉明距离矩阵;每个 query chunk 保留所有
距离 <= match_threshold 的候选 target 分片(与帧匹配判定同一阈值)。
2. 时序一致贪心对齐:沿 query 时序推进,run 内优先选择与上一匹配帧
目标序号连贯(|delta| <= neighbor_window+1,允许 ±1 邻接/时序偏移
对齐——1s 密集采样下相邻帧 pHash 接近,最近邻在目标相邻帧间
正/反向跳变均属正常,缓解场景切割切点、取帧错位、局部倒退)的
候选;同距时偏好小索引(最早对齐位置)。
3. 连贯匹配中允许 <= max_gap 帧间隙桥接;断裂后另起新 run——天然
支持局部片段复用(复用片段可出现在任意时序位置,各成独立片段)。
4. 连续匹配帧数 >= min_consecutive 的 run 报为重复片段。短视频自适应:
min_consecutive = min(5, max(2, len(query_chunks)//2))n=1 时
不形成片段,由调用方匹配帧回退兜底。
Args:
query_chunks: 查询视频的分片列表(FingerprintChunk 或 dict
target_chunks: 目标视频的分片列表
match_threshold: 汉明距离匹配阈值
min_consecutive: 最少连续匹配帧数
match_threshold: 汉明距离匹配阈值(统一常量 PHASH_THRESHOLD
min_consecutive: 最少连续匹配帧数None 时按短视频自适应
max_gap: 允许的最大间隙帧数
neighbor_window: 时序对齐允许的目标分片序号邻接窗口(正/反向均允许)
Returns:
DuplicateSegment 列表
@@ -321,95 +401,94 @@ def find_duplicate_segments(
if not query_chunks or not target_chunks:
return []
def _get_phash(chunk) -> str:
def _get(chunk, key):
if isinstance(chunk, dict):
return chunk["phash_binary"]
return chunk.phash_binary
return chunk[key]
return getattr(chunk, key)
def _get_start(chunk) -> int:
if isinstance(chunk, dict):
return chunk["start_time_ms"]
return chunk.start_time_ms
n, m = len(query_chunks), len(target_chunks)
q_ph = [_get(c, "phash_binary") for c in query_chunks]
t_ph = [_get(c, "phash_binary") for c in target_chunks]
def _get_end(chunk) -> int:
if isinstance(chunk, dict):
return chunk["end_time_ms"]
return chunk.end_time_ms
# Step 1: 全量距离矩阵。每个 query chunk 保留所有 <= 阈值的候选 target
# 按距离升序;同距时小索引优先(取最早的对齐位置,贪心连贯推进时最保守,
# 不会越过复用片段末端;重复 hash 的连续帧由 Step 2 的连贯性窗口约束)。
candidates: list[list[tuple[int, int]]] = [] # 每 query 帧: [(target_idx, dist), ...]
for i in range(n):
dists = [hamming_distance(q_ph[i], t_ph[j]) for j in range(m)]
cand = [(j, d) for j, d in enumerate(dists) if d <= match_threshold]
cand.sort(key=lambda x: (x[1], x[0]))
candidates.append(cand)
# Step 1: 逐帧匹配
frame_matches: list[tuple[bool, int, int]] = [] # (is_match, min_dist, best_target_idx)
for qc in query_chunks:
qc_phash = _get_phash(qc)
best_dist = 64
best_idx = 0
for j, tc in enumerate(target_chunks):
d = hamming_distance(qc_phash, _get_phash(tc))
if d < best_dist:
best_dist = d
best_idx = j
frame_matches.append((best_dist <= match_threshold, best_dist, best_idx))
# 短视频自适应连续匹配门槛(Issue #1702 工单公式):
# MIN_CONSECUTIVE_MATCHES = min(5, max(2, 分片数//2))。
# n=1 时门槛为 2 不形成片段,由 _evaluate_candidate 的匹配帧回退
# temporal_coverage 按匹配帧占比估计)兜底检出,不回归。
if min_consecutive is None:
min_consecutive = min(MIN_CONSECUTIVE_MATCHES, max(2, n // 2))
# Step 2: 找连续匹配的 runs
runs: list[tuple[int, int]] = [] # list of (start_idx, end_idx)
run_start = None
# Step 2: 时序一致贪心对齐。
# run 内偏好与上一匹配帧目标序号连贯(|delta| <= neighbor_window+1
# 支持 ±1 邻接窗口/时序偏移对齐,正反向抖动均允许)的候选;
# 无连贯候选时关闭旧 run。
# 这天然支持局部片段复用:同一 query 视频中多个复用片段各自形成独立 run。
frame_matches: list[tuple[bool, int, int]] = []
runs: list[tuple[int, int]] = []
run_start: Optional[int] = None
run_last_t: Optional[int] = None
gap_count = 0
for i, (is_match, _dist, _idx) in enumerate(frame_matches):
if is_match:
def _matching_count(a: int, b: int) -> int:
return sum(1 for k in range(a, b + 1) if frame_matches[k][0])
def _close_run(a: int, b: int) -> None:
if b >= a and _matching_count(a, b) >= min_consecutive:
runs.append((a, b))
for i in range(n):
cand = candidates[i]
if run_last_t is None:
chosen = cand[0] if cand else None
else:
chosen = next(
(c for c in cand if abs(c[0] - run_last_t) <= neighbor_window + 1),
None,
)
if chosen is not None:
tidx, dist = chosen
frame_matches.append((True, dist, tidx))
if run_start is None:
run_start = i
gap_count = 0 # 重置间隙
gap_count = 0
run_last_t = tidx
else:
frame_matches.append((False, match_threshold + 1, -1))
if run_start is not None:
gap_count += 1
if gap_count > max_gap:
# 中断当前 run
run_end = i - gap_count # 最后一个匹配帧的索引
# 计算 run 内的实际匹配帧数(总跨度 - 间隙数)
total_gaps = sum(1 for k in range(run_start, run_end + 1) if not frame_matches[k][0])
matching_count = (run_end - run_start + 1) - total_gaps
if matching_count >= min_consecutive:
runs.append((run_start, run_end))
run_start = None
gap_count = 0
# 非匹配帧从 i-gap_count+1 开始,run 结束于其前一帧
_close_run(run_start, i - gap_count)
run_start, run_last_t, gap_count = None, None, 0
# 处理末尾 run
if run_start is not None:
last_idx = len(frame_matches) - 1
# 回退找到最后一个匹配帧的位置(跳过尾部非匹配帧)
last_idx = n - 1
while last_idx >= run_start and not frame_matches[last_idx][0]:
last_idx -= 1
if last_idx >= run_start:
# 计算 run 内的总间隙数
total_gaps = sum(1 for k in range(run_start, last_idx + 1) if not frame_matches[k][0])
matching_count = (last_idx - run_start + 1) - total_gaps
if matching_count >= min_consecutive:
runs.append((run_start, last_idx))
_close_run(run_start, last_idx)
# Step 3: 构建 DuplicateSegment
segments: list[DuplicateSegment] = []
for start, end in runs:
query_start = _get_start(query_chunks[start])
query_end = _get_end(query_chunks[end])
# 取目标范围(按最佳匹配的目标 chunk 时间范围)
target_indices = [frame_matches[k][2] for k in range(start, end + 1) if frame_matches[k][0]]
if target_indices:
t_min = min(target_indices)
t_max = max(target_indices)
target_start = _get_start(target_chunks[t_min])
target_end = _get_end(target_chunks[t_max])
else:
target_start = _get_start(target_chunks[0])
target_end = _get_end(target_chunks[-1])
avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1)) / (end - start + 1)
t_min, t_max = min(target_indices), max(target_indices)
avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1) if frame_matches[k][0]) / len(target_indices)
segments.append(
DuplicateSegment(
query_start_ms=query_start,
query_end_ms=query_end,
target_start_ms=target_start,
target_end_ms=target_end,
query_start_ms=_get(query_chunks[start], "start_time_ms"),
query_end_ms=_get(query_chunks[end], "end_time_ms"),
target_start_ms=_get(target_chunks[t_min], "start_time_ms"),
target_end_ms=_get(target_chunks[t_max], "end_time_ms"),
avg_distance=avg_dist,
)
)
@@ -423,7 +502,9 @@ def find_duplicate_segments(
class VideoDeduplicator:
"""Video deduplication using multiple fingerprint methods."""
PHASH_THRESHOLD = 8 # Issue #1658: pHash 汉明距离阈值由 10 收紧到 8,降低不同视频误判率
# Issue #1702: 阈值统一来源为模块常量 PHASH_THRESHOLD#1658 曾收紧到 8
# 后经 staging 真实同源/异源指纹分布重新校准,见 test_phash_threshold_calibration_1702)。
PHASH_THRESHOLD = PHASH_THRESHOLD
HISTOGRAM_THRESHOLD = 0.85
@staticmethod
@@ -447,27 +528,36 @@ class VideoDeduplicator:
# 单帧不视为坏指纹(短视频或抽帧不足)
if len(phashes) == 1:
return False
# 多帧但所有 phash 完全相同 → 黑屏/纯色视频
# Issue #1702: 旧逻辑"所有 phash 完全相同即判黑屏"会误杀短视频——
# 11s 视频只有几个不同镜头时,相邻 1s 采样帧可能 phash 完全一致(内容
# 连续但非黑屏)。黑屏的特征是「大量帧全部无内容」,要求至少 8 帧
# 且相同帧占比 >=80% 才判坏;短视频(<8 帧)只有真正单值时交给
# _bhattacharyya/融合分兜底,不因"帧都一样"直接跳过。
if len(phashes) < 8:
return False
unique = set(phashes)
if len(unique) == 1:
same_ratio = sum(1 for x in phashes if x == phashes[0]) / len(phashes)
if len(unique) == 1 and same_ratio >= 0.8:
return True
# 多帧但所有 phash 之间的汉明距离都极小(<3)→ 近似黑屏
# 多帧但所有唯一 phash 之间的汉明距离都极小(<3)且占比 >=80% → 近似黑屏
phash_list = list(unique)
if len(phash_list) >= 2:
all_distances = []
for i in range(len(phash_list)):
for j in range(i + 1, len(phash_list)):
all_distances.append(hamming_distance(phash_list[i], phash_list[j]))
if len(phash_list) >= 2 and same_ratio >= 0.8:
all_distances = [
hamming_distance(phash_list[i], phash_list[j])
for i in range(len(phash_list))
for j in range(i + 1, len(phash_list))
]
if all_distances and max(all_distances) < 3:
return True
return False
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
"""Compute video fingerprint using dynamic keyframe detection.
"""Compute video fingerprint using dense uniform sampling.
使用 detect_keyframe_timestamps() 检测内容感知关键帧,
在每个关键帧处取帧计算 pHash + color_histogram。
同时保留 MD5 计算和分片数据结构。
Issue #1702: 使用 sample_fingerprint_timestamps() 固定 1s 间隔密集均匀
采样(替代动态场景检测抽帧),保证两个同源视频复用片段的帧时刻天然
对齐;每帧取中心 90% 区域(center_crop_frame)计算 pHash + color_histogram
绕开 random_edge_crop 降重裁剪污染;MD5 仍基于原始帧。
"""
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
@@ -481,8 +571,8 @@ class VideoDeduplicator:
cap.release()
# 1. 检测关键帧时间戳
keyframe_times = detect_keyframe_timestamps(video_path)
# 1. 固定间隔密集采样(Issue #1702:替代动态场景检测,保证跨视频时序对齐)
keyframe_times = sample_fingerprint_timestamps(duration)
if not keyframe_times:
return VideoFingerprint(
@@ -506,12 +596,15 @@ class VideoDeduplicator:
if not ret:
continue
# MD5 计算
# MD5 计算(基于原始帧,指纹文件级去重不受裁剪影响)
_, buffer = cv2.imencode(".jpg", frame)
md5_hash.update(buffer)
phash = compute_phash(frame)
hist = compute_color_histogram(frame)
# Issue #1702: pHash / 颜色直方图基于中心 90% 区域,绕开 random_edge_crop
# 降重裁剪对指纹的污染(降重只服务外部平台,不污染自查重)。
fp_frame = center_crop_frame(frame)
phash = compute_phash(fp_frame)
hist = compute_color_histogram(fp_frame)
# 计算分片时间范围(从前一个关键帧到下一个关键帧的中点)
prev_boundary = keyframe_times[i - 1] * 1000 if i > 0 else 0
@@ -564,12 +657,22 @@ class VideoDeduplicator:
@staticmethod
def _bhattacharyya_coefficient(hist_a: list[float], hist_b: list[float]) -> float:
"""Bhattacharyya 系数:Σ √(a[i] * b[i]),范围 [0, 1]1=完全相同。"""
"""Bhattacharyya 系数(概率分布版,范围 [0,1]1=完全相同
Issue #1702: compute_color_histogram 输出 3 通道拼接、每通道独立 NORM_L1
(单通道 Σ=1,三通道拼接向量 Σ=3)。旧实现直接 Σ√(a*b) 对三通道拼接向量
算出 ~3(旧 L2 归一化更是算出 ~14.9),不是合法的概率系数。
这里按两个直方图各自的总量归一:BC = Σ√(a*b) / √(Σa·Σb)。
- 单通道概率分布(Σa=Σb=1):分母 1,与旧测试/教科书定义一致;
- 三通道拼接(Σa=Σb=3):分母 3,结果在 [0,1]。
"""
min_len = min(len(hist_a), len(hist_b))
a = hist_a[:min_len]
b = hist_b[:min_len]
# 纯标准库计算(不依赖 numpy);max(0.0, ...) 防御上游异常负值导致 sqrt domain error
return float(sum(math.sqrt(max(0.0, ai * bi)) for ai, bi in zip(a, b, strict=False)))
a = [max(0.0, float(x)) for x in hist_a[:min_len]]
b = [max(0.0, float(x)) for x in hist_b[:min_len]]
# max(0.0, ...) 防御上游异常负值导致 sqrt domain error
coeff = sum(math.sqrt(ai * bi) for ai, bi in zip(a, b, strict=False))
norm = math.sqrt(sum(a) * sum(b))
return float(coeff / norm) if norm > 0 else 0.0
@staticmethod
def _compute_histogram_similarity(
@@ -611,6 +714,72 @@ class VideoDeduplicator:
hist_similarity = VideoDeduplicator._compute_histogram_similarity(hist_a, hist_b) if hist_b else 0.5
return PHASH_WEIGHT * phash_similarity + HISTOGRAM_WEIGHT * hist_similarity
@staticmethod
def _evaluate_candidate(
fingerprint: VideoFingerprint,
existing_phashes: list[str],
existing_histograms: list,
existing_chunk_objects: list,
*,
query_duration_sec: float,
) -> dict:
"""评估新视频指纹与单个候选视频的相似度(Issue #1702 共享逻辑)。
指标:
- min_distances / frame_match_rate:每个新分片到候选视频全局最近邻的汉明距离,
分母取两视频分片数的较小值(支持局部片段复用:短视频复用长视频片段时不被长视频分母稀释)。
- temporal_coverage:时序一致连续匹配片段总时长 / 新视频时长(局部复用主指标)。
- fusionpHash 中位数距离 + 颜色直方图的加权融合分。
Returns:
{frame_match_rate, temporal_coverage, segments, median_distance,
fusion, matching_frames, min_distances}
"""
query_phashes = fingerprint.keyframe_phashes or []
if not query_phashes or not existing_phashes:
return {
"frame_match_rate": 0.0,
"temporal_coverage": 0.0,
"segments": [],
"median_distance": 64,
"fusion": 0.0,
"matching_frames": 0,
"min_distances": [],
}
min_distances = [min(hamming_distance(ph, ep) for ep in existing_phashes) for ph in query_phashes]
matching_frames = sum(1 for d in min_distances if d <= PHASH_THRESHOLD)
# 分母取 min(两视频分片数):局部复用时(如 B 的 5 片复用 A 9 片中的若干片)
# 命中帧占比不因候选视频更长而被稀释。
frame_match_rate = matching_frames / min(len(query_phashes), len(existing_phashes))
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
duration_ms = query_duration_sec * 1000 if query_duration_sec else 0
if duration_ms > 0 and segments:
covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments)
temporal_coverage = min(covered_ms / duration_ms, 1.0)
elif matching_frames > 0:
# 无连续片段(时序连贯性不足)时,按匹配帧占比估计覆盖:
# 密集 1s 采样下每个分片≈1s 等权时间片,匹配帧数≈命中秒数。
temporal_coverage = min(frame_match_rate, 1.0)
else:
temporal_coverage = 0.0
median_distance = statistics.median(min_distances) if min_distances else 64
fusion = VideoDeduplicator._compute_fusion_score(
median_distance, fingerprint.color_histograms, existing_histograms
)
return {
"frame_match_rate": frame_match_rate,
"temporal_coverage": temporal_coverage,
"segments": segments,
"median_distance": median_distance,
"fusion": fusion,
"matching_frames": matching_frames,
"min_distances": min_distances,
}
def check_duplicate(
self,
fingerprint: VideoFingerprint,
@@ -620,6 +789,7 @@ class VideoDeduplicator:
scope: str = "project",
user_id: str = "",
duration_sec: float = 0,
exclude_video_id: str | None = None,
) -> Optional[dict]:
"""检查视频是否与已有视频重复。
@@ -636,6 +806,9 @@ class VideoDeduplicator:
scope: "project" 项目内查重(默认),"user" 跨项目全局查重
user_id: 用户 IDscope="user" 时使用)
duration_sec: 视频时长(秒),用于时长预过滤 ±15%
exclude_video_id: 排除的视频 ID(查重自身时用)。recompute-dedup
重算时视频记录已存在,不排除会自匹配(距离 0 分最高)导致
duplicate_of 指向自己(Issue #1702 连带修复)。
Returns:
重复信息字典(含 duplicate, duplicate_of, reason, similarity, duplicate_segments),
@@ -643,13 +816,22 @@ class VideoDeduplicator:
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
if scope == "user" and user_id:
dur_min = duration_sec * 0.85 if duration_sec > 0 else 0
dur_max = duration_sec * 1.15 if duration_sec > 0 else 0
existing_videos = video_repo.list_by_user(user_id, duration_min=dur_min, duration_max=dur_max)
# Issue #1702: 不做 ±15% 时长预过滤。旧逻辑按 duration_sec 缩小候选窗口,
# 但局部片段复用的两个视频时长必然不同(证据视频 20s vs 11s,差 42%),
# ±15% 窗口让同源视频互相不可见 → is_duplicate 恒 False。
# 全量遍历同用户视频(与 compute_duplicate_rate 口径一致),异源视频由
# fusion/temporal_coverage 阈值天然过滤(校准:异源最小汉明距离 24)。
existing_videos = video_repo.list_by_user(user_id)
else:
existing_videos = video_repo.list_by_project(project_id)
best_score = 0.0
best_result: Optional[dict] = None
for existing in existing_videos:
# 排除自身(recompute 时当前视频已在候选列表里,否则自匹配距离 0 必最高分)
if exclude_video_id and existing.id == exclude_video_id:
continue
if not existing.video_fingerprint:
continue
@@ -677,61 +859,70 @@ class VideoDeduplicator:
if not existing_phashes:
continue
# 计算每个新关键帧到已有关键帧的最小汉明距离
min_distances = []
for phash in fingerprint.keyframe_phashes:
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
min_distances.append(min(distances))
# 帧匹配比例检查
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
match_ratio = matching_frames / len(min_distances) if min_distances else 0
if match_ratio < MATCH_RATIO_THRESHOLD:
continue
# 中位数距离
median_distance = statistics.median(min_distances) if min_distances else 64
if median_distance >= self.PHASH_THRESHOLD:
continue
# 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
# 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
existing_chunk_objects = chunk_data
else:
existing_histograms = ef.get("color_histograms") or []
existing_chunk_objects = [
{"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
]
combined_score = self._compute_fusion_score(
median_distance, fingerprint.color_histograms, existing_histograms
# Issue #1702: 统一评估每个候选(含局部片段复用),不再用
# "frame_match_rate<0.7 整条跳过" 的硬门槛——局部复用(如 B 结尾 2s
# ≈ A 中间 2s)帧比例天然低,但 coverage 能检出。
ev = self._evaluate_candidate(
fingerprint,
existing_phashes,
existing_histograms,
existing_chunk_objects,
query_duration_sec=fingerprint.duration,
)
logger.debug(
"check_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f "
"temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
existing.id,
ev["min_distances"],
ev["frame_match_rate"],
ev["temporal_coverage"],
ev["median_distance"],
ev["fusion"],
len(ev["segments"]),
)
if combined_score < DUPLICATE_THRESHOLD:
continue
# 滑动窗口时序匹配:获取具体重复片段
existing_chunk_objects = (
chunk_data
if chunk_data
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
# 全片重复判定:融合分过阈 且(帧匹配比例 >=70% 或 局部覆盖 >=50%
is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and (
ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
)
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "phash_histogram_fusion",
"similarity": combined_score,
"duplicate_segments": [
{
"query_start_ms": s.query_start_ms,
"query_end_ms": s.query_end_ms,
"target_start_ms": s.target_start_ms,
"target_end_ms": s.target_end_ms,
"avg_distance": round(s.avg_distance, 2),
}
for s in segments
],
}
if is_full_duplicate and ev["fusion"] > best_score:
best_score = ev["fusion"]
best_result = {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "phash_histogram_fusion",
"similarity": ev["fusion"],
"duplicate_segments": [
{
"query_start_ms": s.query_start_ms,
"query_end_ms": s.query_end_ms,
"target_start_ms": s.target_start_ms,
"target_end_ms": s.target_end_ms,
"avg_distance": round(s.avg_distance, 2),
}
for s in ev["segments"]
],
}
if best_result:
return best_result
logger.info(
"check_duplicate no match (project=%s scope=%s): %d candidates evaluated, best_fusion=%.3f",
project_id,
scope,
len(existing_videos),
best_score,
)
return None
def check_batch_duplicate(
@@ -763,6 +954,9 @@ class VideoDeduplicator:
video_repo = SQLAlchemyGeneratedVideoRepository(session)
batch_videos = video_repo.list_by_batch(batch_id)
best_score = 0.0
best_result: Optional[dict] = None
for existing in batch_videos:
if existing.id == current_video_id:
continue
@@ -796,59 +990,59 @@ class VideoDeduplicator:
if not existing_phashes:
continue
min_distances = []
for phash in fingerprint.keyframe_phashes:
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
min_distances.append(min(distances))
# 帧匹配比例检查
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
match_ratio = matching_frames / len(min_distances) if min_distances else 0
if match_ratio < MATCH_RATIO_THRESHOLD:
continue
median_distance = statistics.median(min_distances) if min_distances else 64
if median_distance >= self.PHASH_THRESHOLD:
continue
# 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
existing_chunk_objects = chunk_data
else:
existing_histograms = ef.get("color_histograms") or []
existing_chunk_objects = [
{"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
]
combined_score = self._compute_fusion_score(
median_distance, fingerprint.color_histograms, existing_histograms
ev = self._evaluate_candidate(
fingerprint,
existing_phashes,
existing_histograms,
existing_chunk_objects,
query_duration_sec=fingerprint.duration,
)
logger.debug(
"check_batch_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f "
"temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
existing.id,
ev["min_distances"],
ev["frame_match_rate"],
ev["temporal_coverage"],
ev["median_distance"],
ev["fusion"],
len(ev["segments"]),
)
if combined_score < DUPLICATE_THRESHOLD:
continue
# 滑动窗口时序匹配
existing_chunk_objects = (
chunk_data
if chunk_data
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and (
ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
)
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "batch_phash_histogram_fusion",
"similarity": combined_score,
"duplicate_segments": [
{
"query_start_ms": s.query_start_ms,
"query_end_ms": s.query_end_ms,
"target_start_ms": s.target_start_ms,
"target_end_ms": s.target_end_ms,
"avg_distance": round(s.avg_distance, 2),
}
for s in segments
],
}
if is_full_duplicate and ev["fusion"] > best_score:
best_score = ev["fusion"]
best_result = {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "batch_phash_histogram_fusion",
"similarity": ev["fusion"],
"duplicate_segments": [
{
"query_start_ms": s.query_start_ms,
"query_end_ms": s.query_end_ms,
"target_start_ms": s.target_start_ms,
"target_end_ms": s.target_end_ms,
"avg_distance": round(s.avg_distance, 2),
}
for s in ev["segments"]
],
}
if best_result:
return best_result
logger.info("check_batch_duplicate no match (batch=%s): best_fusion=%.3f", batch_id, best_score)
return None
def compute_duplicate_rate(
@@ -897,8 +1091,7 @@ class VideoDeduplicator:
max_duplicate_rate = 0.0
max_visual_similarity = 0.0
match_count = 0
total_duration_ms = fingerprint.duration if fingerprint.duration else 0
evaluated = 0
for existing in existing_videos:
if current_video_id and existing.id == current_video_id:
@@ -933,57 +1126,63 @@ class VideoDeduplicator:
if not existing_phashes or not fingerprint.keyframe_phashes:
continue
min_distances = []
for phash in fingerprint.keyframe_phashes:
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
min_distances.append(min(distances))
# frame_match_rate
total_frames = len(min_distances)
if total_frames == 0:
continue
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
frame_match_rate = matching_frames / total_frames
# 帧匹配比例太低则跳过
if frame_match_rate < 0.3:
continue
# temporal_coverage_rate via find_duplicate_segments
existing_chunk_objects = (
chunk_data
if chunk_data
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
)
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
if total_duration_ms > 0 and segments:
covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments)
temporal_coverage_rate = min(covered_ms / total_duration_ms, 1.0)
else:
temporal_coverage_rate = 0.0
# duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage_rate
dup_rate = (frame_match_rate * 0.4 + temporal_coverage_rate * 0.6) * 100
# visual_similarity (融合相似度,归一化 0~1)
median_distance = statistics.median(min_distances) if min_distances else 64
# 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
existing_chunk_objects = chunk_data
else:
# JSON NULL 显式回退空列表
existing_histograms = ef.get("color_histograms") or []
existing_chunk_objects = [
{"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
]
visual_sim = self._compute_fusion_score(median_distance, fingerprint.color_histograms, existing_histograms)
# Issue #1702: 统一评估;frame_match_rate 分母为 min(两视频分片数),
# temporal_coverage 时长量纲在 _evaluate_candidate 内统一为毫秒。
ev = self._evaluate_candidate(
fingerprint,
existing_phashes,
existing_histograms,
existing_chunk_objects,
query_duration_sec=fingerprint.duration,
)
evaluated += 1
logger.debug(
"compute_duplicate_rate candidate=%s min_distances=%s frame_match_rate=%.3f "
"temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
existing.id,
ev["min_distances"],
ev["frame_match_rate"],
ev["temporal_coverage"],
ev["median_distance"],
ev["fusion"],
len(ev["segments"]),
)
# 判定是否为重复(融合分数超过阈值)
if visual_sim >= DUPLICATE_THRESHOLD:
# Issue #1702: 去掉 "frame_match_rate<0.3 整条跳过" 硬门槛——
# 局部片段复用帧比例天然低;coverage 为主指标,0 匹配自然得 0 分。
# duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage
dup_rate = (min(ev["frame_match_rate"], 1.0) * 0.4 + ev["temporal_coverage"] * 0.6) * 100
# 全片重复计数与 check_duplicate 判定口径一致
if ev["fusion"] >= DUPLICATE_THRESHOLD and (
ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
):
match_count += 1
if dup_rate > max_duplicate_rate:
max_duplicate_rate = dup_rate
max_visual_similarity = visual_sim
max_visual_similarity = ev["fusion"]
logger.info(
"compute_duplicate_rate done (project=%s scope=%s): evaluated=%d max_rate=%.2f%% "
"max_visual_sim=%.3f matches=%d",
project_id,
scope,
evaluated,
max_duplicate_rate,
max_visual_similarity,
match_count,
)
return {
"duplicate_rate": round(max(max_duplicate_rate, 0.0), 2),
"visual_similarity": round(max_visual_similarity, 4),
@@ -999,18 +1198,20 @@ def _save_fingerprint_chunks(
session: Session,
) -> None:
"""将指纹分片数据批量写入 video_fingerprint_chunks 表。幂等:已有数据时跳过。"""
# 幂等检查:已有分片数据则跳过
existing_count = (
session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count()
)
if existing_count > 0:
logger.debug("Fingerprint chunks already exist for video %s (%d chunks), skipping", video_id, existing_count)
return
if not fingerprint.chunks:
logger.warning("No chunks in fingerprint for video %s, skipping chunk save", video_id)
return
# Issue #1702: recompute-dedup 重算时指纹算法已变(中心裁剪 + 新阈值),
# 旧分片必须替换而非跳过(旧实现"有数据就跳过"导致重算不刷新分片表)。
deleted = (
session.query(VideoFingerprintChunkModel)
.filter(VideoFingerprintChunkModel.video_id == video_id)
.delete(synchronize_session=False)
)
if deleted:
logger.info("Replaced %d stale fingerprint chunks for video %s", deleted, video_id)
chunk_models = fingerprint.to_chunk_models(video_id, project_id, user_id)
session.bulk_save_objects(chunk_models)
logger.info("Saved %d fingerprint chunks for video %s", len(chunk_models), video_id)
@@ -1032,9 +1233,17 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
raise ValueError(f"Generated video {generated_video_id} not found")
local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4")
storage_service.download_file(
f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path
)
# Issue #1702: recompute 走的是 OSS 重新下载路径(正常生成流程用本地渲染文件,
# 不经此任务)。成片真实 OSS key 是生成时的
# generated/projects/{pid}/tasks/{task_id}/rendered_*.mp4(见 generation.py
# _upload_and_record),旧代码硬编码 projects/{pid}/generated/{vid}/{vid}.mp4
# 这个从不存在的 key,导致所有 recompute 任务下载 404、查重数据永远无法重算。
# 优先从 file_url 解析真实 key,旧 key 模式仅作回退。
download_key = getattr(video, "file_url", "") or ""
if not download_key:
download_key = f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4"
logger.warning("video %s has no file_url, falling back to legacy key %s", generated_video_id, download_key)
storage_service.download_file(download_key, local_path)
fingerprint = deduplicator.compute_fingerprint(local_path)
@@ -1045,7 +1254,10 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
session,
scope="user",
user_id=video.user_id,
duration_sec=fingerprint.duration / 1000 if fingerprint.duration else 0,
# Issue #1702: fingerprint.duration 单位已经是秒,旧代码 /1000 导致
# ±15% 时长预过滤窗口缩到 ~0.013sscope=user 的跨项目查重永远返回 None。
duration_sec=fingerprint.duration if fingerprint.duration else 0,
exclude_video_id=generated_video_id,
)
video.video_fingerprint = fingerprint.to_dict()
@@ -92,7 +92,8 @@ def create_video_record_and_dedup(
logger.warning("Failed to save fingerprint chunks for %s: %s", video_id, chunk_err)
# (a) 历史成片查重(跨项目全局 + 时长预过滤)
duration_sec = fingerprint.duration / 1000 if fingerprint.duration else 0
# Issue #1702: fingerprint.duration 单位是秒,旧代码 /1000 让时长预过滤失效
duration_sec = fingerprint.duration if fingerprint.duration else 0
duplicate_result = deduplicator.check_duplicate(
fingerprint,
project_id,
@@ -100,6 +101,7 @@ def create_video_record_and_dedup(
scope="user",
user_id=user_id,
duration_sec=duration_sec,
exclude_video_id=video_id,
)
# (b) 批次内查重(仅当有 batch_id 时)
+22 -2
View File
@@ -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",
@@ -22,10 +34,18 @@ celery_app.conf.imports = (
)
# Celery Beat 定时任务调度
# 注:worker 单实例内嵌 beatentrypoint-worker.sh -B),定时任务不会重复执行
celery_app.conf.beat_schedule = {
# pending 任务超时清理:worker 停止消费后,卡 pending 的任务 15 分钟内释放限流名额
"cleanup-stale-pending-tasks": {
"task": "worker.cleanup_stale_pending_tasks",
"schedule": 600.0, # 每 10 分钟(秒)
"options": {"expires": 300}, # 5 分钟过期,避免堆积
"schedule": 300.0, # 每 5 分钟(秒)
"options": {"expires": 240}, # 4 分钟过期,避免堆积
},
# running 孤儿任务巡检:容器重启/进程被杀后卡 running 的任务,20 分钟无更新则判失败
"cleanup-stale-running-tasks": {
"task": "worker.cleanup_stale_running_tasks",
"schedule": 300.0, # 每 5 分钟(秒)
"options": {"expires": 240},
},
}
+117 -8
View File
@@ -7,11 +7,83 @@ from worker_app.db import SessionLocal
logger = logging.getLogger(__name__)
# 孤儿任务超时阈值:渲染任务超过此时间未更新则视为卡死
ORPHAN_TASK_TIMEOUT_MINUTES = 10
# Pending 任务超时阈值:pending 任务在队列中等待超过此时间则自动清理
PENDING_TASK_TIMEOUT_MINUTES = 30
def cleanup_stale_running_with_session(repo, timeout_minutes: int) -> int:
"""清理超时未更新的 running GenerationTask(可注入 repo 的纯核心,便于单测)。
Returns:
清理的任务数量
"""
return len(cleanup_stale_running_with_session_ids(repo, timeout_minutes))
def cleanup_stale_running_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
"""同 cleanup_stale_running_with_session,返回 [(task_id, celery_task_id), ...]。"""
fn = getattr(repo, "cleanup_stale_running_with_ids", None)
if fn is not None:
return fn(timeout_minutes)
# 旧仓储无 _with_ids 方法:降级为计数,无法撤销消息(执行前状态守卫兜底)
count = repo.cleanup_stale_running(timeout_minutes)
return [("", "") for _ in range(count)]
def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int:
"""清理超时 pending GenerationTask(可注入 repo 的纯核心,便于单测)。
Returns:
清理的任务数量
"""
return len(cleanup_stale_pending_with_session_ids(repo, timeout_minutes))
def cleanup_stale_pending_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
"""同 cleanup_stale_pending_with_session,返回 [(task_id, celery_task_id), ...]。"""
fn = getattr(repo, "cleanup_stale_pending_with_ids", None)
if fn is not None:
return fn(timeout_minutes)
count = repo.cleanup_stale_pending(timeout_minutes)
return [("", "") for _ in range(count)]
def _revoke_and_purge_stale_messages(items: list[tuple[str, str]]) -> int:
"""把清理掉的任务对应的 Celery 消息撤销并从 Redis 队列清除(#1714)。
防止「DB 已标 failed,但队列消息还在 → 重投执行 → 非法状态转换 → 半成品」。
失败不阻断清理流程(执行前状态守卫是第二道防线)。
"""
biz_ids = [tid for tid, _ in items if tid]
celery_ids = [cid for _, cid in items if cid]
if not biz_ids and not celery_ids:
return 0
try:
from worker_app.celery_app import celery_app as app
from worker_app.core.config import get_settings
from packages.shared.celery_orphan_guard import revoke_and_purge
broker_url = get_settings().broker_url
return revoke_and_purge(
app,
broker_url,
business_task_ids=biz_ids,
celery_task_ids=celery_ids,
)
except Exception as e: # noqa: BLE001
logger.error("撤销作废任务队列消息失败(执行前守卫仍会兜底): %s", e, exc_info=True)
return 0
# 孤儿任务超时阈值:running 任务超过此时间无进度更新则视为卡死。
# 依据:worker.generate_video 硬超时 time_limit=11 分钟,正常任务不可能超过;
# 20 分钟阈值覆盖硬超时 + 重试 + 余量,绝不误杀正常任务。
ORPHAN_TASK_TIMEOUT_MINUTES = 20
# Pending 任务超时阈值:任务创建后超过此时间仍未开始执行则判死。
# 注意区分 running 孤儿阈值(20 分钟):pending 是「排队等待」时间,
# 队列积压(如 20+ 转码任务)时视频生成可能正常排队较久,阈值必须放宽,
# 避免正常排队任务被误杀。队列隔离(#1714)后 generation 队列独占 worker
# 理论上排队极短;保留 45 分钟作为兜底,覆盖 worker 短暂停止消费的场景。
PENDING_TASK_TIMEOUT_MINUTES = 45
def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover
@@ -32,11 +104,16 @@ def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) ->
try:
session = SessionLocal()
repo = SQLAlchemyGenerationTaskRepository(session)
count = repo.cleanup_stale_running(timeout_minutes)
session.close()
try:
repo = SQLAlchemyGenerationTaskRepository(session)
items = cleanup_stale_running_with_session_ids(repo, timeout_minutes)
finally:
session.close()
count = len(items)
if count > 0:
logger.warning("清理了 %d 个超时的孤儿 GenerationTask(超过 %d 分钟未更新)", count, timeout_minutes)
purged = _revoke_and_purge_stale_messages(items)
logger.info("孤儿任务对应队列消息撤销/清除完成: %d", purged)
else:
logger.info("无孤儿 GenerationTask 需要清理")
return count
@@ -70,7 +147,9 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
.all()
)
count = 0
stale_items: list[tuple[str, str]] = []
for model in stale_jobs:
stale_items.append((model.id, getattr(model, "celery_task_id", "") or ""))
model.status = JobStatus.FAILED.value
model.error_message = f"任务执行中断(超过 {timeout_minutes} 分钟未更新)"
count += 1
@@ -80,6 +159,8 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
else:
logger.info("无孤儿 Job 需要清理")
session.close()
if count > 0:
_revoke_and_purge_generation(stale_items)
return count
except Exception as e:
logger.error("清理孤儿 Job 失败: %s", e, exc_info=True)
@@ -105,9 +186,12 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
session = SessionLocal()
try:
repo = SQLAlchemyGenerationTaskRepository(session)
count = repo.cleanup_stale_pending(timeout_minutes)
items = cleanup_stale_pending_with_session_ids(repo, timeout_minutes)
count = len(items)
if count > 0:
logger.warning("清理了 %d 个超时的 pending GenerationTask(超过 %d 分钟未处理)", count, timeout_minutes)
purged = _revoke_and_purge_stale_messages(items)
logger.info("超时 pending 任务对应队列消息撤销/清除完成: %d", purged)
else:
logger.info("无超时 pending GenerationTask 需要清理")
return count
@@ -118,6 +202,31 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
session.close()
def _revoke_and_purge_generation(items: list[tuple[str, str]]) -> int:
"""撤销 Job 表孤儿任务(TTS/配音等)的队列消息,队列覆盖全部已知队列。"""
biz_ids = [tid for tid, _ in items if tid]
celery_ids = [cid for _, cid in items if cid]
if not biz_ids and not celery_ids:
return 0
try:
from worker_app.celery_app import celery_app as app
from worker_app.core.config import get_settings
from packages.shared.celery_orphan_guard import revoke_and_purge
broker_url = get_settings().broker_url
return revoke_and_purge(
app,
broker_url,
business_task_ids=biz_ids,
celery_task_ids=celery_ids,
queue_names=("generation", "transcode", "celery"),
)
except Exception as e: # noqa: BLE001
logger.error("撤销 Job 队列消息失败: %s", e, exc_info=True)
return 0
def cleanup_all_stale_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict: # pragma: no cover
"""统一清理所有超时的孤儿任务。
+40 -4
View File
@@ -1,14 +1,18 @@
"""定期清理任务 — Celery Beat 调度。
包含:
- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks
- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasksworker 停止消费时占位)
- cleanup_stale_running_tasks: 定期清理卡在 running 超时的 generation_tasks(容器重启/进程被杀后的孤儿)
"""
import logging
from celery import shared_task
from worker_app.tasks._startup import (
ORPHAN_TASK_TIMEOUT_MINUTES,
PENDING_TASK_TIMEOUT_MINUTES,
cleanup_orphan_tasks,
cleanup_stale_jobs,
cleanup_stale_pending_tasks,
)
@@ -19,12 +23,14 @@ logger = logging.getLogger(__name__)
def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINUTES) -> dict:
"""Celery Beat 调度的定期任务:清理超时的 pending 任务。
10 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
5 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
查找所有 status='pending' 且 created_at < NOW() - timeout_minutes
的 generation_tasks,批量更新为 failed
的 generation_tasks,批量更新为 failed,释放限流名额;同时 revoke 并清除
Redis 队列中对应的 Celery 消息,杜绝作废消息重投执行(#1714)。
Args:
timeout_minutes: 超时时间(分钟),默认 30 分钟
timeout_minutes: 超时时间(分钟),默认 45 分钟pending 排队阈值放宽,
与 running 孤儿 20 分钟区分,避免正常排队任务被误杀)
Returns:
{"cleaned": int}
@@ -33,3 +39,33 @@ def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_
if count > 0:
logger.info("[Beat] 清理了 %d 个超时 pending 任务(超时阈值 %d 分钟)", count, timeout_minutes)
return {"cleaned": count}
@shared_task(name="worker.cleanup_stale_running_tasks")
def scheduled_cleanup_stale_running(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict:
"""Celery Beat 调度的定期任务:清理超时的 running 孤儿任务。
每 5 分钟执行一次。worker_ready 信号只在 worker 启动时清一次,
若 worker 没重启但任务卡死(上传挂起、进程 OOM 被内核杀掉等),
任务会永久卡在 running 占位。此任务做持续兜底:
查找 status='running' 且 updated_at < NOW() - timeout_minutes 的任务,
标记为 failed(原因:容器重启/超时中断),同时清理 Job 表孤儿。
Args:
timeout_minutes: 超时时间(分钟),默认 20 分钟
worker.generate_video 硬超时 11 分钟,正常任务不可能超过 20 分钟)
Returns:
{"generation_tasks": int, "jobs": int}
"""
gen_count = cleanup_orphan_tasks(timeout_minutes)
job_count = cleanup_stale_jobs(timeout_minutes)
total = gen_count + job_count
if total > 0:
logger.warning(
"[Beat] 清理孤儿任务: running GenerationTask=%d, Job=%d(超时阈值 %d 分钟)",
gen_count,
job_count,
timeout_minutes,
)
return {"generation_tasks": gen_count, "jobs": job_count}
+30 -2
View File
@@ -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:
+138 -27
View File
@@ -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/<id>/IMG_2282_h264.MOV
占位 asset 以原始 key uploads/<id>/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 占位 assetIssue #1714)。
HEVC 转码成功后 job.storage_key 会被改写为 *_h264 新 key,旧实现用新 key
回查占位必然落空,进而兜底新建一条 READY 记录,导致原占位永久卡 processing。
查找优先级:
1. job.asset_idcomplete 派单时透传的占位 id,最可靠,不依赖 key);
2. 原始 storage_key(占位记录以原始 key 创建);
3. 当前 job.storage_key(未转码/降级场景与原始 key 相同)。
找不到返回 None(旧链路兼容,由调用方兜底新建并告警)。
"""
asset_id = getattr(job, "asset_id", "") or ""
if asset_id:
try:
found = asset_repo.find_by_id(asset_id)
if found is not None:
return found
except Exception as find_err:
logger.warning("占位 asset 按 id 查询失败 asset_id=%s: %s", asset_id, find_err)
for key in (original_storage_key, getattr(job, "storage_key", "")):
if not key:
continue
try:
found = asset_repo.find_by_storage_key(key)
except Exception:
logger.warning("find_by_storage_key not available, trying fallback lookup")
found = None
if found is not None:
return found
return None
@celery_app.task(name="worker.ingest_asset")
def ingest_asset(job_id: str) -> dict:
"""
@@ -380,6 +428,23 @@ def ingest_asset(job_id: str) -> dict:
if job is None:
return {"status": "failed", "error": "job not found"}
# ── 执行前状态守卫(#1714):job 已终态(失败/完成)说明这是作废消息
# (超时清理标记 failed 后旧消息重投、或重复投递),直接丢弃不执行,
# 避免重复转码、重复回写。processing 是本任务自己第一次置位前的旧消息
# 极少见,保守起见也丢弃(processing 的占位由恢复流程处理)。
current_status = job.status.value if hasattr(job.status, "value") else str(job.status)
if current_status in ("failed", "completed"):
logger.warning(
"[ingest job_id=%s] 任务状态已为 %s,丢弃作废消息,不执行转码",
job_id,
current_status,
)
return {"status": "discarded", "job_id": job_id, "reason": f"job already terminal: {current_status}"}
# 记录原始 storage_keyHEVC 转码成功后 job.storage_key 会改写为 *_h264
# 而 complete 阶段的占位 asset 始终以原始 key 创建,关联回写必须保留它。
original_storage_key = job.storage_key
# Update job status to PROCESSING
job.status = IngestJobStatus.PROCESSING
job.updated_at = datetime.now(timezone.utc)
@@ -624,22 +689,41 @@ def ingest_asset(job_id: str) -> dict:
error_reason,
)
asset = Asset.create(
project_id=job.project_id,
library_id=job.library_id,
name=filename,
storage_key=job.storage_key,
mime_type=mime_type,
metadata={"source": "upload", "ingest_error": error_reason},
file_size=int(metadata.get("size_bytes", 0)),
duration=float(metadata.get("duration", 0)),
width=int(metadata.get("width", 0)),
height=int(metadata.get("height", 0)),
codec=metadata.get("codec") or None,
status=AssetStatus.ERROR,
file_hash=job.file_hash,
)
asset_repo.create(asset)
placeholder = _resolve_placeholder_asset(asset_repo, job, original_storage_key)
if placeholder is not None:
# 回写占位记录:标 ERROR(Issue #1714:禁止新建第二条导致占位孤儿)
asset = placeholder
asset.mime_type = mime_type
asset.metadata = {"source": "upload", "ingest_error": error_reason}
asset.file_size = int(metadata.get("size_bytes", 0))
asset.duration = float(metadata.get("duration", 0)) or None
asset.width = int(metadata.get("width", 0)) or None
asset.height = int(metadata.get("height", 0)) or None
codec_val = metadata.get("codec")
if codec_val:
asset.codec = str(codec_val)
asset.status = AssetStatus.ERROR
asset.updated_at = datetime.now(timezone.utc)
asset_repo.update(asset)
else:
# 旧链路兜底:无占位记录(如历史 job 重跑)才新建
logger.warning("无效素材且未找到占位记录,兜底新建 ERROR asset: job_id=%s", job_id)
asset = Asset.create(
project_id=job.project_id,
library_id=job.library_id,
name=filename,
storage_key=job.storage_key,
mime_type=mime_type,
metadata={"source": "upload", "ingest_error": error_reason},
file_size=int(metadata.get("size_bytes", 0)),
duration=float(metadata.get("duration", 0)),
width=int(metadata.get("width", 0)),
height=int(metadata.get("height", 0)),
codec=metadata.get("codec") or None,
status=AssetStatus.ERROR,
file_hash=job.file_hash,
)
asset_repo.create(asset)
# Update job status to FAILED
job.status = IngestJobStatus.FAILED
@@ -656,16 +740,19 @@ def ingest_asset(job_id: str) -> dict:
"error": error_reason,
}
# 查找已存在的 Asset 记录(由 API 端在上传完成时立即创建为 PROCESSING 状态)
existing_asset = None
try:
existing_asset = asset_repo.find_by_storage_key(job.storage_key)
except Exception:
logger.warning("find_by_storage_key not available, trying fallback lookup")
# 查找 complete 阶段创建的占位 Asset 记录(Issue #1714)。
# 必须用原始 storage_key / job.asset_id 关联——HEVC 转码后 job.storage_key
# 已改写为 *_h264,用新 key 回查占位必然落空,旧实现因此兜底新建 READY 记录,
# 导致原 PROCESSING 占位永久卡住(每个 HEVC 视频产生两条记录)。
existing_asset = _resolve_placeholder_asset(asset_repo, job, original_storage_key)
if existing_asset is None:
# 兜底:如果 API 端没有预先创建 Asset(旧版本兼容),则创建新记录
logger.info("No pre-created asset found for storage_key=%s, creating new", job.storage_key)
# 兜底:仅当确实没有占位记录(旧版本 API / 历史 job 重跑)才新建。
logger.warning(
"No placeholder asset found for job_id=%s original_key=%s, creating new",
job_id,
original_storage_key,
)
metadata["source"] = "upload"
asset = Asset.create(
project_id=job.project_id,
@@ -685,8 +772,13 @@ def ingest_asset(job_id: str) -> dict:
)
asset_repo.create(asset)
else:
# 更新已有的 Asset 记录补充元数据并将状态改为 READY
# 更新占位记录补充元数据、置 READY。转码成功时 storage_key 同步改写为
# *_h264(播放/下载走转码产物),原始 key 记入 metadata 可溯源。
asset = existing_asset
if job.storage_key != asset.storage_key:
metadata["original_storage_key"] = asset.storage_key
metadata["hevc_transcoded"] = True
asset.storage_key = job.storage_key
asset.mime_type = mime_type
metadata["source"] = "upload"
asset.metadata = metadata
@@ -737,9 +829,28 @@ def ingest_asset(job_id: str) -> dict:
job_repo.update(job)
# 将上传时创建的占位 AssetPROCESSING/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,
+7
View File
@@ -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
+7
View File
@@ -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
+6 -3
View File
@@ -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:
+2 -1
View File
@@ -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 \
+2 -1
View File
@@ -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 \
+54 -9
View File
@@ -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 beatbeat 负责定期触发 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
@@ -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)
@@ -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)
@@ -38,6 +38,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
bgm_config=dict(getattr(model, "bgm_config", {}) or {}),
is_preview=bool(getattr(model, "is_preview", False)),
source_task_id=getattr(model, "source_task_id", "") or "",
celery_task_id=getattr(model, "celery_task_id", "") or "",
output_width=getattr(model, "output_width", 1280) or 1280,
output_height=getattr(model, "output_height", 720) or 720,
cover_url=getattr(model, "cover_url", "") or "",
@@ -82,6 +83,7 @@ class SQLAlchemyGenerationTaskRepository:
bgm_config=task.bgm_config or {},
is_preview=task.is_preview or False,
source_task_id=task.source_task_id or "",
celery_task_id=getattr(task, "celery_task_id", "") or "",
output_width=task.output_width,
output_height=task.output_height,
cover_url=task.cover_url or "",
@@ -138,6 +140,52 @@ class SQLAlchemyGenerationTaskRepository:
.count()
)
def count_running_by_user(self, user_id: str) -> int:
"""统计指定用户处于 running 状态的任务数(用于限流提示展示)。"""
return (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.created_by_user_id == user_id,
GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value,
)
.count()
)
def count_running_total(self) -> int:
"""统计全局处于 running 状态的任务数(worker 实际在执行的任务数)。"""
return (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value)
.count()
)
def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float:
"""估算最近完成任务的平均耗时(秒),用于 429 限流提示的等待预估。
取最近 N completed 任务的 (completed_at - started_at) 平均值
无足够历史数据时返回 default_seconds
Python 侧计算差值避免 SQLite/PostgreSQL 方言差异
"""
rows = (
self.session.query(GenerationTaskModel.started_at, GenerationTaskModel.completed_at)
.filter(
GenerationTaskModel.status == GenerationTaskStatus.COMPLETED.value,
GenerationTaskModel.started_at.isnot(None),
GenerationTaskModel.completed_at.isnot(None),
)
.order_by(GenerationTaskModel.completed_at.desc())
.limit(limit)
.all()
)
durations = [
(completed - started).total_seconds()
for started, completed in rows
if completed and started and (completed - started).total_seconds() > 0
]
if not durations:
return default_seconds
return sum(durations) / len(durations)
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
@@ -269,6 +317,7 @@ class SQLAlchemyGenerationTaskRepository:
if hasattr(model, "is_preview"):
model.is_preview = task.is_preview or False
model.source_task_id = task.source_task_id or ""
model.celery_task_id = getattr(task, "celery_task_id", "") or model.celery_task_id or ""
model.output_width = task.output_width
model.output_height = task.output_height
model.cover_url = task.cover_url or ""
@@ -280,12 +329,14 @@ class SQLAlchemyGenerationTaskRepository:
def cleanup_stale_running(self, timeout_minutes: int = 10) -> int:
"""清理超时未更新的 running 任务(孤儿任务)。
status=running updated_at 超过 timeout_minutes 分钟未更新的任务
标记为 failederror_message 标记为任务执行中断
Returns:
清理的任务数量
清理的任务数量仅计数保持旧签名兼容
"""
items = self.cleanup_stale_running_with_ids(timeout_minutes)
return len(items)
def cleanup_stale_running_with_ids(self, timeout_minutes: int = 10) -> list[tuple[str, str]]:
"""同 cleanup_stale_running,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
from datetime import timedelta
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
@@ -298,8 +349,10 @@ class SQLAlchemyGenerationTaskRepository:
.all()
)
if not models:
return 0
return []
result: list[tuple[str, str]] = []
for model in models:
result.append((model.id, getattr(model, "celery_task_id", "") or ""))
model.status = GenerationTaskStatus.FAILED.value
model.error_message = "任务执行中断(worker重启/超时)"
model.error_info = {
@@ -309,43 +362,43 @@ class SQLAlchemyGenerationTaskRepository:
}
model.completed_at = datetime.now(timezone.utc)
self.session.commit()
return len(models)
return result
def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int:
"""清理超时的 pending 任务(未被 Worker 拉取的任务)。
全局任务队列有 pending 数量上限长期卡在 pending 的任务会占满队列
导致新用户无法创建任务将超时的 pending 任务标记为 failed
Args:
timeout_minutes: 超时时间分钟默认 30 分钟
Returns:
清理的任务数量
清理的任务数量仅计数保持旧签名兼容
"""
items = self.cleanup_stale_pending_with_ids(timeout_minutes)
return len(items)
def cleanup_stale_pending_with_ids(self, timeout_minutes: int = 30) -> list[tuple[str, str]]:
"""同 cleanup_stale_pending,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
from datetime import timedelta
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
error_info = {
"error_type": "PendingTimeout",
"message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
"failed_at": datetime.now(timezone.utc).isoformat(),
}
count = (
models = (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
GenerationTaskModel.created_at < cutoff,
)
.update(
{
GenerationTaskModel.status: GenerationTaskStatus.FAILED.value,
GenerationTaskModel.error_message: "pending timeout: auto cleanup",
GenerationTaskModel.error_info: error_info,
GenerationTaskModel.completed_at: datetime.now(timezone.utc),
},
synchronize_session=False,
)
.all()
)
if not models:
return []
error_info = {
"error_type": "PendingTimeout",
"message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
"failed_at": datetime.now(timezone.utc).isoformat(),
}
result: list[tuple[str, str]] = []
for model in models:
result.append((model.id, getattr(model, "celery_task_id", "") or ""))
model.status = GenerationTaskStatus.FAILED.value
model.error_message = "pending timeout: auto cleanup"
model.error_info = error_info
model.completed_at = datetime.now(timezone.utc)
self.session.commit()
return count
return result
@@ -18,6 +18,8 @@ class SQLAlchemyIngestJobRepository:
error_message=job.error_message,
result_asset_id=job.result_asset_id,
file_hash=job.file_hash,
asset_id=job.asset_id or "",
celery_task_id=getattr(job, "celery_task_id", "") or "",
created_at=job.created_at,
updated_at=job.updated_at,
)
@@ -38,6 +40,8 @@ class SQLAlchemyIngestJobRepository:
error_message=model.error_message,
result_asset_id=model.result_asset_id,
file_hash=model.file_hash or "",
asset_id=getattr(model, "asset_id", "") or "",
celery_task_id=getattr(model, "celery_task_id", "") or "",
created_at=model.created_at,
updated_at=model.updated_at,
)
@@ -54,6 +58,12 @@ class SQLAlchemyIngestJobRepository:
model.error_message = job.error_message
model.result_asset_id = job.result_asset_id
model.file_hash = job.file_hash
model.storage_key = job.storage_key
if job.asset_id:
model.asset_id = job.asset_id
celery_tid = getattr(job, "celery_task_id", "")
if celery_tid:
model.celery_task_id = celery_tid
model.updated_at = job.updated_at
self.session.commit()
return job
@@ -38,6 +38,7 @@ class UserModel(Base):
phone = Column(String(32), nullable=True, unique=True, index=True)
phone_verified = Column(Boolean, nullable=False, default=False)
binding_completed_at = Column(DateTime, nullable=True)
profile_completed = Column(Boolean, nullable=False, default=True, server_default="true")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -94,6 +95,7 @@ class AssetModel(Base):
quality_score = Column(Float, nullable=True)
uploaded_by_user_id = Column(String(36), nullable=False)
file_hash = Column(String(64), nullable=True, index=True)
client_upload_id = Column(String(64), nullable=True, index=True)
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc), index=True)
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -244,6 +246,8 @@ class IngestJobModel(Base):
error_message = Column(Text, nullable=False, default="")
result_asset_id = Column(String(36), nullable=False, default="")
file_hash = Column(String(64), nullable=True, index=True)
asset_id = Column(String(36), nullable=False, default="", index=True)
celery_task_id = Column(String(64), nullable=False, default="", server_default="")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -296,6 +300,7 @@ class GenerationTaskModel(Base):
resolution = Column(String(20), nullable=False, default="")
is_preview = Column(Boolean, nullable=False, default=False, index=True)
source_task_id = Column(String(32), nullable=False, default="", index=True)
celery_task_id = Column(String(64), nullable=False, default="", server_default="")
output_width = Column(Integer, nullable=False, default=1280)
output_height = Column(Integer, nullable=False, default=720)
cover_url = Column(String(1000), nullable=False, default="")

Some files were not shown because too many files have changed in this diff Show More