Compare commits

..

1 Commits

Author SHA1 Message Date
xiaoxia b203d84956 fix: WebCodecs 解码失败时自动 fallback 到原生 video 播放
- useCanvasPlayer 新增 hasDecodeError/errorMessage 状态和 onError 回调
- decodeSegment catch 块不再静默失败,更新状态并通知上层
- VideoDecoder error 回调同步报告错误状态
- 连续 3 次解码失败时报告 hasDecodeError
- FrontendPreviewPlayer 检测 hasDecodeError 后自动切换 video fallback
- 解码失败时显示具体错误信息而非静默'暂无可播放素材'
- 新增 normalizeCodecString 函数规范化 codec 字符串
- configure 调用增加 codec 字符码调试日志
2026-08-21 14:43:01 +08:00
77 changed files with 2063 additions and 3455 deletions
@@ -1,26 +0,0 @@
"""Add title_config to generation_tasks
Revision ID: 057_title_config
Revises: 056_fix_cover_templates_config
Create Date: 2026-08-23
"""
import sqlalchemy as sa
from alembic import op
revision = "057_title_config"
down_revision = "056_fix_cover_templates_config"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"generation_tasks",
sa.Column("title_config", sa.JSON(), nullable=False, server_default="{}"),
)
def downgrade() -> None:
op.drop_column("generation_tasks", "title_config")
+39 -224
View File
@@ -31,6 +31,8 @@ logger = logging.getLogger(__name__)
router = APIRouter(tags=["Generation"])
# ── Schemas ──────────────────────────────────────────────────────────────
@@ -63,83 +65,6 @@ class GenerateCoverResponse(BaseModel):
# ── Route ────────────────────────────────────────────────────────────────
def _persist_cover_frame(
frame_url: str,
plan_id: str,
title_text: str = "",
*,
title_color: str = "#ffffff",
title_position: str = "bottom",
title_font_size: int | None = None,
) -> str:
"""下载 MediaKit 返回的临时帧图,可选叠加标题后转存到 OSS covers/ 路径。
Args:
frame_url: MediaKit 返回的临时帧图 URL
plan_id: 剪辑计划 ID(生成 OSS key
title_text: 非空时用 Pillow 在帧上叠加标题(用于 E2 从源素材抽帧,
因为源素材本身没有烧录标题)
title_color: 标题字体颜色(#RRGGBB
title_position: 标题位置 top/center/bottom
title_font_size: 标题字号,None 时自动计算
"""
import tempfile
import uuid
from pathlib import Path
tmp_path: str | None = None
try:
import httpx
resp = httpx.get(frame_url, timeout=30, follow_redirects=True)
resp.raise_for_status()
if not resp.content:
return frame_url
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp.write(resp.content)
tmp_path = tmp.name
# E2 从源素材抽帧时,源素材无标题,叠加标题文字
if title_text and title_text.strip():
try:
from packages.shared.title_overlay import apply_title_to_image
applied = apply_title_to_image(
tmp_path,
title_text,
color=title_color,
position=title_position,
font_size=title_font_size,
)
if applied:
logger.info("[封面生成] E2 帧图已叠加标题: plan_id=%s", plan_id)
except Exception:
logger.warning(
"[封面生成] E2 标题叠加失败(返回无标题帧): plan_id=%s",
plan_id,
exc_info=True,
)
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
cover_key = f"covers/{plan_id}/cover_{uuid.uuid4().hex[:8]}.jpg"
storage.upload_file(
file_or_path=tmp_path,
storage_key=cover_key,
content_type="image/jpeg",
)
public_url = storage.get_url(cover_key)
return public_url or frame_url
except Exception:
logger.warning("封面帧转存失败,返回原始 URL: plan_id=%s", plan_id, exc_info=True)
return frame_url
finally:
if tmp_path:
Path(tmp_path).unlink(missing_ok=True)
@router.post("/generate-cover", response_model=GenerateCoverResponse)
def generate_cover(
body: GenerateCoverRequest,
@@ -273,31 +198,44 @@ def generate_cover(
exc_info=True,
)
# 使用裸 URLrendered/* 已配置公开读);找不到渲染视频时不立即报错,
# 因为步骤 E 可以直接从源素材抽帧(历史数据或 Worker 抽帧失败时的兜底)
# 仍然找不到才报 400
if not rendered_storage_key:
logger.error("[封面生成] ❌ 找不到预览视频: plan_id=%s", plan_id)
raise HTTPException(
status_code=400,
detail="请先生成预览视频,再生成封面",
)
# 回写到 plan.config
plan_svc.update_plan_config(plan_id, {"rendered_storage_key": rendered_storage_key})
# 使用裸 URLrendered/* 已配置公开读)
primary_video_url = None
if rendered_storage_key:
plan_svc.update_plan_config(plan_id, {"rendered_storage_key": rendered_storage_key})
try:
if rendered_storage_key.startswith("http"):
primary_video_url = rendered_storage_key
else:
from packages.shared.storage import get_shared_storage_service
try:
if rendered_storage_key.startswith("http"):
primary_video_url = rendered_storage_key
else:
from packages.shared.storage import get_shared_storage_service
storage_svc = get_shared_storage_service()
primary_video_url = storage_svc.get_url(rendered_storage_key)
if primary_video_url:
import re as _re
storage_svc = get_shared_storage_service()
primary_video_url = storage_svc.get_url(rendered_storage_key)
# 防御性规范化:合并路径中的双斜杠(// -> /),但保留协议头的 ://
# 历史数据中 project_id 为空时会产生 projects//tasks/ 路径,
# MediaKit 的 HTTP 客户端会规范化 URL 导致 404
if primary_video_url:
import re as _re
primary_video_url = _re.sub(r"(?<!:)//", "/", primary_video_url)
logger.info(
"获取预览视频URL用于封面生成: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
except Exception as e:
logger.warning("获取预览视频URL失败: plan_id=%s err=%s", plan_id, e)
primary_video_url = None
primary_video_url = _re.sub(r"(?<!:)//", "/", primary_video_url)
logger.info(
"获取预览视频URL用于封面生成: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
except Exception as e:
raise HTTPException(
status_code=500,
detail=f"获取预览视频URL失败: {e}",
) from e
# 统一封面管道:优先从 GenerationTask.cover_url 读取渲染后视频抽帧的封面
# 多步查找 cover_url,和查找视频 URL 一样的 fallback 逻辑
@@ -372,129 +310,6 @@ def generate_cover(
exc_info=True,
)
# 步骤 D:从 plan.config.cover_candidates 读取(Worker 渲染时写入)
if not cover_url_from_task:
_candidates = (plan.config or {}).get("cover_candidates") or []
if isinstance(_candidates, list) and _candidates:
_first = _candidates[0]
if isinstance(_first, dict):
cover_url_from_task = _first.get("image_url") or _first.get("url") or ""
if cover_url_from_task:
logger.info(
"[封面生成] 统一管道封面(步骤D-cover_candidates): plan_id=%s url=%s",
plan_id,
cover_url_from_task[:80],
)
# 步骤 E1:如果有已渲染的预览视频 URL 但 cover_url 未持久化(历史数据),
# 直接从渲染视频抽帧
if not cover_url_from_task and primary_video_url:
try:
from packages.shared.mediakit_client import get_mediakit_client
mk_client = get_mediakit_client()
if mk_client.is_available:
logger.info(
"[封面生成] 步骤E1-从渲染视频抽帧: plan_id=%s url=%s",
plan_id,
primary_video_url[:80],
)
snapshots = mk_client.extract_frames(
video_url=primary_video_url,
strategy="SpecifiedFrames",
max_frames=1,
poll_interval=2.0,
max_poll_attempts=5,
max_retries=0,
)
if snapshots:
raw = snapshots[0].get("image_url") or snapshots[0].get("url") or ""
if raw:
cover_url_from_task = _persist_cover_frame(raw, plan_id)
logger.info(
"[封面生成] 统一管道封面(步骤E1-rendered-video): plan_id=%s url=%s",
plan_id,
cover_url_from_task[:80],
)
except Exception:
logger.warning(
"[封面生成] 步骤E1从渲染视频抽帧失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 步骤 E2:当 A/B/C/D/E1 均未命中(如历史预览任务无 cover_url)时,
# 直接从用户选择的第一个视频素材中抽取封面帧作为兜底。API 请求内短超时,不阻塞。
if not cover_url_from_task and body.asset_ids:
from packages.adapters.sqlalchemy_impl.asset_repository import (
SQLAlchemyAssetRepository,
)
from packages.shared.mediakit_client import get_mediakit_client
from packages.shared.storage import get_shared_storage_service
asset_repo = SQLAlchemyAssetRepository(db)
storage_svc = get_shared_storage_service()
mk_client = get_mediakit_client()
# 从 plan.config 读取完整标题样式,E2 从源素材抽帧时叠加(源素材本身无标题)
_e2_title_cfg = (plan.config or {}).get("title", {}) or {}
if not isinstance(_e2_title_cfg, dict):
_e2_title_cfg = {}
_e2_title_text = (_e2_title_cfg.get("text", "") or "").strip() if _e2_title_cfg.get("enabled", True) else ""
# 读取标题样式:前端可能传 color 或 font_color,都兼容
_e2_title_color = _e2_title_cfg.get("color") or _e2_title_cfg.get("font_color") or "#ffffff"
_e2_title_position = _e2_title_cfg.get("position", "bottom") or "bottom"
_e2_title_font_size = _e2_title_cfg.get("font_size") or _e2_title_cfg.get("size")
if mk_client.is_available:
for aid in body.asset_ids:
try:
asset = asset_repo.get(aid)
if not asset or asset.file_type != "video":
continue
sk = asset.storage_key or ""
if not sk:
continue
src_url = sk if sk.startswith("http") else storage_svc.get_url(sk)
if not src_url:
continue
logger.info(
"[封面生成] 步骤E-从素材抽帧: plan_id=%s asset_id=%s url=%s",
plan_id,
aid,
src_url[:80],
)
snapshots = mk_client.extract_frames(
video_url=src_url,
strategy="SpecifiedFrames",
max_frames=1,
poll_interval=2.0,
max_poll_attempts=5,
max_retries=0,
)
if snapshots:
raw = snapshots[0].get("image_url") or snapshots[0].get("url") or ""
if raw:
cover_url_from_task = _persist_cover_frame(
raw,
plan_id,
title_text=_e2_title_text,
title_color=_e2_title_color,
title_position=_e2_title_position,
title_font_size=_e2_title_font_size,
)
logger.info(
"[封面生成] 统一管道封面(步骤E-source-asset): plan_id=%s url=%s",
plan_id,
cover_url_from_task[:80],
)
break
except Exception:
logger.warning(
"[封面生成] 步骤E从素材抽帧失败: plan_id=%s asset_id=%s",
plan_id,
aid,
exc_info=True,
)
if cover_url_from_task:
# 标题已在预览视频渲染时烧录(ASS字幕),封面帧自然包含标题
cover_data = {
@@ -510,13 +325,13 @@ def generate_cover(
return GenerateCoverResponse(plan_id=plan_id, cover=cover_data)
logger.warning(
"[封面生成] 统一管道未找到 cover_url (A/B/C/D均未命中): plan_id=%s",
"[封面生成] 统一管道未找到 cover_url: plan_id=%s",
plan_id,
)
# ai_frame/ai_regenerate 类型必须从渲染管道获取,不再回退到 AI 服务
raise HTTPException(
status_code=400,
detail="封面生成失败:未找到可抽帧的视频素材,请确认已上传视频素材后重试",
detail="封面尚未生成,请先重新生成预览视频以触发封面自动提取",
)
from packages.shared.ai_service import run_generate_cover
@@ -16,7 +16,6 @@ from app.core.task_enqueue import (
from app.dependencies import (
get_asset_library_repository,
get_asset_repository,
get_db_session,
get_generated_video_repository,
get_generation_task_repository,
get_project_repository,
@@ -33,7 +32,6 @@ from app.schemas.generation_task import (
ListGenerationTasksResponse,
)
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from packages.application import (
CreateGenerationTaskCommand,
@@ -71,7 +69,6 @@ def _to_generation_task_response(task) -> GenerationTaskResponse:
output_height=getattr(task, "output_height", 720),
cover_url=getattr(task, "cover_url", ""),
custom_title=getattr(task, "custom_title", ""),
title_config=getattr(task, "title_config", {}) or {},
logs=getattr(task, "logs", "[]"),
status=task.status,
progress=task.progress,
@@ -145,54 +142,6 @@ def _select_assets_from_library(
return [a.id for a in ready_video_assets]
def _writeback_edit_plan_config(
plan_id: str,
task_id: str,
title_config: dict | None,
db: Session,
) -> None:
"""任务入队成功后,回写 EditPlan.configgeneration_task_id + title_config。
用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。
失败只记日志,不影响任务创建。
"""
if not plan_id:
return
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
plan_model = db.query(EditPlanModel).filter(EditPlanModel.id == plan_id).first()
if plan_model is None:
logger.warning("[生成任务] 回写plan.config失败: plan不存在 plan_id=%s", plan_id)
return
current_config = plan_model.config if isinstance(plan_model.config, dict) else {}
merged = dict(current_config)
merged["generation_task_id"] = task_id
if title_config:
merged["title_config"] = title_config
plan_model.config = merged
db.commit()
logger.info(
"[生成任务] 回写plan.config成功: plan_id=%s task_id=%s keys=%s",
plan_id,
task_id,
list(merged.keys()),
)
except Exception as e:
logger.warning(
"[生成任务] 回写plan.config异常(不影响任务创建): plan_id=%s error=%s",
plan_id,
e,
exc_info=True,
)
try:
db.rollback()
except Exception:
pass
def _resolve_project_and_library(
request: CreateGenerationTaskRequest,
project_repository: Any,
@@ -238,7 +187,6 @@ def create_generation_task(
project_repository: Any = Depends(get_project_repository),
asset_library_repository: Any = Depends(get_asset_library_repository),
asset_repository: Any = Depends(get_asset_repository),
db: Session = Depends(get_db_session),
) -> BatchGenerationTaskResponse:
logger.info(
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%s, count=%d",
@@ -356,7 +304,6 @@ def create_generation_task(
output_height=request.output_height,
cover_url=request.cover_url,
custom_title=request.custom_title,
title_config=request.title_config or {},
)
)
try:
@@ -368,15 +315,6 @@ def create_generation_task(
log_task_status=True,
):
created_tasks.append(task)
# 只在首个成功任务时回写一次 plan.config
# 避免批量生成时循环覆盖 generation_task_id
if request.source_edit_plan_id and len(created_tasks) == 1:
_writeback_edit_plan_config(
plan_id=request.source_edit_plan_id,
task_id=task.id,
title_config=request.title_config,
db=db,
)
else:
failed_tasks.append(task)
except UserPendingLimitExceeded as _e:
@@ -17,8 +17,6 @@ from fastapi import APIRouter, Depends, HTTPException, Query, status
from .dependencies import get_draft_plan_id, get_editor_services
from .schemas import (
EditorClipBatchUpdateRequest,
EditorClipBatchUpdateResponse,
EditorDraftResponse,
EditorPublishResponse,
EditorRollbackRequest,
@@ -128,7 +126,11 @@ def list_template_versions(
clip_count=len(v.clip_configs),
change_note=v.change_note,
published_by=v.published_by,
created_at=(v.created_at.isoformat() if hasattr(v.created_at, "isoformat") else str(v.created_at)),
created_at=(
v.created_at.isoformat()
if hasattr(v.created_at, "isoformat")
else str(v.created_at)
),
)
for v in versions
]
@@ -160,35 +162,3 @@ def rollback_template(
new_version=tpl.version,
clip_count=len(clip_configs),
)
@router.put("/clips", response_model=EditorClipBatchUpdateResponse)
def batch_update_clips(
template_id: str,
req: EditorClipBatchUpdateRequest,
plan_id: str = Depends(get_draft_plan_id),
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
_: AuthenticatedUser = Depends(get_current_user),
):
"""批量替换草稿clips(全量覆盖,用于前端选择素材后同步片段)
事务保证:清空→创建→标记ready 在同一数据库事务内完成,
任何步骤失败时自动回滚,避免数据不一致。
"""
_, plan_svc = services
plan_svc.get_plan_or_raise(plan_id)
clips_data = []
for clip_item in req.clips:
item = {
"asset_id": clip_item.asset_id,
"start_time": clip_item.start_time,
"duration": clip_item.duration,
}
if clip_item.order is not None:
item["order"] = clip_item.order
clips_data.append(item)
plan_svc.replace_all_clips_transactional(plan_id, clips_data)
return EditorClipBatchUpdateResponse(plan_id=plan_id, clip_count=len(req.clips))
@@ -56,7 +56,6 @@ logger = logging.getLogger(__name__)
router = APIRouter(tags=["Template Editor"])
# DEPRECATED: 前端已改用 /generation/tasks 体系,此路由保留仅供旧版兼容,计划下线
@router.post("/generate", response_model=EditPlanGenerateResponse)
def generate_editor_draft(
template_id: str,
@@ -197,7 +196,7 @@ def generate_editor_draft(
plan_svc.update_plan_config(plan_id, {"generation_task_id": gen_task.id})
plan_svc.transition_status(plan_id, EditPlanStatus.RENDERING)
celery_app.send_task("worker.generate_video", args=[gen_task.id])
celery_app.send_task("worker.render_edit_plan", args=[plan_id])
updated_plan = plan_svc.get_plan_or_raise(plan_id)
@@ -286,7 +285,6 @@ def _get_task_output_url(task, gen_task_repo, db) -> str:
return ""
# DEPRECATED: 前端已改用 /generation/tasks 体系,此路由保留仅供旧版兼容,计划下线
@router.get("/generation-status", response_model=EditPlanGenerationStatusResponse)
def get_editor_generation_status(
template_id: str,
@@ -46,7 +46,6 @@ class EditPlanGenerationStatusResponse(BaseModel):
class EditPlanGenerateRequest(BaseModel):
"""模板编辑器触发生成请求体"""
title_config: Optional[Dict[str, Any]] = Field(
default_factory=dict,
description="标题配置(可选),渲染时烧录到视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow",
@@ -76,8 +75,12 @@ class AIRecommendRequest(BaseModel):
"""AI 推荐片段方案请求体"""
asset_ids: List[str] = Field(default_factory=list, description="素材 ID 列表")
editing_mode: str = Field(default="one_take", description="剪辑模式: one_take / pip / voice_over / voice_pip")
target_duration: float = Field(default=30.0, ge=1.0, le=600.0, description="目标时长(秒)")
editing_mode: str = Field(
default="one_take", description="剪辑模式: one_take / pip / voice_over / voice_pip"
)
target_duration: float = Field(
default=30.0, ge=1.0, le=600.0, description="目标时长(秒)"
)
class AIRecommendClipItem(BaseModel):
@@ -104,6 +107,8 @@ class AIRecommendResponse(BaseModel):
confidence: float = Field(..., ge=0.0, le=1.0, description="AI 推荐置信度 (0~1)")
# ── BGM ────────────────────────────────────────────────────────────────────
@@ -219,7 +224,9 @@ class ClipBatchDeleteResponse(BaseModel):
class ClipsFromAssetsRequest(BaseModel):
"""从素材批量创建片段请求"""
asset_ids: List[str] = Field(..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾")
asset_ids: List[str] = Field(
..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾"
)
clip_type: str = Field(default="main", description="片段类型,默认 main")
@@ -494,28 +501,6 @@ class EditorClipUpdateRequest(BaseModel):
config: Optional[dict[str, Any]] = None
class EditorClipBatchItem(BaseModel):
"""批量更新clips的单个片段"""
asset_id: str = Field(default="", max_length=100, description="关联素材ID,可为空(占位片段)")
start_time: float = Field(default=0.0, ge=0.0)
duration: float = Field(default=0.0, ge=0.0)
order: Optional[int] = Field(default=None, ge=0, description="排序,None表示按数组顺序")
class EditorClipBatchUpdateRequest(BaseModel):
"""批量替换clips请求(全量覆盖)"""
clips: List[EditorClipBatchItem] = Field(default_factory=list)
class EditorClipBatchUpdateResponse(BaseModel):
"""批量更新clips响应"""
plan_id: str
clip_count: int
class EditorPublishResponse(BaseModel):
"""发布草稿响应"""
-6
View File
@@ -33,11 +33,6 @@ class CreateGenerationTaskRequest(BaseModel):
voice_ids: list[str] = Field(default_factory=list)
# ── 来源剪辑计划 ──
source_edit_plan_id: str = ""
# ── 标题配置(结构化,优先于 custom_title 纯文本)──
title_config: dict | None = Field(
default=None,
description="标题样式对象,包含 text/font/font_size/font_color/position/bold/stroke/shadow 等。为空时不影响现有行为。",
)
# ── 视频标题 ──
video_title: str = Field(default="", description="生成视频的标题/名称,为空则使用默认命名")
# ── 批量生成 ──
@@ -114,7 +109,6 @@ class GenerationTaskResponse(BaseModel):
output_height: int = 720
cover_url: str = ""
custom_title: str = ""
title_config: dict = Field(default_factory=dict)
status: str
progress: float
result_count: int
@@ -371,85 +371,6 @@ class EditPlanService:
logger.info("删除所有片段: plan_id=%s count=%d", plan_id, count)
return count
def replace_all_clips_transactional(
self,
plan_id: str,
clips_data: list[dict],
) -> int:
"""事务性地替换所有片段:清空→创建→标记ready,单事务保证原子性。
Args:
plan_id: 计划 ID
clips_data: 片段数据列表,每项包含 asset_id/start_time/duration/order
Returns:
int: 创建的片段数量
Raises:
Exception: 任何步骤失败时自动回滚
"""
from packages.adapters.sqlalchemy_impl.models import EditPlanClipModel
db = self._clip_repo.session
try:
# 1. 清空现有 clips(不 commit
deleted_count = db.query(EditPlanClipModel).filter(EditPlanClipModel.plan_id == plan_id).delete()
# 2. 批量创建新 clips(不 commit
for i, clip_item in enumerate(clips_data):
order = clip_item.get("order") or i
clip = EditPlanClip.create(
plan_id=plan_id,
clip_type="main",
order=order,
asset_id=clip_item.get("asset_id", ""),
start_time=clip_item.get("start_time", 0.0),
duration=clip_item.get("duration", 0.0),
)
model = EditPlanClipModel(
id=clip.id,
plan_id=clip.plan_id,
clip_type=clip.clip_type,
order=clip.order,
asset_id=clip.asset_id,
text_content=clip.text_content,
start_time=clip.start_time,
duration=clip.duration,
transition_effect=clip.transition_effect,
transition_duration=clip.transition_duration,
playback_speed=clip.playback_speed,
status=clip.status.value,
config=clip.config,
)
db.add(model)
# 3. 标记有 asset_id 的 clips 为 ready(不 commit
pending_with_asset = (
db.query(EditPlanClipModel)
.filter(
EditPlanClipModel.plan_id == plan_id,
EditPlanClipModel.status == "pending",
EditPlanClipModel.asset_id != "",
)
.all()
)
for m in pending_with_asset:
m.status = "ready"
# 4. 一次性提交
db.commit()
logger.info(
"事务性替换片段: plan_id=%s deleted=%d created=%d",
plan_id,
deleted_count,
len(clips_data),
)
return len(clips_data)
except Exception:
db.rollback()
logger.exception("事务性替换片段失败: plan_id=%s", plan_id)
raise
# ── 片段分割与合并 ──────────────────────────────────────────────────────
def split_clip(self, clip_id: str, split_time: float) -> Dict[str, Any]:
+4 -1
View File
@@ -231,7 +231,10 @@ test.describe("Core generation flow", () => {
(response) => {
const url = response.url()
const path = new URL(url).pathname
return response.request().method() === "POST" && path.endsWith("/generation/tasks")
return (
response.request().method() === "POST" &&
path.endsWith("/generation/tasks")
)
},
{ timeout: 30_000 },
)
+3 -12
View File
@@ -4,20 +4,11 @@
import apiClient from "../client"
import type { BgmPreset, BgmPresetsQuery } from "./types"
/**
* 获取 BGM 预设列表
* @param templateId 模板/草稿 ID
* @param params 分类/关键词筛选
*/
export const getBgmPresets = async (
templateId: string,
params?: BgmPresetsQuery,
): Promise<BgmPreset[]> => {
/** 获取 BGM 预设列表 */
export const getBgmPresets = async (params?: BgmPresetsQuery): Promise<BgmPreset[]> => {
const searchParams: Record<string, string> = {}
if (params?.category) searchParams.category = params.category
if (params?.keyword) searchParams.keyword = params.keyword
const res = await apiClient.get(`/templates/${templateId}/editor/bgm/presets`, {
params: searchParams,
})
const res = await apiClient.get("/bgm/presets", { params: searchParams })
return res.data?.data ?? res.data ?? []
}
-13
View File
@@ -1,22 +1,9 @@
import apiClient from "../client"
export interface GenerateCoverTitleConfig {
text?: string
font?: string
font_size?: number
font_color?: string
position?: string
bold?: boolean
stroke?: boolean
shadow?: boolean
}
export interface GenerateCoverRequest {
asset_ids: string[]
cover_type?: "ai_frame" | "manual" | "upload" | "ai_regenerate"
frame_time?: number
/** 标题样式,用于在封面上叠加标题文字 */
title_config?: GenerateCoverTitleConfig
}
export interface GenerateCoverResponse {
-7
View File
@@ -6,7 +6,6 @@ import apiClient from "../client"
import type {
CreateGenerationTaskRequest,
CreateGenerationTaskResponse,
GenerationTaskDetail,
TaskItem,
TaskListParams,
TaskListResponse,
@@ -20,12 +19,6 @@ export const createGenerationTask = async (
return data
}
/** 获取单个生成任务详情(轮询用) */
export const getGenerationTask = async (taskId: string): Promise<GenerationTaskDetail> => {
const { data } = await apiClient.get<GenerationTaskDetail>(`/generation/tasks/${taskId}`)
return data
}
/** 获取任务列表(支持分页和筛选) */
export const getTasks = async (params?: TaskListParams): Promise<TaskListResponse> => {
const { data } = await apiClient.get<TaskListResponse>("/tasks", {
+2 -14
View File
@@ -82,12 +82,10 @@ export interface CreateGenerationTaskRequest {
stroke?: boolean
shadow?: boolean
}
/** 关联的草稿 ID(编辑流程数据链路用) */
source_edit_plan_id?: string
}
/** 单个生成任务详情(对齐后端 GenerationTaskResponse */
export interface GenerationTaskDetail {
/** 创建生成任务响应(对齐后端 GenerationTaskResponse */
export interface CreateGenerationTaskResponse {
id: string
project_id: string
asset_library_id: string
@@ -97,18 +95,8 @@ export interface GenerationTaskDetail {
asset_ids: string[]
title_ids: string[]
voice_ids: string[]
source_edit_plan_id?: string
status: string
progress: number
result_count: number
error_message: string
error_info?: TaskErrorInfo
created_at?: string | null
updated_at?: string | null
}
/** 创建生成任务响应(后端返回批量结构 {items, total} */
export interface CreateGenerationTaskResponse {
items: GenerationTaskDetail[]
total: number
}
+42 -3
View File
@@ -4,29 +4,51 @@
import apiClient from "../client"
import type {
EditPlan,
EditPlanListParams,
EditPlanListResponse,
CreateEditPlanRequest,
UpdateEditPlanRequest,
GenerateResponse,
GenerationStatusResponse,
EditPlanGeneration,
GeneratedVideo,
CopyEditPlanRequest,
} from "./types"
/** 获取模板草稿列表(支持分页和筛选) */
export async function getEditPlans(params?: EditPlanListParams): Promise<EditPlanListResponse> {
const response = await apiClient.get<EditPlanListResponse>("/templates/drafts", {
params,
})
return response.data
}
/** 获取单个模板草稿 */
export async function getEditPlan(templateId: string): Promise<EditPlan> {
const response = await apiClient.get(`/templates/${templateId}/editor`)
return response.data
}
/** 更新模板草稿(支持传入 AbortSignal 用于自动保存竞态取消) */
/** 创建模板草稿 */
export async function createEditPlan(data: CreateEditPlanRequest): Promise<EditPlan> {
const response = await apiClient.post("/templates/drafts", data)
return response.data
}
/** 更新模板草稿 */
export async function updateEditPlan(
templateId: string,
data: UpdateEditPlanRequest,
signal?: AbortSignal,
): Promise<EditPlan> {
const response = await apiClient.put(`/templates/${templateId}/editor`, data, { signal })
const response = await apiClient.put(`/templates/${templateId}/editor`, data)
return response.data
}
/** 删除模板草稿 */
export async function deleteEditPlan(templateId: string): Promise<void> {
await apiClient.delete(`/templates/${templateId}/editor`)
}
/** 触发生成 */
export async function generateEditPlan(templateId: string): Promise<GenerateResponse> {
const response = await apiClient.post(`/templates/${templateId}/editor/generate`)
@@ -50,3 +72,20 @@ export async function getGenerationTaskResults(taskId: string): Promise<Generate
const response = await apiClient.get(`/generation/tasks/${taskId}/results`)
return response.data.items || response.data || []
}
/** 取消生成任务 */
export async function cancelGeneration(templateId: string): Promise<void> {
await apiClient.post(`/templates/${templateId}/editor/cancel`)
}
/** 复制模板草稿(含所有片段配置) */
export async function copyEditPlan(
templateId: string,
data?: CopyEditPlanRequest,
): Promise<EditPlan> {
const response = await apiClient.post<EditPlan>(
`/templates/${templateId}/editor/copy`,
data || {},
)
return response.data
}
@@ -15,7 +15,10 @@ export type {
EditPlanSegment,
EditPlanConfig,
EditPlan,
CreateEditPlanRequest,
UpdateEditPlanRequest,
EditPlanListParams,
EditPlanListResponse,
GenerateResponse,
EditPlanGeneration,
ClipStatusItem,
@@ -34,6 +37,7 @@ export type {
ClipReorderResponse,
ClipBatchDeleteResponse,
ClipsFromAssetsResponse,
CopyEditPlanRequest,
TransitionEffect,
MediaAsset,
} from "./types"
@@ -49,12 +53,17 @@ export {
// 模板草稿 CRUD + 生成
export {
getEditPlans,
getEditPlan,
createEditPlan,
updateEditPlan,
deleteEditPlan,
generateEditPlan,
getGenerationStatus,
getEditPlanGenerations,
getGenerationTaskResults,
cancelGeneration,
copyEditPlan,
} from "./editPlans"
// 片段 CRUD + 批量操作
-11
View File
@@ -118,17 +118,6 @@ export interface EditPlanConfig {
generate_count?: number
/** 素材模式 */
material_mode?: string
/** 前端标题设置(Step4 自动保存,与 title_config 字段分离,不影响后端渲染) */
title?: {
text?: string
font?: string
font_size?: number
color?: string
position?: string
bold?: boolean
stroke?: boolean
shadow?: boolean
}
/** 预览视频 URL(封面生成用) */
rendered_storage_key?: string
/** 生成任务 ID */
+3
View File
@@ -9,6 +9,8 @@ export type {
TemplateSegment,
TemplateListParams,
TemplateListResponse,
GenerateFromTemplateRequest,
GenerateFromTemplateResponse,
CopyTemplateResponse,
} from "./types"
@@ -22,4 +24,5 @@ export {
getTemplate,
toggleFavoriteTemplate,
copyTemplate,
generateFromTemplate,
} from "./templates"
+14
View File
@@ -5,6 +5,8 @@
import apiClient from "../client"
import type {
CopyTemplateResponse,
GenerateFromTemplateRequest,
GenerateFromTemplateResponse,
TemplateItem,
TemplateListParams,
TemplateListResponse,
@@ -43,3 +45,15 @@ export const copyTemplate = async (templateId: string): Promise<CopyTemplateResp
const response = await apiClient.post<CopyTemplateResponse>(`/templates/${templateId}/copy`)
return response.data
}
/** 从模板生成 */
export const generateFromTemplate = async (
templateId: string,
data?: GenerateFromTemplateRequest,
): Promise<GenerateFromTemplateResponse> => {
const response = await apiClient.post<GenerateFromTemplateResponse>(
`/templates/${templateId}/generate`,
data,
)
return response.data
}
@@ -16,17 +16,9 @@ interface BgmSelectorProps {
onClose: () => void
config: BgmMixConfig
onChange: (config: BgmMixConfig) => void
/** 模板/草稿 ID,用于请求 BGM 预设 */
templateId?: string
}
const BgmSelector: React.FC<BgmSelectorProps> = ({
open,
onClose,
config,
onChange,
templateId,
}) => {
const BgmSelector: React.FC<BgmSelectorProps> = ({ open, onClose, config, onChange }) => {
const {
presets,
loading,
@@ -38,7 +30,7 @@ const BgmSelector: React.FC<BgmSelectorProps> = ({
loadPresets,
handlePreview,
stopPreview,
} = useBgmSelector(open, templateId)
} = useBgmSelector(open)
/* ── 选中 BGM ── */
const handleSelect = useCallback(
@@ -19,7 +19,7 @@ export const CATEGORY_LIST: {
* BGM 选择器数据与交互 Hook
* 封装列表加载、搜索、分类筛选、试听播放逻辑
*/
export function useBgmSelector(open: boolean, templateId?: string) {
export function useBgmSelector(open: boolean) {
const [presets, setPresets] = useState<BgmPreset[]>([])
const [loading, setLoading] = useState(false)
const [activeCategory, setActiveCategory] = useState<BgmCategory | "all">("all")
@@ -30,23 +30,19 @@ export function useBgmSelector(open: boolean, templateId?: string) {
/* ── 加载 BGM 列表 ── */
const loadPresets = useCallback(async () => {
if (!templateId) {
setPresets([])
return
}
setLoading(true)
try {
const params: { category?: string; keyword?: string } = {}
if (activeCategory !== "all") params.category = activeCategory
if (keyword.trim()) params.keyword = keyword.trim()
const data = await getBgmPresets(templateId, params)
const data = await getBgmPresets(params)
setPresets(data)
} catch {
message.error("加载 BGM 列表失败")
} finally {
setLoading(false)
}
}, [activeCategory, keyword, templateId])
}, [activeCategory, keyword])
useEffect(() => {
if (open) loadPresets()
@@ -171,7 +171,6 @@ const GeneratePage: React.FC = () => {
autoSubtitles,
bgm,
generateCount,
sourceEditPlanId: editPlanId,
})
/* ================================================================
@@ -17,7 +17,7 @@ import {
import type { AssetItem } from "@/api/assets"
import type { EditingTemplate } from "@/api/editing-planner"
import { useSegmentScheduler, type PlaybackSegment } from "../hooks/useSegmentScheduler"
import { useCanvasPlayer } from "../hooks/useCanvasPlayer"
import { useCanvasPlayer, isWebCodecsSupported } from "../hooks/useCanvasPlayer"
interface FrontendPreviewPlayerProps {
assets: AssetItem[]
@@ -82,9 +82,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
titleSettings,
}) => {
const segments = useMemo(() => buildPlaybackSegments(assets, template), [assets, template])
// 默认走原生 video 播放(浏览器硬件解码,独立线程,不阻塞 UI)
// WebCodecs 仅在明确需要时启用(保留代码作为兜底)
const useWebCodecs = false
const useWebCodecs = isWebCodecsSupported()
// ── 两条路径共用同一个 canvas ref(fallback 路径不使用) ──
const canvasRef = useRef<HTMLCanvasElement>(null)
@@ -124,10 +122,9 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
const { state: canvasState, controls: canvasControls } = useCanvasPlayer(
canvasRef,
useWebCodecs && !forceVideoFallback ? canvasSegments : [],
canvasSegments,
useWebCodecs && !forceVideoFallback ? canvasTitle : undefined,
handleCanvasError,
useWebCodecs && !forceVideoFallback,
)
// WebCodecs 报告解码失败时自动切换到 video fallback
@@ -198,9 +195,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
const audio = audioRef.current
if (!audio || !audio.src || !isPlaying) return
audio.currentTime = currentTime
// 注意:不要把 currentTime 放进依赖数组,否则每200ms会重置音频位置导致卡顿
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [segmentSyncKey, isPlaying])
}, [segmentSyncKey, isPlaying, currentTime])
const handleSeekTo = useCallback(
(time: number) => {
@@ -270,8 +265,32 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
const progressPercent = totalDuration > 0 ? (currentTime / totalDuration) * 100 : 0
// ── Canvas 容器 ref(保留声明,WebCodecs 兜底路径仍引用) ──
// ── Canvas ResizeObserver ──
const canvasContainerRef = useRef<HTMLDivElement>(null)
useEffect(() => {
if (!effectiveUseWebCodecs || !canPlay) return
const container = canvasContainerRef.current
const canvas = canvasRef.current
if (!container || !canvas) return
// 立即设置一次 canvas 像素分辨率,避免默认 300×150 导致首帧变形
const initRect = container.getBoundingClientRect()
if (initRect.width > 0 && initRect.height > 0) {
const dpr = window.devicePixelRatio || 1
canvas.width = initRect.width * dpr
canvas.height = initRect.height * dpr
}
const ro = new ResizeObserver((entries) => {
for (const entry of entries) {
const { width, height } = entry.contentRect
if (width > 0 && height > 0) {
canvas.width = width * window.devicePixelRatio
canvas.height = height * window.devicePixelRatio
}
}
})
ro.observe(container)
return () => ro.disconnect()
}, [effectiveUseWebCodecs, canPlay])
// ── 未就绪 ──
if (!ready || !assets.length) {
@@ -366,7 +385,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
</div>
)}
{/* ── Video 渲染层(默认路径,浏览器原生硬件解码 ── */}
{/* ── Video 渲染层(fallback 路径,或 WebCodecs 解码失败时自动切换 ── */}
{!effectiveUseWebCodecs &&
segments.map((seg, i) => (
<video
@@ -375,7 +394,9 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
ref={(el) => {
videoRefs.current[i] = el
}}
preload="auto"
preload={
i === videoCurrentSegIdx ? "auto" : i === videoCurrentSegIdx + 1 ? "metadata" : "none"
}
src={seg.videoUrl}
style={{
position: "absolute",
@@ -434,7 +455,11 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
zIndex: 10,
}}
>
{`片段 ${videoCurrentSegIdx + 1}/${segments.length}`}
{effectiveUseWebCodecs
? "Canvas"
: forceVideoFallback
? "Canvas 解码失败,已切换原生播放"
: `片段 ${videoCurrentSegIdx + 1}/${segments.length}`}
</div>
{/* 控制条 */}
@@ -141,7 +141,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
onSelectedMaterialsChange={onSelectedMaterialsChange}
smartSelectedIds={smartSelectedIds}
onSmartSelectedIdsChange={onSmartSelectedIdsChange}
selectedTemplate={selectedTemplate}
/>
)
case 3:
@@ -157,7 +156,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
<Step4TitleSettings
titleSettings={titleSettings}
onTitleSettingsChange={onTitleSettingsChange}
selectedTemplate={selectedTemplate}
/>
)
case 5:
@@ -184,7 +182,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
duration={duration}
assetIds={materialMode === "auto" ? smartSelectedIds : selectedMaterials}
selectedTemplate={selectedTemplate}
titleSettings={titleSettings}
/>
)
case 7:
@@ -16,7 +16,6 @@ import { LoadingOutlined } from "@ant-design/icons"
import type { AssetItem } from "@/api/assets"
import type { EditingTemplate } from "@/api/editing-planner"
import type { TitleSettings } from "../types"
import { getFontFamily } from "../constants"
import FrontendPreviewPlayer from "./FrontendPreviewPlayer"
interface PreviewVideoPanelProps {
@@ -88,7 +87,7 @@ function buildTitleStyle(settings: TitleSettings, containerHeight: number): Reac
: (Math.min(settings.size, 96) / ASS_VIDEO_HEIGHT) * 400 // fallback
const base: React.CSSProperties = {
fontFamily: getFontFamily(settings.font),
fontFamily: settings.font || "思源黑体",
fontSize: `${fontSizePx}px`,
color: settings.color || "#ffffff",
fontWeight: settings.bold ? 700 : 400,
@@ -177,12 +176,7 @@ const TitleOverlay: React.FC<{ titleSettings: TitleSettings }> = ({ titleSetting
position: "absolute",
}}
>
{displayTitle.split("/").map((part, i) => (
<span key={i}>
{i > 0 && <br />}
{part}
</span>
))}
{displayTitle}
</div>
</div>
)
@@ -15,8 +15,6 @@ interface Step2MaterialSelectProps {
onSelectedMaterialsChange: (ids: string[]) => void
smartSelectedIds: string[]
onSmartSelectedIdsChange: (ids: string[]) => void
/** 当前选中的模板/草稿 ID,用于自动保存 */
selectedTemplate?: string
}
const Step2MaterialSelect: React.FC<Step2MaterialSelectProps> = (props) => {
@@ -12,8 +12,6 @@ import AiTitleGenerator from "./title/AiTitleGenerator"
interface Step4TitleSettingsProps {
titleSettings: TitleSettings
onTitleSettingsChange: (settings: TitleSettings) => void
/** 当前选中的模板/草稿 ID,用于自动保存 */
selectedTemplate?: string
}
const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
@@ -17,16 +17,6 @@ interface Step5VoiceSelectProps {
}
/** 格式化时长 mm:ss */
/** 获取素材实际时长(优先顶层 durationfallback 到 metadata.duration */
const getDuration = (item: AssetItem): number => {
return item.duration ?? (item.metadata?.duration as number) ?? 0
}
/** 获取素材实际文件大小 */
const getFileSize = (item: AssetItem): number => {
return item.file_size ?? (item.metadata?.file_size as number) ?? 0
}
const formatDuration = (seconds?: number): string => {
if (!seconds || seconds <= 0) return "00:00"
const m = Math.floor(seconds / 60)
@@ -99,7 +89,7 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
// 如果启用了时长校验,且配音时长不足
if (totalVideoDuration > 0) {
const material = materials.find((m) => m.id === id)
if (material && getDuration(material) < totalVideoDuration) {
if (material && (material.duration || 0) < totalVideoDuration) {
setPendingVoiceId(id)
setDurationWarningOpen(true)
return
@@ -281,24 +271,25 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
}}
>
<span style={{ display: "flex", alignItems: "center", gap: 4 }}>
{formatDuration(getDuration(item))}
{totalVideoDuration > 0 && getDuration(item) < Number(totalVideoDuration) && (
<span
style={{
color: "#ff4d4f",
fontSize: 11,
fontWeight: 500,
display: "inline-flex",
alignItems: "center",
gap: 2,
}}
>
<WarningOutlined />
</span>
)}
{formatDuration(item.duration)}
{totalVideoDuration > 0 &&
(Number(item.duration) || 0) < Number(totalVideoDuration) && (
<span
style={{
color: "#ff4d4f",
fontSize: 11,
fontWeight: 500,
display: "inline-flex",
alignItems: "center",
gap: 2,
}}
>
<WarningOutlined />
</span>
)}
</span>
<span>{formatFileSize(getFileSize(item))}</span>
<span>{formatFileSize(item.file_size)}</span>
</div>
</div>
)
@@ -327,9 +318,7 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
return (
<p>
<strong>
{pendingMaterial ? formatDuration(getDuration(pendingMaterial)) : "--"}
</strong>
<strong>{pendingMaterial ? formatDuration(pendingMaterial.duration) : "--"}</strong>
<strong>{formatDuration(totalVideoDuration)}</strong>
@@ -14,8 +14,6 @@ interface Step6CoverSettingsProps {
assetIds?: string[]
/** 当前选中的模板 ID */
selectedTemplate?: string
/** Step4 标题设置,用于预览视频烧录标题 & 封面叠加标题 */
titleSettings?: import("../types").TitleSettings
}
const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
@@ -42,7 +40,6 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
duration: props.duration,
assetIds: props.assetIds,
selectedTemplate: props.selectedTemplate,
titleSettings: props.titleSettings,
})
const handleAutoGenerate = () => {
@@ -2,7 +2,6 @@
* 标题预设样式网格
*/
import React from "react"
import { getFontFamily } from "../../constants"
interface TitlePresetItem {
key: string
@@ -36,7 +35,7 @@ const TitlePresetsGrid: React.FC<TitlePresetsGridProps> = ({
>
<span
className="xx-title-preset-preview-text"
style={{ ...p.previewStyle, fontFamily: getFontFamily(fontFamily || "思源黑体") }}
style={{ ...p.previewStyle, ...(fontFamily ? { fontFamily } : {}) }}
>
</span>
-15
View File
@@ -56,21 +56,6 @@ export const FONT_OPTIONS = [
"华康俪金黑",
]
/* ── 标题字体 CSS font-family 映射(中文显示名 → 浏览器可识别的字体栈) ── */
export const FONT_FAMILY_MAP: Record<string, string> = {
: '"Source Han Sans SC", "Noto Sans SC", "PingFang SC", "Microsoft YaHei", sans-serif',
: '"Source Han Serif SC", "Noto Serif SC", "Songti SC", "SimSun", serif',
: '"PingFang SC", -apple-system, "Helvetica Neue", sans-serif',
PingFang: '"PingFang SC", -apple-system, "Helvetica Neue", sans-serif',
: '"Microsoft YaHei", "PingFang SC", sans-serif',
: '"KaiTi", "STKaiti", "DFKai-SB", serif',
: '"华康俪金黑", "DFLiJinHei-W8", "Source Han Sans SC", "Microsoft YaHei", sans-serif',
}
export function getFontFamily(font: string): string {
return FONT_FAMILY_MAP[font] || FONT_FAMILY_MAP["思源黑体"]
}
/* ── 标题样式预设 ── */
export const TITLE_PRESETS = [
{
-2
View File
@@ -2309,8 +2309,6 @@
.xx-cover-preview-box {
position: relative;
aspect-ratio: 9 / 16;
max-width: 180px;
margin: 0 auto;
background: var(--bg-tertiary);
border-radius: var(--radius-md);
overflow: hidden;
@@ -19,8 +19,6 @@ export interface UseGenerateVideoProps {
autoSubtitles: boolean
bgm: boolean
generateCount: number
/** 当前草稿 IDURL 参数 edit_plan_id,用于后端回写任务关联) */
sourceEditPlanId?: string | null
}
/** 生成阶段 */
@@ -1,146 +1,87 @@
import { useRef, useCallback } from "react"
import { message } from "antd"
import axios from "axios"
import { getGenerationTask } from "@/api/tasks/tasks"
import { getGenerationTaskResults } from "@/api/template-editor"
import { getGenerationStatus, getGenerationTaskResults } from "@/api/template-editor"
import { safeExtractError } from "./errorUtils"
interface UseGenerationPollingOptions {
templateId: string
onProgress: (progress: number) => void
onComplete: (videos: unknown[]) => void
onFailed: (errorMsg: string) => void
}
/** 最大连续错误次数(仅对可重试错误),超过后终止轮询 */
const MAX_RETRYABLE_ERRORS = 10
/** 获取结果的最大重试次数 */
const MAX_RESULTS_RETRIES = 3
/**
* 生成状态轮询 Hookv2 — 改用 /generation/tasks/{task_id}
*
* 旧版轮询 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
* 生成状态轮询 Hook
* 轮询生成状态,更新进度,处理完成/失败
*/
export const useGenerationPolling = ({
templateId,
onProgress,
onComplete,
onFailed,
}: UseGenerationPollingOptions) => {
const progressTimer = useRef<ReturnType<typeof setTimeout>>()
const cancelledRef = useRef(false)
const clearTimer = useCallback(() => {
cancelledRef.current = true
if (progressTimer.current) {
clearTimeout(progressTimer.current)
progressTimer.current = undefined
}
}, [])
/** 任务完成后拉取结果列表,带重试 */
const fetchResultsWithRetry = useCallback(
async (taskId: string, attempt = 0): Promise<unknown[] | null> => {
const startPolling = useCallback(() => {
const poll = async () => {
try {
return await getGenerationTaskResults(taskId)
} catch (err) {
if (cancelledRef.current) return null
console.error(`[获取生成结果失败] 第 ${attempt + 1}`, err)
if (attempt < MAX_RESULTS_RETRIES - 1) {
await new Promise((resolve) => setTimeout(resolve, 1000 * (attempt + 1)))
return fetchResultsWithRetry(taskId, attempt + 1)
}
return null
}
},
[],
)
const data = await getGenerationStatus(templateId)
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
if (data.plan_status === "completed") {
onProgress(100)
// 获取生成的视频结果
let videos: unknown[] = []
if (data.generation_task_id) {
try {
videos = await getGenerationTaskResults(data.generation_task_id)
} catch (err) {
console.error("[获取生成结果失败]", err)
}
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) {
if (cancelledRef.current) return
console.error("[轮询出错] taskId:", taskId, pollErr)
// 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
}
consecutiveErrors += 1
if (consecutiveErrors >= MAX_RETRYABLE_ERRORS) {
const errorMsg = "任务状态查询连续失败,请稍后在任务列表查看结果"
onFailed(errorMsg)
message.error(errorMsg)
return
}
progressTimer.current = setTimeout(poll, 3000)
onComplete(videos)
message.success("视频生成完成!")
return
}
if (data.plan_status === "failed") {
const dataAny = data as unknown as Record<string, unknown>
const rawMsg =
dataAny.error_message ||
dataAny.error ||
dataAny.message ||
(Array.isArray(data.clips)
? (data.clips as { status: string; error_message?: string }[]).find(
(c) => c.status === "failed",
)?.error_message
: undefined) ||
"视频生成失败,请联系管理员或重试"
const errorMsg = safeExtractError(rawMsg)
console.error("[生成失败] templateId:", templateId, "响应:", data)
onFailed(errorMsg)
message.error(errorMsg)
return
}
}
progressTimer.current = setTimeout(poll, 1500)
},
[onProgress, onComplete, onFailed, fetchResultsWithRetry],
)
const clips = data.clips || []
const total = clips.length || 1
const done = (clips as { status: string }[]).filter((c) => c.status === "completed").length
onProgress(Math.round((done / total) * 100))
progressTimer.current = setTimeout(poll, 2000)
} catch (pollErr) {
console.error("[轮询出错] templateId:", templateId, pollErr)
progressTimer.current = setTimeout(poll, 3000)
}
}
progressTimer.current = setTimeout(poll, 2000)
}, [templateId, onProgress, onComplete, onFailed])
return { startPolling, clearTimer }
}
@@ -9,6 +9,8 @@ import { createFile } from "mp4box"
import type { Movie, Sample } from "mp4box"
// ── 常量 ──
/** 初始化预解码最大帧数(约 2 秒 @30fps),后续帧通过 decodeAroundPosition 按需解码 */
const MAX_INIT_FRAMES = 60
/**
* 规范化 mp4box 提取的 codec 字符串为 WebCodecs 兼容格式
@@ -109,7 +111,7 @@ class FrameQueue {
private frames: FrameEntry[] = []
private maxSize: number
constructor(maxSize = 200) {
constructor(maxSize = 5) {
this.maxSize = maxSize
}
@@ -121,36 +123,24 @@ class FrameQueue {
this.frames.push(entry)
}
/** 获取当前时间戳应显示的帧(二分查找,O(log n) */
/** 获取当前时间戳应显示的帧 */
getCurrentFrame(timestamp: number): VideoFrame | null {
if (this.frames.length === 0) return null
const target = timestamp + 0.01
// 找到最后一个 pts <= target 的帧(右边界)
let lo = 0,
hi = this.frames.length - 1,
bestIdx = -1
while (lo <= hi) {
const mid = (lo + hi) >> 1
if (this.frames[mid].pts <= target) {
bestIdx = mid
lo = mid + 1
} else {
hi = mid - 1
let best: FrameEntry | null = null
let bestIdx = -1
for (let i = 0; i < this.frames.length; i++) {
const f = this.frames[i]
if (f.pts <= timestamp + 0.01) {
best = f
bestIdx = i
}
}
if (bestIdx < 0) return null
// 关闭并移除 bestIdx 之前的所有已播放帧
for (let i = 0; i < bestIdx; i++) {
this.frames[i].frame.close()
}
this.frames.splice(0, bestIdx)
// 此时 bestIdx 对应帧已在索引 0
return this.frames[0]?.frame ?? null
if (bestIdx >= 0) {
this.frames = this.frames.slice(bestIdx)
}
return best?.frame ?? null
}
clear() {
@@ -239,10 +229,9 @@ export function useCanvasPlayer(
shadow?: boolean
},
onError?: (error: Error) => void,
enabled: boolean = true,
) {
const [state, setState] = useState<CanvasPlayerState>({
hasSupport: enabled && isWebCodecsSupported(),
hasSupport: isWebCodecsSupported(),
isPlaying: false,
currentTime: 0,
duration: 0,
@@ -254,13 +243,9 @@ export function useCanvasPlayer(
// ── 内部引用 ──
const decoderRef = useRef<VideoDecoder | null>(null)
const frameQueueRef = useRef(new FrameQueue(200))
/** 每个片段持久化解码器,避免每次新建导致关键帧错误 */
const segmentDecodersRef = useRef(new Map<number, VideoDecoder>())
/** 每个片段已送入解码器的 sample 游标(用于续解码) */
const segmentSampleCursorRef = useRef(new Map<number, number>())
/** 后台补充解码是否正在运行(防重入) */
const isFeedingRef = useRef(false)
const frameQueueRef = useRef(new FrameQueue(600))
/** 已解码的片段索引集合,用于按需解码(先标记防重入,失败时移除允许重试) */
const decodedSegmentsRef = useRef(new Set<number>())
/** 解码代数计数器,seek 时递增以作废正在进行的异步解码 */
const decodeGenerationRef = useRef(0)
const rafRef = useRef<number>(0)
@@ -490,91 +475,115 @@ export function useCanvasPlayer(
[segments, extractCodecDescription],
)
// ── 解码片段的一批帧(使用持久化解码器,支持从断点续解码) ──
/**
* @param segIdx 片段索引
* @param maxFrames 本次最多解码多少帧
*/
const decodeSegmentBatch = useCallback(
async (segIdx: number, maxFrames: number = 60): Promise<number> => {
if (isDestroyedRef.current) return 0
const metas = segmentMetaRef.current
const meta = metas[segIdx]
if (!meta) return 0
const buffer = segmentDataRef.current.get(meta.assetId)
if (!buffer) return 0
// ── 初始化 VideoDecoder 并解码指定片段 ──
const decodeSegment = useCallback(
async (_buffer: ArrayBuffer, meta: SegmentMeta, maxFrames?: number): Promise<void> => {
if (isDestroyedRef.current) return
const gen = decodeGenerationRef.current
let decoder = segmentDecodersRef.current.get(segIdx)
let cursor = segmentSampleCursorRef.current.get(segIdx) ?? 0
const samples = meta.samples
let decoderReady = false
// 如果还没有解码器,新建一个(从关键帧开始,不会报 key frame 错误
if (!decoder || decoder.state === "closed") {
decoder = new VideoDecoder({
output: (frame: VideoFrame) => {
if (videoDimRef.current.width === 0 || videoDimRef.current.height === 0) {
videoDimRef.current = { width: frame.codedWidth, height: frame.codedHeight }
}
const localTime = frame.timestamp / 1_000_000
const globalTime = localTime + meta.globalStartTime
frameQueueRef.current.push({
frame,
pts: globalTime,
duration: (frame.duration ?? 0) / 1_000_000,
})
},
error: (e: DOMException) => {
console.error(`[useCanvasPlayer] Segment ${segIdx} decoder error:`, e)
// 重置该片段的解码器和游标,允许重试
segmentDecodersRef.current.delete(segIdx)
segmentSampleCursorRef.current.set(segIdx, 0)
},
})
try {
await decoder.configure({
codec: meta.codec,
...(meta.description ? { description: meta.description } : {}),
// 配置解码器(每个片段可能需要不同的 codec/分辨率
const decoder = new VideoDecoder({
output: (frame: VideoFrame) => {
// 从第一帧获取实际尺寸
if (videoDimRef.current.width === 0 || videoDimRef.current.height === 0) {
videoDimRef.current = { width: frame.codedWidth, height: frame.codedHeight }
console.log(
`[useCanvasPlayer] Actual frame size: ${frame.codedWidth}x${frame.codedHeight}`,
)
}
const localTime = frame.timestamp / 1_000_000
const globalTime = localTime + meta.globalStartTime
frameQueueRef.current.push({
frame,
pts: globalTime,
duration: (frame.duration ?? 0) / 1_000_000,
})
} catch (err) {
console.error(`[useCanvasPlayer] Segment ${segIdx} configure failed:`, err)
const error = err instanceof Error ? err : new Error(String(err))
},
error: (e: DOMException) => {
console.error("[useCanvasPlayer] Decoder error callback:", e)
const error = new Error(`VideoDecoder error: ${e.message || e.name || "unknown"}`)
setState((s) => ({
...s,
isBuffering: false,
hasDecodeError: true,
errorMessage: `视频解码失败: ${error.message || "不支持的编解码器"}`,
errorMessage: `视频解码器错误: ${e.message || "解码异常"}`,
}))
onErrorRef.current?.(error)
return 0
}
},
})
segmentDecodersRef.current.set(segIdx, decoder)
cursor = 0
// 标记缓冲结束
console.log("[useCanvasPlayer] configure:", {
codec: meta.codec,
description: meta.description,
descriptionByteLength: meta.description?.byteLength,
videoWidth: meta.videoWidth,
videoHeight: meta.videoHeight,
codecCharCodes: meta.codec.split("").map((c) => c.charCodeAt(0)),
})
try {
await decoder.configure({
codec: meta.codec,
...(meta.description ? { description: meta.description } : {}),
})
decoderRef.current = decoder
decoderReady = true
// 标记缓冲结束,让 UI 开始渲染
setState((s) => ({ ...s, isBuffering: false }))
} catch (err) {
const error = err instanceof Error ? err : new Error(String(err))
console.error(
"[useCanvasPlayer] Decoder configure failed for segment:",
meta.assetId,
error,
)
console.error("[useCanvasPlayer] Failed codec config:", {
codec: meta.codec,
descriptionByteLength: meta.description?.byteLength,
videoWidth: meta.videoWidth,
videoHeight: meta.videoHeight,
})
setState((s) => ({
...s,
isBuffering: false,
isReady: false,
hasDecodeError: true,
errorMessage: `视频解码失败: ${error.message || "不支持的编解码器"}`,
}))
onErrorRef.current?.(error)
return
}
if (decoder.state !== "configured") return 0
if (!decoderReady) return
// 从 cursor 继续喂 sample(流水线批量提交,不 await 单个 decode
let decoded = 0
let si = cursor
while (si < samples.length && decoded < maxFrames) {
if (decodeGenerationRef.current !== gen || isDestroyedRef.current) break
if ((decoder.state as string) === "closed") break
// 帧队列快满时停止提交(这才是真正的背压)
if (frameQueueRef.current.size >= 180) break
// 解码器内部队列积压过多时短暂让出线程(阈值64,给硬件足够流水线深度)
if (decoder.decodeQueueSize > 64) {
await new Promise((r) => setTimeout(r, 5))
// 使用 demuxSegment 中已提取并过滤的 samples(前端切片
const samplesCollected = meta.samples
console.log(
`[useCanvasPlayer] Segment ${meta.assetId}: ${samplesCollected.length} samples to decode`,
)
if (samplesCollected.length === 0) {
console.warn("[useCanvasPlayer] No samples to decode for segment", meta.assetId)
return
}
// 送入解码器
let decodedCount = 0
let skippedCount = 0
let decodeErrors = 0
for (const sample of samplesCollected) {
if (!sample.data || isDestroyedRef.current) {
skippedCount++
continue
}
const sample = samples[si]
si++
if (!sample.data) continue
if (decoder.state === "closed") break
// 初始化阶段限制解码帧数,避免帧缓冲溢出
if (maxFrames && decodedCount >= maxFrames) {
console.log(
`[useCanvasPlayer] Segment ${meta.assetId}: init decode limited to ${maxFrames} frames`,
)
break
}
const chunk = new EncodedVideoChunk({
type: sample.is_sync ? "key" : "delta",
@@ -584,73 +593,93 @@ export function useCanvasPlayer(
})
try {
decoder.decode(chunk)
decoded++
await decoder.decode(chunk) // 修复:await 捕获异步错误
decodedCount++
} catch (e) {
console.warn(`[useCanvasPlayer] Segment ${segIdx} decode error:`, e)
break
}
}
cursor = si
segmentSampleCursorRef.current.set(segIdx, cursor)
// 等待解码器输出帧(最多500ms)
if (decoded > 0 && (decoder.state as string) === "configured") {
let waited = 0
while (frameQueueRef.current.size < Math.min(decoded, 10) && waited < 500) {
await new Promise((r) => setTimeout(r, 20))
waited += 20
if (isDestroyedRef.current) break
decodeErrors++
console.warn(`[useCanvasPlayer] Decode chunk error (${decodeErrors}):`, e)
// 连续 3 次解码失败,放弃当前片段并报告错误
if (decodeErrors >= 3) {
console.error("[useCanvasPlayer] Too many decode errors, aborting segment")
const error = new Error(`视频解码连续失败 ${decodeErrors} 次,片段: ${meta.assetId}`)
setState((s) => ({
...s,
isBuffering: false,
hasDecodeError: true,
errorMessage: `视频解码失败: 连续 ${decodeErrors} 次错误`,
}))
onErrorRef.current?.(error)
break
}
}
}
console.log(
`[useCanvasPlayer] Segment ${segIdx} decoded ${decoded} frames, queue size: ${frameQueueRef.current.size}`,
`[useCanvasPlayer] Segment ${meta.assetId}: decoded ${decodedCount}, skipped ${skippedCount}, errors ${decodeErrors}, decoder.state=${decoder.state}`,
)
return decoded
// flush 仅在解码器状态正常时执行
if (decoder.state === "configured") {
try {
await decoder.flush()
console.log(`[useCanvasPlayer] Segment ${meta.assetId}: flush complete`)
} catch (e) {
console.warn("[useCanvasPlayer] Decoder flush error:", e)
}
}
},
[],
)
// ── 后台持续补充帧 ──
/**
* 根据当前播放时间,确保队列中有足够缓冲
* 播放循环每 200ms 调用一次
* 按需解码当前播放位置 ±1 个片段。
* 在渲染循环中定期调用,避免一次性解码所有片段导致环形缓冲区溢出丢帧。
* 使用"先标记再解码"模式防止并发重复解码,失败时移除标记允许重试。
*/
const feedFrames = useCallback(
const decodeAroundPosition = useCallback(
async (currentTime: number) => {
if (isFeedingRef.current) return
isFeedingRef.current = true
try {
const metas = segmentMetaRef.current
if (!metas || metas.length === 0) return
const metas = segmentMetaRef.current
if (!metas || metas.length === 0) return
// 队列帧数充足时不解码(目标:保持 >= 80 帧缓冲)
if (frameQueueRef.current.size >= 80) return
// 记录当前代数,seek 后代数变化则中止
const gen = decodeGenerationRef.current
// 找到当前播放的片段
let targetIdx = 0
let acc = 0
for (let i = 0; i < metas.length; i++) {
const dur = metas[i].globalEndTime - metas[i].globalStartTime
if (currentTime < acc + dur) {
targetIdx = i
break
}
acc += dur
let targetIdx = -1
let acc = 0
for (let i = 0; i < metas.length; i++) {
const dur = metas[i].globalEndTime - metas[i].globalStartTime
if (currentTime < acc + dur) {
targetIdx = i
break
}
acc += dur
}
if (targetIdx === -1) targetIdx = metas.length - 1
// 依次补充:当前片段 → 下一个片段 → 再下一个
for (let offset = 0; offset <= 2; offset++) {
const idx = targetIdx + offset
if (idx >= metas.length) break
if (frameQueueRef.current.size >= 180) break
await decodeSegmentBatch(idx, 60)
for (
let i = Math.max(0, targetIdx - 1);
i <= Math.min(metas.length - 1, targetIdx + 1);
i++
) {
// seek 已作废当前解码任务
if (decodeGenerationRef.current !== gen) return
if (decodedSegmentsRef.current.has(i)) continue
const meta = metas[i]
const buffer = segmentDataRef.current.get(meta.assetId)
if (!buffer) continue
// 先标记为解码中,防止下一帧渲染时重复发起解码
decodedSegmentsRef.current.add(i)
try {
await decodeSegment(buffer, meta, 300)
} catch (e) {
// 解码失败则移除标记,允许后续重试
decodedSegmentsRef.current.delete(i)
console.warn(`[useCanvasPlayer] 按需解码片段 ${i} 失败:`, e)
}
} finally {
isFeedingRef.current = false
// await 后再次检查代数,seek 期间不更新标记
if (decodeGenerationRef.current !== gen) return
}
},
[decodeSegmentBatch],
[decodeSegment],
)
// ── 标题绘制 ──
@@ -667,6 +696,7 @@ export function useCanvasPlayer(
// 按 "/" 分割为多行("/" 作为手动换行符)
const lines = title.text.split(/[//⁄∕]/)
console.log("[drawTitle] 原始标题:", JSON.stringify(title.text), "分割后:", lines)
const lineHeight = fontSize * 1.3
const totalHeight = lines.length * lineHeight
@@ -778,10 +808,8 @@ export function useCanvasPlayer(
}
return s
})
// 后台补充帧:队列不足时自动续解码
if (frameQueueRef.current.size < 80) {
void feedFrames(currentTime)
}
// 按需解码当前 ±1 片段
decodeAroundPosition(currentTime)
}
if (currentTime >= totalDuration) {
@@ -790,46 +818,29 @@ export function useCanvasPlayer(
}
rafRef.current = requestAnimationFrame(renderFrame)
}, [canvasRef, totalDuration, titleSettings, drawTitle, computeDrawRect, feedFrames])
}, [canvasRef, totalDuration, titleSettings, drawTitle, computeDrawRect, decodeAroundPosition])
// ── 播放控制 ──
const play = useCallback(async () => {
if (!state.hasSupport || isDestroyedRef.current) return
// 重播:必须关闭旧解码器、清空队列、重置游标,从头重新解码
if (state.currentTime >= totalDuration - 0.1 || state.currentTime <= 0.1) {
// 重播场景:currentTime 已回到起点但 decodedSegmentsRef 仍有旧标记
// 此时 FrameQueue 中旧帧已被淘汰,需清空标记让 decodeAroundPosition 重新解码
if (state.currentTime <= 0.1 && decodedSegmentsRef.current.size > 0) {
decodeGenerationRef.current++
// 关闭所有持久化解码器
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
decodedSegmentsRef.current.clear()
// 同步清空帧缓冲,避免旧帧残留导致 getCurrentFrame 返回 null
frameQueueRef.current.clear()
playStartOffsetRef.current = 0
setState((s) => ({ ...s, currentTime: 0 }))
// 重新初始化解码
const metas = segmentMetaRef.current
const initialDecodeCount = Math.min(metas.length, 2)
for (let i = 0; i < initialDecodeCount; i++) {
await decodeSegmentBatch(i, 60)
}
}
setState((s) => ({ ...s, isPlaying: true }))
playStartRef.current = performance.now()
if (state.currentTime < 0.1) {
playStartOffsetRef.current = 0
} else {
playStartOffsetRef.current = state.currentTime
}
playStartOffsetRef.current = state.currentTime
lastProgressUpdateRef.current = 0
rafRef.current = requestAnimationFrame(renderFrame)
}, [state.hasSupport, state.currentTime, totalDuration, renderFrame, decodeSegmentBatch])
// 立即触发一次按需解码,不等渲染循环 200ms 节流
decodeAroundPosition(state.currentTime)
}, [state.hasSupport, state.currentTime, renderFrame, decodeAroundPosition])
const pause = useCallback(() => {
setState((s) => ({ ...s, isPlaying: false }))
@@ -839,61 +850,33 @@ export function useCanvasPlayer(
const seek = useCallback(
async (time: number) => {
const clampedTime = Math.max(0, Math.min(time, totalDuration))
decodeGenerationRef.current++
// 关闭所有解码器、清空队列、重置游标
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
frameQueueRef.current.clear()
setState((s) => ({ ...s, currentTime: clampedTime }))
playStartOffsetRef.current = clampedTime
playStartRef.current = performance.now()
// 找到 seek 目标片段,从该片段开始解码
const metas = segmentMetaRef.current
let targetIdx = 0,
acc = 0
for (let i = 0; i < metas.length; i++) {
const dur = metas[i].globalEndTime - metas[i].globalStartTime
if (clampedTime < acc + dur) {
targetIdx = i
break
}
acc += dur
}
await decodeSegmentBatch(targetIdx, 60)
await decodeSegmentBatch(Math.min(targetIdx + 1, metas.length - 1), 60)
// seek 时递增解码代数,作废正在进行的异步解码
decodeGenerationRef.current++
// 清空帧队列(clear 内部会 close 所有帧)+ 清空已解码标记
frameQueueRef.current.clear()
decodedSegmentsRef.current.clear()
await decodeAroundPosition(clampedTime)
},
[totalDuration, decodeSegmentBatch],
[totalDuration, decodeAroundPosition],
)
const destroy = useCallback(() => {
isDestroyedRef.current = true
cancelAnimationFrame(rafRef.current)
// 关闭所有持久化解码器
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
if (decoderRef.current && decoderRef.current.state !== "closed") {
decoderRef.current.close()
}
// 递增代数中止进行中的异步解码,清空帧队列(clear 内部 close 所有帧)
decodeGenerationRef.current++
frameQueueRef.current.clear()
segmentDataRef.current.clear()
segmentMetaRef.current = []
decodedSegmentsRef.current.clear()
}, [])
// ── 预加载下一个片段的数据 ──
@@ -910,7 +893,6 @@ export function useCanvasPlayer(
// ── 初始化:加载并解码所有片段 ──
useEffect(() => {
if (!enabled) return
if (!state.hasSupport || segments.length === 0) {
console.log("[useCanvasPlayer] Skip init:", {
hasSupport: state.hasSupport,
@@ -920,24 +902,10 @@ export function useCanvasPlayer(
}
let cancelled = false
console.log("[useCanvasPlayer] Init start, segments:", segments.length)
const init = async () => {
// ✅ 关键修复:重置销毁标记,允许新的 init 周期正常工作
// destroy() 在 useEffect cleanup 中被调用,将 isDestroyedRef 设为 true
// 如果不重置,后续的 loadSegment / decodeSegment 会立即 return
isDestroyedRef.current = false
// ✅ Strict Mode 修复:init 不再递增 generation
// seek() 和 play() 仍保留 generation 递增用于中止异步解码
// 重置错误状态,避免上一轮的解码错误影响新的 init 周期
setState((s) => ({
...s,
isBuffering: true,
hasDecodeError: false,
errorMessage: "",
isReady: false,
}))
console.log("[useCanvasPlayer] Init start v2_DIAG, segments:", segments.length)
setState((s) => ({ ...s, isBuffering: true }))
// 1. 加载所有片段数据
for (const seg of segments) {
@@ -979,44 +947,32 @@ export function useCanvasPlayer(
segmentMetaRef.current = metas
// 3. 初始化解码:关闭旧解码器,2 个片段各解 60 帧
// 后续由 feedFrames 后台补充
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
frameQueueRef.current.clear()
// 3. 按需解码:初始只解码3 个片段,后续通过 decodeAroundPosition 动态加载
// 避免一次性全量解码导致 frameQueue 环形缓冲区旧帧被丢弃引发黑屏
decodedSegmentsRef.current.clear()
const initGen = decodeGenerationRef.current
const initialDecodeCount = Math.min(metas.length, 2)
console.log(
`[useCanvasPlayer] Starting init decode: ${initialDecodeCount} segments, metas: ${metas.length}`,
)
const initialDecodeCount = Math.min(metas.length, 3)
for (let i = 0; i < initialDecodeCount; i++) {
if (cancelled) break
// seek 或 destroy 已作废当前初始化
if (decodeGenerationRef.current !== initGen) break
console.log(`[DIAG_v2] Init decode segment ${i}...`)
const meta = metas[i]
const buffer = segmentDataRef.current.get(meta.assetId)
if (!buffer) continue
// 先标记为解码中,防止重复解码
decodedSegmentsRef.current.add(i)
try {
await decodeSegmentBatch(i, 60)
console.log(`[DIAG_v2] Init decode segment ${i} done`)
await decodeSegment(buffer, meta, MAX_INIT_FRAMES)
} catch (e) {
// 解码失败则移除标记,允许后续重试
decodedSegmentsRef.current.delete(i)
console.warn(`[useCanvasPlayer] 初始化解码片段 ${i} 失败:`, e)
}
if (cancelled) break
}
console.log(`[useCanvasPlayer] Init decode finished, cancelled:`, cancelled)
if (!cancelled) {
console.log("[useCanvasPlayer] Init complete, isReady = true, duration:", totalDuration)
console.log("[useCanvasPlayer] Init complete, isReady = true")
setState((s) => ({ ...s, duration: totalDuration, isReady: true, isBuffering: false }))
} else {
console.warn("[useCanvasPlayer] Init was cancelled before completion")
}
}
@@ -1,124 +0,0 @@
/**
* 草稿自动保存工具 Hook
*
* 背景:后端 PUT /templates/{id}/editor 的 config 是「整体替换」语义,
* 直接发送 { config: { asset_ids } } 会把 title 等其他字段覆盖掉。
* 本 Hook 统一执行「GET 当前 config → 浅合并新字段 → PUT 回去」,
* 并用串行队列 + AbortController 保证:
* - 同一时刻只有一个保存请求在飞
* - 快速连续变化时只提交最后一次
* - 组件卸载时取消未完成请求
* - 保存失败时保留补丁,自动重试(指数退避,最多 5 次)
*
* 保存失败只 console.warn,不弹窗、不阻塞。
*/
import { useCallback, useEffect, useRef } from "react"
import { getEditPlan, updateEditPlan } from "@/api/template-editor"
type ConfigPatch = Record<string, unknown>
/** 最大自动重试次数 */
const MAX_RETRIES = 5
/** 初始重试延迟(ms),每次翻倍 */
const BASE_RETRY_DELAY = 1000
export function useDraftAutoSave(templateId?: string) {
const timerRef = useRef<ReturnType<typeof setTimeout>>()
const abortRef = useRef<AbortController | null>(null)
// 待合并的补丁队列(解决「保存进行中又来了新变化」)
const pendingPatchRef = useRef<ConfigPatch | null>(null)
const savingRef = useRef(false)
const templateIdRef = useRef(templateId)
templateIdRef.current = templateId
const flush = useCallback(async (retryCount = 0) => {
const tid = templateIdRef.current
if (!tid) return
// 已有保存在飞:把新补丁暂存,等当前请求结束后再合并一次
if (savingRef.current) return
// 快照当前补丁,但先不清空 —— 成功后才清除,失败时保留以便重试
const patchToSave = pendingPatchRef.current
if (!patchToSave) {
savingRef.current = false
return
}
savingRef.current = true
const controller = new AbortController()
abortRef.current = controller
try {
// 1. 读当前 config(拿最新,避免覆盖别人/别的步骤写入的字段)
const current = await getEditPlan(tid)
if (controller.signal.aborted) return
const merged = { ...(current.config || {}), ...patchToSave }
// 2. 写回完整合并后的 config
await updateEditPlan(tid, { config: merged }, controller.signal)
// 3. 保存成功才清除已保存的补丁
// (保存期间可能有新补丁进来,只清除我们已经保存的部分)
pendingPatchRef.current = null
} catch (err) {
const name = (err as { name?: string })?.name
if (name === "CanceledError" || name === "AbortError") {
// 组件卸载或新请求取消,不重试
return
}
console.warn("[useDraftAutoSave] 自动保存草稿失败:", err)
// 保存失败:把本次尝试保存的补丁合并回 pendingPatchRef
// (保存期间可能有新补丁,新补丁优先)
pendingPatchRef.current = {
...patchToSave,
...(pendingPatchRef.current || {}),
}
// 指数退避重试
if (retryCount < MAX_RETRIES && !controller.signal.aborted) {
const delay = BASE_RETRY_DELAY * Math.pow(2, retryCount)
timerRef.current = setTimeout(() => {
void flush(retryCount + 1)
}, delay)
}
// 超过最大重试次数后,补丁仍保留在 pendingPatchRef 中,
// 下次 scheduleSave 触发时会一起带上
} finally {
savingRef.current = false
// 保存期间又积累了新变化(且不是在重试路径中),再触发一次
if (pendingPatchRef.current && !controller.signal.aborted && retryCount === 0) {
timerRef.current = setTimeout(() => {
void flush()
}, 50)
}
}
}, [])
/**
* 调度一次自动保存(防抖)
* @param patch 要合并进 config 的局部字段
* @param delay 防抖毫秒数
*/
const scheduleSave = useCallback(
(patch: ConfigPatch, delay = 500) => {
const tid = templateIdRef.current
if (!tid) return
if (timerRef.current) clearTimeout(timerRef.current)
// 累计补丁(同一周期内多次变化合并成一次写入)
pendingPatchRef.current = { ...(pendingPatchRef.current || {}), ...patch }
timerRef.current = setTimeout(() => {
void flush()
}, delay)
},
[flush],
)
useEffect(() => {
return () => {
if (timerRef.current) clearTimeout(timerRef.current)
if (abortRef.current) abortRef.current.abort()
}
}, [])
return { scheduleSave }
}
export default useDraftAutoSave
@@ -34,6 +34,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
}, [])
const { startPolling, clearTimer } = useGenerationPolling({
templateId: selectedTemplate,
onProgress: handleProgress,
onComplete: handleComplete,
onFailed: handleFailed,
@@ -89,20 +90,16 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
const assetIds =
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
// 封面 URL:优先 AI 生成缩略图,兜底用户上传
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
// 直接创建正式生成任务
const taskResp = await createGenerationTask({
await createGenerationTask({
template_id: selectedTemplate,
asset_ids: assetIds,
output_width: outputWidth,
output_height: outputHeight,
cover_url: coverUrl,
cover_url: props.coverSettings?.upload_url || "",
custom_title: props.titleSettings?.title || "",
duration: props.duration || undefined,
video_ratio: props.videoRatio,
...(props.sourceEditPlanId ? { source_edit_plan_id: props.sourceEditPlanId } : {}),
...(props.titleSettings?.title
? {
title_config: {
@@ -119,12 +116,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
: {}),
})
// 从创建响应直接拿 task_id,改用新接口轮询
const taskId = taskResp.items?.[0]?.id
if (!taskId) {
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
}
startPolling(taskId)
startPolling()
} catch (err: unknown) {
console.error("[handleGenerate] 生成失败:", err)
setGenerating(false)
@@ -1,37 +1,55 @@
/**
* 素材片段调度器 Hook(多 video 元素方案 v3
*
* v3 修复:
* - 所有动态状态存入 ref,tick 为稳定函数,彻底消除 RAF 闭包陷阱
* - 片段切换时先启动下一个 video 再切可见性,消除冻屏间隔
* - 进度更新 200ms 节流
* 素材片段调度器 Hook(多 video 元素方案 v2
* 每个片段对应一个独立 <video> 元素,全部预加载,通过 display 切换实现无缝播放
* 替代单 video + 切 src 方案,消除片段切换延迟
*/
import { useState, useRef, useCallback, useEffect, useMemo } from "react"
/** 单个播放片段 */
export interface PlaybackSegment {
/** 素材 ID */
assetId: string
/** 素材视频 URL */
videoUrl: string
/** 片段在素材中的入点(秒) */
startTime: number
/** 片段在素材中的出点(秒) */
endTime: number
/** 片段在时间线中的顺序 */
order: number
}
/** 调度器返回 */
export interface SegmentSchedulerState {
/** 是否正在播放 */
isPlaying: boolean
/** 当前播放的全局时间(秒) */
currentTime: number
/** 总时长(秒) */
totalDuration: number
/** 当前片段索引 */
currentSegmentIndex: number
/** 当前片段的本地播放时间 */
segmentLocalTime: number
/** 是否已播完 */
isEnded: boolean
/** 是否可以播放(至少有 1 个片段) */
canPlay: boolean
/** 播放 */
play: () => void
/** 暂停 */
pause: () => void
/** 切换播放/暂停 */
togglePlayPause: () => void
/** 跳转到全局时间 */
seekTo: (time: number) => void
/** 每个片段对应的 video 元素 ref 数组 */
videoRefs: React.MutableRefObject<(HTMLVideoElement | null)[]>
}
/**
* 根据全局时间定位对应的片段和本地时间
*/
function findSegmentAtTime(
segments: PlaybackSegment[],
globalTime: number,
@@ -48,6 +66,9 @@ function findSegmentAtTime(
return { index: segments.length - 1, localTime: segments[segments.length - 1].endTime }
}
/**
* 计算每个片段的全局起始时间
*/
function buildTimeline(segments: PlaybackSegment[]): number[] {
const starts: number[] = []
let acc = 0
@@ -58,242 +79,224 @@ function buildTimeline(segments: PlaybackSegment[]): number[] {
return starts
}
/**
* useSegmentScheduler — 多 video 元素版素材片段调度器
*
* 核心改变:
* - 每个片段对应一个独立 <video> 元素(由组件渲染,ref 传入)
* - 所有 video 在挂载时即设置 src + preload="auto",浏览器自动预加载
* - 切换片段仅改 currentSegmentIndex + display,无需重新 load
* - 实现无缝切换,无加载延迟
*/
export function useSegmentScheduler(segments: PlaybackSegment[]): SegmentSchedulerState {
/** 每个片段对应的 video 元素 ref(由组件 JSX 渲染并绑定) */
const videoRefs = useRef<(HTMLVideoElement | null)[]>([])
const [isPlaying, setIsPlaying] = useState(false)
const [currentTime, setCurrentTime] = useState(0)
const [currentSegmentIndex, setCurrentSegmentIndex] = useState(0)
const [isEnded, setIsEnded] = useState(false)
const rafRef = useRef(0)
const rafRef = useRef<number>(0)
const isSeekingRef = useRef(false)
const lastTimeUpdateRef = useRef(0)
// 所有动态值存入 ref,tick 始终读取最新值,不依赖闭包
const segIdxRef = useRef(0)
const segmentsRef = useRef(segments)
const timelineStartsData = useMemo(() => buildTimeline(segments), [segments])
const totalDurationData = useMemo(
// 计算时间线
const timelineStarts = useMemo(() => buildTimeline(segments), [segments])
const totalDuration = useMemo(
() => segments.reduce((sum, seg) => sum + (seg.endTime - seg.startTime), 0),
[segments],
)
const timelineStartsRef = useRef(timelineStartsData)
const totalDurationRef = useRef(totalDurationData)
const isPlayingRef = useRef(false)
segmentsRef.current = segments
timelineStartsRef.current = timelineStartsData
totalDurationRef.current = totalDurationData
const canPlay = segments.length > 0
useEffect(() => {
segIdxRef.current = currentSegmentIndex
}, [currentSegmentIndex])
useEffect(() => {
isPlayingRef.current = isPlaying
}, [isPlaying])
const waitForReady = useCallback((video: HTMLVideoElement, timeout = 3000): Promise<void> => {
if (video.readyState >= 3) return Promise.resolve()
return new Promise((resolve) => {
const onCanPlay = () => {
video.removeEventListener("canplay", onCanPlay)
clearTimeout(timer)
resolve()
}
const timer = setTimeout(() => {
video.removeEventListener("canplay", onCanPlay)
resolve()
}, timeout)
video.addEventListener("canplay", onCanPlay)
})
}, [])
// 当前片段信息
const currentSegment = segments[currentSegmentIndex] || null
const segmentLocalTime = currentSegment
? currentTime - (timelineStarts[currentSegmentIndex] || 0) + currentSegment.startTime
: 0
/**
* 切换到指定片段
* 不改变 src(video 已在 JSX 中设置),仅 seek + 等待可播
*/
const switchToSegment = useCallback(
async (index: number, seekToLocalTime?: number) => {
const segs = segmentsRef.current
const video = videoRefs.current[index]
if (!video || index >= segs.length) return
(index: number, seekToLocalTime?: number): Promise<void> => {
return new Promise((resolve) => {
// 暂停当前视频
const prevVideo = videoRefs.current[currentSegmentIndex]
if (prevVideo) prevVideo.pause()
const seg = segs[index]
const localTime = seekToLocalTime ?? seg.startTime
const oldIdx = segIdxRef.current
const oldVideo = videoRefs.current[oldIdx]
const video = videoRefs.current[index]
if (!video || index >= segments.length) {
resolve()
return
}
if (oldVideo && oldVideo !== video) oldVideo.pause()
const seg = segments[index]
const localTime = seekToLocalTime ?? seg.startTime
if (!video.src && seg.videoUrl) {
video.src = seg.videoUrl
video.load()
}
if (Math.abs(video.currentTime - localTime) > 0.05) {
// 设置播放位置
video.currentTime = localTime
}
segIdxRef.current = index
setCurrentSegmentIndex(index)
// 如果已有足够帧数据,直接 resolve
if (video.readyState >= 2) {
setCurrentSegmentIndex(index)
resolve()
return
}
await waitForReady(video)
// 等待 canplay 事件
const onCanPlay = () => {
video.removeEventListener("canplay", onCanPlay)
clearTimeout(timeoutId)
setCurrentSegmentIndex(index)
resolve()
}
// 10 秒超时保护
const timeoutId = setTimeout(() => {
video.removeEventListener("canplay", onCanPlay)
console.warn(
`[useSegmentScheduler] 片段 ${index} 预加载超时 (10s), readyState=${video.readyState}`,
)
setCurrentSegmentIndex(index)
resolve()
}, 10000)
video.addEventListener("canplay", onCanPlay)
})
},
[waitForReady],
[segments, currentSegmentIndex],
)
// 稳定的 tick 函数,空依赖,所有值从 ref 读取
/** 播放循环 — 检测片段边界并切换 */
const tick = useCallback(() => {
const segs = segmentsRef.current
const idx = segIdxRef.current
const video = videoRefs.current[idx]
const video = videoRefs.current[currentSegmentIndex]
if (!video || isSeekingRef.current) {
rafRef.current = requestAnimationFrame(tick)
return
}
const seg = segs[idx]
const seg = segments[currentSegmentIndex]
if (!seg) return
// 预加载下一个片段
const nextIndex = idx + 1
if (nextIndex < segs.length) {
const nextVideo = videoRefs.current[nextIndex]
if (nextVideo) {
const timeToEnd = seg.endTime - video.currentTime
if (timeToEnd <= 2 && nextVideo.readyState < 3) {
const nextSeg = segs[nextIndex]
if (Math.abs(nextVideo.currentTime - nextSeg.startTime) > 0.5) {
nextVideo.currentTime = nextSeg.startTime
// 检查是否到达出点(容差 0.15s
if (video.currentTime >= seg.endTime - 0.15) {
video.pause()
const nextIndex = currentSegmentIndex + 1
if (nextIndex < segments.length) {
switchToSegment(nextIndex).then(() => {
setIsPlaying(true)
rafRef.current = requestAnimationFrame(tick)
const nextVideo = videoRefs.current[nextIndex]
if (nextVideo) {
const canPlay = () => {
nextVideo
.play()
.catch((e) =>
console.warn("[useSegmentScheduler] auto-play next segment failed:", e),
)
}
if (nextVideo.readyState >= 3) {
canPlay()
} else {
const timeout = setTimeout(canPlay, 300)
nextVideo.addEventListener(
"canplay",
() => {
clearTimeout(timeout)
canPlay()
},
{ once: true },
)
}
}
}
}
}
// 检测片段边界
if (video.currentTime >= seg.endTime - 0.1) {
if (nextIndex < segs.length) {
const nextVideo = videoRefs.current[nextIndex]
const nextSeg = segs[nextIndex]
const accumulatedTime =
(timelineStartsRef.current[idx] || 0) + (seg.endTime - seg.startTime)
if (nextVideo) {
if (Math.abs(nextVideo.currentTime - nextSeg.startTime) > 0.1) {
nextVideo.currentTime = nextSeg.startTime
}
// 先启动下一个视频(muted,可安全同时播放)
nextVideo
.play()
.catch((e) => console.warn("[useSegmentScheduler] next segment play failed:", e))
}
// 立即切换可见性
segIdxRef.current = nextIndex
setCurrentSegmentIndex(nextIndex)
setCurrentTime(accumulatedTime)
lastTimeUpdateRef.current = 0
setIsPlaying(true)
// 下一帧暂停旧视频(让新视频先渲染,避免冻屏)
const oldVideo = video
requestAnimationFrame(() => {
oldVideo.pause()
})
rafRef.current = requestAnimationFrame(tick)
return
const accumulatedTime =
(timelineStarts[currentSegmentIndex] || 0) + (seg.endTime - seg.startTime)
setCurrentTime(accumulatedTime)
} else {
video.pause()
setIsPlaying(false)
setIsEnded(true)
setCurrentTime(totalDurationRef.current)
setCurrentTime(totalDuration)
return
}
}
const globalTime = (timelineStartsRef.current[idx] || 0) + (video.currentTime - seg.startTime)
const now = performance.now()
if (now - lastTimeUpdateRef.current >= 200) {
lastTimeUpdateRef.current = now
setCurrentTime(Math.max(0, Math.min(globalTime, totalDurationRef.current)))
} else {
const globalTime =
(timelineStarts[currentSegmentIndex] || 0) + (video.currentTime - seg.startTime)
setCurrentTime(Math.max(0, Math.min(globalTime, totalDuration)))
}
rafRef.current = requestAnimationFrame(tick)
}, [])
}, [segments, currentSegmentIndex, timelineStarts, totalDuration, switchToSegment])
/** 播放 */
const play = useCallback(async () => {
if (!canPlay) return
setIsEnded(false)
const idx = segIdxRef.current
const video = videoRefs.current[idx]
if (!video) return
if (idx === 0 && video.readyState < 2) {
if (!video.src && segmentsRef.current[0]?.videoUrl) {
video.src = segmentsRef.current[0].videoUrl
video.load()
}
await waitForReady(video)
setIsEnded(false)
// 确保第一段可播放
const firstVideo = videoRefs.current[0]
if (firstVideo && currentSegmentIndex === 0 && firstVideo.readyState < 2) {
await switchToSegment(0)
}
const video = videoRefs.current[currentSegmentIndex]
if (!video) return
try {
await video.play()
const playPromise = video.play()
if (playPromise !== undefined) {
await playPromise
}
setIsPlaying(true)
cancelAnimationFrame(rafRef.current)
rafRef.current = requestAnimationFrame(tick)
} catch (err) {
console.warn("[useSegmentScheduler] 播放失败:", err)
}
}, [canPlay, waitForReady, tick])
}, [canPlay, switchToSegment, tick, currentSegmentIndex])
/** 暂停 */
const pause = useCallback(() => {
const video = videoRefs.current[segIdxRef.current]
const video = videoRefs.current[currentSegmentIndex]
if (video) video.pause()
setIsPlaying(false)
cancelAnimationFrame(rafRef.current)
}, [])
}, [currentSegmentIndex])
/** 切换播放/暂停 */
const togglePlayPause = useCallback(() => {
if (isPlayingRef.current) {
if (isPlaying) {
pause()
} else {
if (isEnded) {
// 播放结束后再次播放,从头开始
setIsEnded(false)
lastTimeUpdateRef.current = 0
const firstVideo = videoRefs.current[0]
if (firstVideo) {
videoRefs.current.forEach((v, i) => {
if (v && i !== 0) v.pause()
})
firstVideo.currentTime = segmentsRef.current[0]?.startTime || 0
segIdxRef.current = 0
setCurrentSegmentIndex(0)
setCurrentTime(0)
firstVideo
.play()
.then(() => {
setIsPlaying(true)
cancelAnimationFrame(rafRef.current)
rafRef.current = requestAnimationFrame(tick)
})
.catch((e) => console.warn("[useSegmentScheduler] restart failed:", e))
}
switchToSegment(0, segments[0]?.startTime).then(() => {
const video = videoRefs.current[0]
if (video) {
video.play().catch((e) => console.warn("[useSegmentScheduler] restart play failed:", e))
setIsPlaying(true)
setCurrentTime(0)
rafRef.current = requestAnimationFrame(tick)
}
})
} else {
play()
}
}
}, [isEnded, pause, play, tick])
}, [isPlaying, isEnded, pause, play, switchToSegment, segments, tick])
/** 跳转到指定全局时间 */
const seekTo = useCallback(
async (time: number) => {
if (!canPlay) return
const clampedTime = Math.max(0, Math.min(time, totalDurationRef.current))
const { index, localTime } = findSegmentAtTime(segmentsRef.current, clampedTime)
const clampedTime = Math.max(0, Math.min(time, totalDuration))
const { index, localTime } = findSegmentAtTime(segments, clampedTime)
isSeekingRef.current = true
cancelAnimationFrame(rafRef.current)
if (index !== segIdxRef.current) {
if (index !== currentSegmentIndex) {
await switchToSegment(index, localTime)
} else {
const video = videoRefs.current[index]
@@ -302,54 +305,48 @@ export function useSegmentScheduler(segments: PlaybackSegment[]): SegmentSchedul
setCurrentTime(clampedTime)
setIsEnded(false)
lastTimeUpdateRef.current = 0
if (isPlayingRef.current) {
const video = videoRefs.current[index]
if (video) {
video.play().catch(() => {})
}
rafRef.current = requestAnimationFrame(tick)
}
setTimeout(() => {
isSeekingRef.current = false
}, 200)
},
[canPlay, switchToSegment, tick],
[canPlay, totalDuration, segments, currentSegmentIndex, switchToSegment],
)
// 确保 videoRefs 数组长度与 segments 一致 + 强制预加载
useEffect(() => {
videoRefs.current = videoRefs.current.slice(0, segments.length)
while (videoRefs.current.length < segments.length) {
videoRefs.current.push(null)
}
// 强制预加载:所有 video 元素挂载后,调用 load() 确保浏览器真正开始加载数据
videoRefs.current.forEach((video) => {
if (video) {
video.load()
}
})
}, [segments])
// 组件卸载时清理
useEffect(() => {
return () => {
cancelAnimationFrame(rafRef.current)
}
}, [])
// 片段列表变化时重置
useEffect(() => {
cancelAnimationFrame(rafRef.current)
segIdxRef.current = 0
setIsPlaying(false)
setCurrentTime(0)
setCurrentSegmentIndex(0)
setIsEnded(false)
}, [segments])
const currentSegment = segments[currentSegmentIndex] || null
const segmentLocalTime = currentSegment
? currentTime - (timelineStartsRef.current[currentSegmentIndex] || 0) + currentSegment.startTime
: 0
return {
isPlaying,
currentTime,
totalDuration: totalDurationData,
totalDuration,
currentSegmentIndex,
segmentLocalTime,
isEnded,
@@ -6,7 +6,6 @@ import { useCallback, useEffect, useRef } from "react"
import { formatDuration } from "../utils/formatDuration"
import { useMaterialLibrary } from "./step2-materials/useMaterialLibrary"
import { useSmartMatch } from "./step2-materials/useSmartMatch"
import { useDraftAutoSave } from "./useDraftAutoSave"
interface UseStep2MaterialsProps {
materialMode: "manual" | "auto"
@@ -15,8 +14,6 @@ interface UseStep2MaterialsProps {
onSelectedMaterialsChange: (ids: string[]) => void
smartSelectedIds: string[]
onSmartSelectedIdsChange: (ids: string[]) => void
/** 当前选中的模板/草稿 ID,用于自动保存 */
selectedTemplate?: string
}
export function useStep2Materials({
@@ -26,7 +23,6 @@ export function useStep2Materials({
onSelectedMaterialsChange,
smartSelectedIds,
onSmartSelectedIdsChange,
selectedTemplate,
}: UseStep2MaterialsProps) {
const { libraries, selectedLibraryId, setSelectedLibraryId, materials, materialsLoading } =
useMaterialLibrary()
@@ -57,14 +53,6 @@ export function useStep2Materials({
handleSmartMatch()
}, [selectedLibraryId, materialMode, materialsLoading, materials.items, handleSmartMatch])
/* ── Step2 选择素材后自动保存草稿(防抖 500ms,失败静默) ── */
const { scheduleSave } = useDraftAutoSave(selectedTemplate)
useEffect(() => {
if (!selectedTemplate) return
const ids = materialMode === "auto" ? smartSelectedIds : selectedMaterials
scheduleSave({ asset_ids: ids }, 500)
}, [selectedTemplate, materialMode, selectedMaterials, smartSelectedIds, scheduleSave])
/* ── 手动选择素材 ── */
const handleToggleMaterial = useCallback(
(materialId: string) => {
@@ -4,24 +4,17 @@ import { getTitles } from "@/api/titles"
import type { TitleSettings } from "../../types"
import { useAiTitleGenerator } from "./useAiTitleGenerator"
import { useTitleStyleUpdaters } from "./useTitleStyleUpdaters"
import { useDraftAutoSave } from "../useDraftAutoSave"
interface UseStep4TitleProps {
titleSettings: TitleSettings
onTitleSettingsChange: (settings: TitleSettings) => void
/** 当前选中的模板/草稿 ID,用于自动保存 */
selectedTemplate?: string
}
/**
* Step 4 标题设置 Hook
* 封装 AI 标题生成、标题样式设置等逻辑
*/
export function useStep4Title({
titleSettings,
onTitleSettingsChange,
selectedTemplate,
}: UseStep4TitleProps) {
export function useStep4Title({ titleSettings, onTitleSettingsChange }: UseStep4TitleProps) {
// 标题库数据
const { data: userTitles = [] } = useQuery({
queryKey: ["titles"],
@@ -49,38 +42,6 @@ export function useStep4Title({
const prevAiAutoSelect = useRef(titleSettings.aiAutoSelect)
const isFirstMount = useRef(true)
/* ── Step4 标题内容/样式变化后自动保存草稿(防抖 800ms,失败静默) ── */
const { scheduleSave: scheduleTitleSave } = useDraftAutoSave(selectedTemplate)
useEffect(() => {
if (!selectedTemplate) return
scheduleTitleSave(
{
title: {
text: titleSettings.title,
font: titleSettings.font,
font_size: titleSettings.size,
color: titleSettings.color,
position: titleSettings.position,
bold: titleSettings.bold,
stroke: titleSettings.stroke,
shadow: titleSettings.shadow,
},
},
800,
)
}, [
selectedTemplate,
titleSettings.title,
titleSettings.font,
titleSettings.size,
titleSettings.color,
titleSettings.position,
titleSettings.bold,
titleSettings.stroke,
titleSettings.shadow,
scheduleTitleSave,
])
// 当 AI 自动选择开关打开时,自动生成/选择一个标题填入
// 首次挂载时如果开关已经是 true 且无标题,也需要触发
useEffect(() => {
@@ -7,8 +7,6 @@ import { message } from "antd"
import type { CoverConfig, CoverTemplate } from "../types/cover"
import { generateCover } from "@/api/generation"
import { createPreview, getPreviewStatus } from "@/api/generation/preview"
import { updateEditPlan } from "@/api/template-editor"
import type { TitleSettings } from "../types"
import {
fetchCoverTemplates,
createCoverTemplate,
@@ -24,8 +22,6 @@ interface UseStep6CoverProps {
assetIds?: string[]
/** 当前选中的模板 ID */
selectedTemplate?: string
/** Step4 标题设置,用于预览视频烧录标题 & 封面叠加标题 */
titleSettings?: TitleSettings
}
export function useStep6Cover({
@@ -34,7 +30,6 @@ export function useStep6Cover({
duration,
assetIds = [],
selectedTemplate = "",
titleSettings,
}: UseStep6CoverProps) {
const [generating, setGenerating] = useState(false)
@@ -97,20 +92,6 @@ export function useStep6Cover({
const response = await generateCover(selectedTemplate, {
asset_ids: assetIds,
cover_type: "ai_frame",
...(titleSettings?.title
? {
title_config: {
text: titleSettings.title,
font: titleSettings.font,
font_size: titleSettings.size,
font_color: titleSettings.color,
position: titleSettings.position,
bold: titleSettings.bold,
stroke: titleSettings.stroke,
shadow: titleSettings.shadow,
},
}
: {}),
})
clearTimeout(timeoutId)
const thumbnailUrl = response.cover?.image_url || ""
@@ -155,20 +136,6 @@ export function useStep6Cover({
template_id: selectedTemplate,
asset_ids: assetIds,
duration: duration || 30,
...(titleSettings?.title
? {
title_config: {
text: titleSettings.title,
font: titleSettings.font,
font_size: titleSettings.size,
font_color: titleSettings.color,
position: titleSettings.position,
bold: titleSettings.bold,
stroke: titleSettings.stroke,
shadow: titleSettings.shadow,
},
}
: {}),
})
// 轮询等待预览渲染完成:递归 setTimeout 避免请求重叠 + 120s 超时兜底
await new Promise<void>((resolve, reject) => {
@@ -187,21 +154,6 @@ export function useStep6Cover({
try {
const status = await getPreviewStatus(previewResp.task_id)
if (status.status === "completed") {
// 保存预览视频地址到 plan.config.rendered_storage_key
// 供封面 API 的 E1 兜底路径定位渲染后的视频(含标题烧录)。
// video_url 可能是完整 http(s) URL 或 OSS storage_key,两种格式后端都能处理。
if (status.video_url) {
try {
await updateEditPlan(selectedTemplate, {
config: { rendered_storage_key: status.video_url },
})
} catch (saveErr) {
console.warn(
"[Step6] 保存 rendered_storage_key 失败(不阻塞封面重试):",
saveErr,
)
}
}
done(() => resolve())
} else if (status.status === "failed") {
done(() => reject(new Error(status.error_message || "预览渲染失败")))
@@ -219,20 +171,6 @@ export function useStep6Cover({
const retryResp = await generateCover(selectedTemplate, {
asset_ids: assetIds,
cover_type: "ai_frame",
...(titleSettings?.title
? {
title_config: {
text: titleSettings.title,
font: titleSettings.font,
font_size: titleSettings.size,
font_color: titleSettings.color,
position: titleSettings.position,
bold: titleSettings.bold,
stroke: titleSettings.stroke,
shadow: titleSettings.shadow,
},
}
: {}),
})
const retryUrl = retryResp.cover?.image_url || ""
if (retryUrl) {
@@ -274,15 +212,7 @@ export function useStep6Cover({
clearTimeout(timeoutId)
setGenerating(false)
}
}, [
selectedTemplate,
assetIds,
coverSettings,
onCoverSettingsChange,
generating,
duration,
titleSettings,
])
}, [selectedTemplate, assetIds, coverSettings, onCoverSettingsChange, generating, duration])
// ── 模板操作方法 ──
const handleSelectTemplate = useCallback((id: string) => {
+2 -2
View File
@@ -32,7 +32,7 @@ describe("bgm API", () => {
describe("getBgmPresets", () => {
it("should resolve successfully", async () => {
await expect(getBgmPresets("test-template", { category: "test" })).resolves.not.toThrow()
await expect(getBgmPresets("test-params?")).resolves.not.toThrow()
})
it("should reject on API error", async () => {
@@ -42,7 +42,7 @@ describe("bgm API", () => {
mockDelete.mockRejectedValue(new Error("Network error"))
mockPatch.mockRejectedValue(new Error("Network error"))
await expect(getBgmPresets("test-template", { category: "test" })).rejects.toThrow()
await expect(getBgmPresets("test-params?")).rejects.toThrow()
})
})
})
+85
View File
@@ -1,12 +1,16 @@
import { describe, expect, it, vi, beforeEach } from "vitest"
import {
getEditPlans,
getEditPlan,
createEditPlan,
updateEditPlan,
deleteEditPlan,
generateEditPlan,
getGenerationStatus,
aiRecommendClips,
getEditPlanGenerations,
getGenerationTaskResults,
cancelGeneration,
getEditPlanClips,
getEditPlanClip,
createEditPlanClip,
@@ -15,6 +19,7 @@ import {
reorderEditPlanClips,
batchDeleteEditPlanClips,
createClipsFromAssets,
copyEditPlan,
getMediaAssets,
getMediaAsset,
} from "@/api/template-editor"
@@ -48,6 +53,22 @@ describe("editPlans API", () => {
mockPatch.mockResolvedValue({ data: { success: true, items: [] } })
})
describe("getEditPlans", () => {
it("should resolve successfully", async () => {
await expect(getEditPlans("test-params?")).resolves.not.toThrow()
})
it("should reject on API error", async () => {
mockGet.mockRejectedValue(new Error("Network error"))
mockPost.mockRejectedValue(new Error("Network error"))
mockPut.mockRejectedValue(new Error("Network error"))
mockDelete.mockRejectedValue(new Error("Network error"))
mockPatch.mockRejectedValue(new Error("Network error"))
await expect(getEditPlans("test-params?")).rejects.toThrow()
})
})
describe("getEditPlan", () => {
it("should resolve successfully", async () => {
await expect(getEditPlan("test-planId")).resolves.not.toThrow()
@@ -64,6 +85,22 @@ describe("editPlans API", () => {
})
})
describe("createEditPlan", () => {
it("should resolve successfully", async () => {
await expect(createEditPlan({ name: "test-item" })).resolves.not.toThrow()
})
it("should reject on API error", async () => {
mockGet.mockRejectedValue(new Error("Network error"))
mockPost.mockRejectedValue(new Error("Network error"))
mockPut.mockRejectedValue(new Error("Network error"))
mockDelete.mockRejectedValue(new Error("Network error"))
mockPatch.mockRejectedValue(new Error("Network error"))
await expect(createEditPlan({ name: "test-item" })).rejects.toThrow()
})
})
describe("updateEditPlan", () => {
it("should resolve successfully", async () => {
await expect(updateEditPlan("test-planId")).resolves.not.toThrow()
@@ -80,6 +117,22 @@ describe("editPlans API", () => {
})
})
describe("deleteEditPlan", () => {
it("should resolve successfully", async () => {
await expect(deleteEditPlan("test-planId")).resolves.not.toThrow()
})
it("should reject on API error", async () => {
mockGet.mockRejectedValue(new Error("Network error"))
mockPost.mockRejectedValue(new Error("Network error"))
mockPut.mockRejectedValue(new Error("Network error"))
mockDelete.mockRejectedValue(new Error("Network error"))
mockPatch.mockRejectedValue(new Error("Network error"))
await expect(deleteEditPlan("test-planId")).rejects.toThrow()
})
})
describe("generateEditPlan", () => {
it("should resolve successfully", async () => {
await expect(generateEditPlan("test-planId")).resolves.not.toThrow()
@@ -160,6 +213,22 @@ describe("editPlans API", () => {
})
})
describe("cancelGeneration", () => {
it("should resolve successfully", async () => {
await expect(cancelGeneration("test-planId")).resolves.not.toThrow()
})
it("should reject on API error", async () => {
mockGet.mockRejectedValue(new Error("Network error"))
mockPost.mockRejectedValue(new Error("Network error"))
mockPut.mockRejectedValue(new Error("Network error"))
mockDelete.mockRejectedValue(new Error("Network error"))
mockPatch.mockRejectedValue(new Error("Network error"))
await expect(cancelGeneration("test-planId")).rejects.toThrow()
})
})
describe("getEditPlanClips", () => {
it("should resolve successfully", async () => {
await expect(getEditPlanClips("test-planId")).resolves.not.toThrow()
@@ -288,6 +357,22 @@ describe("editPlans API", () => {
})
})
describe("copyEditPlan", () => {
it("should resolve successfully", async () => {
await expect(copyEditPlan("test-planId")).resolves.not.toThrow()
})
it("should reject on API error", async () => {
mockGet.mockRejectedValue(new Error("Network error"))
mockPost.mockRejectedValue(new Error("Network error"))
mockPut.mockRejectedValue(new Error("Network error"))
mockDelete.mockRejectedValue(new Error("Network error"))
mockPatch.mockRejectedValue(new Error("Network error"))
await expect(copyEditPlan("test-planId")).rejects.toThrow()
})
})
describe("getMediaAssets", () => {
it("should resolve successfully", async () => {
await expect(getMediaAssets("test-libraryId?")).resolves.not.toThrow()
+17
View File
@@ -5,6 +5,7 @@ import {
getTemplate,
toggleFavoriteTemplate,
copyTemplate,
generateFromTemplate,
} from "@/api/templates"
const mockGet = vi.fn()
@@ -115,4 +116,20 @@ describe("templates API", () => {
await expect(copyTemplate("test-templateId")).rejects.toThrow()
})
})
describe("generateFromTemplate", () => {
it("should resolve successfully", async () => {
await expect(generateFromTemplate("test-templateId")).resolves.not.toThrow()
})
it("should reject on API error", async () => {
mockGet.mockRejectedValue(new Error("Network error"))
mockPost.mockRejectedValue(new Error("Network error"))
mockPut.mockRejectedValue(new Error("Network error"))
mockDelete.mockRejectedValue(new Error("Network error"))
mockPatch.mockRejectedValue(new Error("Network error"))
await expect(generateFromTemplate("test-templateId")).rejects.toThrow()
})
})
})
@@ -150,10 +150,12 @@ vi.mock("@/api/template-editor", () => ({
getMediaAssets: vi.fn().mockResolvedValue({ items: [] }),
getEditPlanGenerations: vi.fn().mockResolvedValue({ items: [] }),
getEditPlan: vi.fn().mockResolvedValue({}),
createEditPlan: vi.fn().mockResolvedValue({ id: "test-plan" }),
updateEditPlan: vi.fn().mockResolvedValue({}),
generateEditPlan: vi.fn().mockResolvedValue({ task_id: "test-task" }),
getGenerationStatus: vi.fn().mockResolvedValue({ status: "completed" }),
getGenerationTaskResults: vi.fn().mockResolvedValue({ items: [] }),
cancelGeneration: vi.fn().mockResolvedValue({}),
getEditPlanClips: vi.fn().mockResolvedValue({ items: [] }),
createEditPlanClip: vi.fn().mockResolvedValue({}),
batchDeleteEditPlanClips: vi.fn().mockResolvedValue({}),
@@ -109,6 +109,7 @@ vi.mock("@/api/templates", () => ({
getTemplate: vi.fn().mockResolvedValue({ items: [], total: 0, success: true }),
toggleFavoriteTemplate: vi.fn().mockResolvedValue({ items: [], total: 0, success: true }),
copyTemplate: vi.fn().mockResolvedValue({ items: [], total: 0, success: true }),
generateFromTemplate: vi.fn().mockResolvedValue({ items: [], total: 0, success: true }),
}))
vi.mock("@/pages/templates/TemplateLibrary.css", () => ({}))
+1 -1
View File
@@ -7,7 +7,7 @@ VideoProcessor 等)按需从子模块导入,避免 __init__ 阶段引入
packages / DB 等重依赖。
"""
# 共享工具模块(零外部依赖,供 editing_modes / generation 等复用)
# 共享工具模块(零外部依赖,供 editing_modes / generation / edit_plan_generation 等复用)
from . import dedup_helpers, ffmpeg_utils, oss_helpers, url_security
__all__ = [
@@ -1,6 +1,6 @@
"""查重辅助函数 — 从 generation.py 提取的 GeneratedVideo 记录 + 查重逻辑.
供 generate_video 共同复用,
render_edit_plan 和 generate_video 共同复用,
创建 GeneratedVideo 记录后计算指纹并执行项目级 + 批次内查重。
"""
+1 -1
View File
@@ -1,4 +1,4 @@
"""OSS 工具函数 — 从 generation.py 提取的共享 OSS 操作.
"""OSS 工具函数 — 从 generation.py / edit_plan_generation.py 提取的共享 OSS 操作.
提供 OSS 配置读取、Bucket 创建、素材上传/下载、asset_id → 本地路径解析
等能力,供 render_edit_plan 和 generate_video 共同复用。
@@ -128,7 +128,6 @@ class RenderAdapter:
job_id: str = "",
work_dir: Path | None = None,
progress_cb: ProgressCallback | None = None,
voiceover_audio_path: str | None = None,
) -> RenderAdapterResult:
"""渲染一个 EditPlan。
@@ -143,7 +142,6 @@ class RenderAdapter:
job_id: 关联的 Job ID(用于结果存储路径)
work_dir: 工作目录,不传则使用临时目录
progress_cb: 进度回调函数
voiceover_audio_path: 配音音频本地路径(一键生成场景使用)
Returns:
RenderAdapterResult
@@ -210,7 +208,6 @@ class RenderAdapter:
progress_cb=progress_cb,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
voiceover_audio_path=voiceover_audio_path,
)
except subprocess.CalledProcessError as exc:
@@ -588,11 +585,14 @@ class RenderAdapter:
try:
from video_processing.thumbnail_generator import extract_and_upload_cover_frames
# 已渲染视频在统一渲染阶段已通过 ASS 字幕把标题烧录进画面,
# 抽帧天然带标题,因此这里传空字符串,避免 Pillow 二次叠加导致重影。
# Pillow 叠加仅用于 API 从源素材抽帧(源素材本身无标题)的兜底场景。
# 从 plan config 提取标题文字,叠加到封面候选帧上
_title_cfg = (plan_config or {}).get("title", {}) or {}
if not isinstance(_title_cfg, dict):
_title_cfg = {}
_title_text = (_title_cfg.get("text", "") or "").strip() if _title_cfg.get("enabled", True) else ""
cover_candidates = extract_and_upload_cover_frames(
str(result.output_path), plan_id, num_frames=3, title_text=""
str(result.output_path), plan_id, num_frames=3, title_text=_title_text
)
if cover_candidates:
logger.info(
@@ -1,8 +1,7 @@
"""视频封面抽帧工具 — 从视频中抽取帧作为封面,支持标题文字叠加
"""视频封面抽帧工具 — 从已渲染视频中抽取帧作为封面。
统一封面管道:
- 从已渲染视频抽帧:标题已通过 ASS 字幕烧进视频,帧天然带标题,无需再叠加
- 从源素材抽帧(API E2 兜底):源素材无标题,通过 Pillow 在帧上绘制标题文字。
统一封面管道:视频渲染时标题已通过 ASS 字幕烧进视频,
渲染完成后直接从此视频抽帧,封面天然带标题,无需额外叠加逻辑
"""
from __future__ import annotations
@@ -13,40 +12,6 @@ from pathlib import Path
logger = logging.getLogger(__name__)
# ── 标题叠加(Pillow)──────────────────────────────────────────────────────
# 实现统一放在 packages/shared/title_overlay.pyAPI 和 Worker 共用。
def apply_title_overlay(
image_path: str,
title_text: str,
*,
color: str = "#ffffff",
position: str = "bottom",
font_size: int | None = None,
margin_ratio: float = 0.06,
stroke_width_ratio: float = 0.04,
) -> str:
"""在图片上绘制标题文字(指定颜色 + 黑色描边/阴影)。
委托给 packages.shared.title_overlay.apply_title_to_image
保持 Worker 内调用方式不变。title_text 为空时直接返回原路径。
"""
from packages.shared.title_overlay import apply_title_to_image
if not title_text or not title_text.strip():
return image_path
result = apply_title_to_image(
image_path,
title_text,
color=color,
position=position,
font_size=font_size,
margin_ratio=margin_ratio,
stroke_width_ratio=stroke_width_ratio,
)
return result or image_path
def extract_first_frame(
video_path: str,
@@ -210,9 +175,6 @@ def extract_and_upload_cover_frames(
*,
num_frames: int = 3,
title_text: str = "",
title_color: str = "#ffffff",
title_position: str = "bottom",
title_font_size: int | None = None,
) -> list[dict]:
"""从视频中抽取多帧作为封面候选,上传到 OSS。
@@ -220,11 +182,7 @@ def extract_and_upload_cover_frames(
video_path: 视频文件路径
plan_id: 编辑计划 ID(用于生成 storage key
num_frames: 抽取帧数(默认 3
title_text: 标题文字;非空时用 Pillow 叠加到每帧。
从已渲染视频抽帧时通常传空(标题已烧录);从源素材抽帧时传标题。
title_color: 标题字体颜色(#RRGGBB
title_position: 标题位置 top/center/bottom
title_font_size: 标题字号,None 时自动计算
title_text: 标题文字(当前版本未叠加,预留参数)
Returns:
封面候选列表,每项包含 {"url": str, "position": float}
@@ -250,15 +208,6 @@ def extract_and_upload_cover_frames(
seek_ratio=ratio,
min_seek_seconds=0.5,
)
# 从源素材抽帧时叠加标题文字;已渲染视频标题已烧录时传空字符串跳过
if title_text and title_text.strip():
apply_title_overlay(
frame_path,
title_text,
color=title_color,
position=title_position,
font_size=title_font_size,
)
storage_key = f"covers/{plan_id}/frame_{i}.jpg"
url = upload_to_oss(frame_path, storage_key)
if url:
+2 -10
View File
@@ -14,17 +14,9 @@ celery_app.conf.imports = (
"worker_app.tasks.voice_extraction",
"worker_app.tasks.voice_clone",
"worker_app.tasks.tts_synthesis",
"worker_app.tasks.edit_plan_generation",
"worker_app.tasks.compose_video",
"worker_app.tasks.batch_download",
"worker_app.tasks._startup",
"apps.worker.video_processing.dedup",
"worker_app.tasks.cleanup",
)
# Celery Beat 定时任务调度
celery_app.conf.beat_schedule = {
"cleanup-stale-pending-tasks": {
"task": "worker.cleanup_stale_pending_tasks",
"schedule": 600.0, # 每 10 分钟(秒)
"options": {"expires": 300}, # 5 分钟过期,避免堆积
},
}
+5
View File
@@ -25,6 +25,10 @@ def __getattr__(name: str):
from .voice_extraction import extract_voice_task
return extract_voice_task
elif name == "compose_video":
from .compose_video import compose_video
return compose_video
elif name == "extract_background_task":
from .voice_extraction import extract_background_task
@@ -54,6 +58,7 @@ def __getattr__(name: str):
__all__ = [
"classify_asset",
"compose_video",
"generate_video",
"healthcheck",
"ingest_asset",
+3 -40
View File
@@ -10,9 +10,6 @@ logger = logging.getLogger(__name__)
# 孤儿任务超时阈值:渲染任务超过此时间未更新则视为卡死
ORPHAN_TASK_TIMEOUT_MINUTES = 10
# Pending 任务超时阈值:pending 任务在队列中等待超过此时间则自动清理
PENDING_TASK_TIMEOUT_MINUTES = 30
def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover
"""清理数据库中超时未更新的 running GenerationTask(孤儿任务)。
@@ -86,38 +83,6 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
return 0
def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover
"""清理数据库中卡在 pending 状态超时的 GenerationTask。
全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列,
导致新用户无法创建任务。通过 created_at 超时判断并标记为 failed。
Args:
timeout_minutes: 超时时间(分钟),默认 PENDING_TASK_TIMEOUT_MINUTES
Returns:
清理的任务数量
"""
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
session = SessionLocal()
try:
repo = SQLAlchemyGenerationTaskRepository(session)
count = repo.cleanup_stale_pending(timeout_minutes)
if count > 0:
logger.warning("清理了 %d 个超时的 pending GenerationTask(超过 %d 分钟未处理)", count, timeout_minutes)
else:
logger.info("无超时 pending GenerationTask 需要清理")
return count
except Exception as e:
logger.error("清理超时 pending GenerationTask 失败: %s", e, exc_info=True)
return 0
finally:
session.close()
def cleanup_all_stale_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict: # pragma: no cover
"""统一清理所有超时的孤儿任务。
@@ -128,17 +93,15 @@ def cleanup_all_stale_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES)
"""
gen_count = cleanup_orphan_tasks(timeout_minutes)
job_count = cleanup_stale_jobs(timeout_minutes)
pending_count = cleanup_stale_pending_tasks(PENDING_TASK_TIMEOUT_MINUTES)
total = gen_count + job_count + pending_count
total = gen_count + job_count
if total > 0:
logger.warning(
"任务清理完成: 孤儿 GenerationTask=%d, 孤儿 Job=%d, 超时 pending=%d, 总计=%d",
"孤儿任务清理完成: GenerationTask=%d, Job=%d, 总计=%d",
gen_count,
job_count,
pending_count,
total,
)
return {"generation_tasks": gen_count, "jobs": job_count, "pending": pending_count}
return {"generation_tasks": gen_count, "jobs": job_count}
@worker_ready.connect
-35
View File
@@ -1,35 +0,0 @@
"""定期清理任务 — Celery Beat 调度。
包含:
- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks
"""
import logging
from celery import shared_task
from worker_app.tasks._startup import (
PENDING_TASK_TIMEOUT_MINUTES,
cleanup_stale_pending_tasks,
)
logger = logging.getLogger(__name__)
@shared_task(name="worker.cleanup_stale_pending_tasks")
def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINUTES) -> dict:
"""Celery Beat 调度的定期任务:清理超时的 pending 任务。
每 10 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
查找所有 status='pending' 且 created_at < NOW() - timeout_minutes
的 generation_tasks,批量更新为 failed。
Args:
timeout_minutes: 超时时间(分钟),默认 30 分钟
Returns:
{"cleaned": int}
"""
count = cleanup_stale_pending_tasks(timeout_minutes)
if count > 0:
logger.info("[Beat] 清理了 %d 个超时 pending 任务(超时阈值 %d 分钟)", count, timeout_minutes)
return {"cleaned": count}
@@ -0,0 +1,198 @@
"""视频合成 Celery 任务 — Phase 8 任务 2.10.
使用 JobService 管理任务生命周期,通过 RenderAdapter 调用 UnifiedRenderService 执行合成。
"""
from __future__ import annotations
import os
import subprocess # pragma: no cover
import tempfile
from pathlib import Path
from celery.exceptions import SoftTimeLimitExceeded # pragma: no cover
from celery.utils.log import get_task_logger
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
logger = get_task_logger(__name__)
# 任务超时时间(秒):超过此时间 Celery 会抛出 SoftTimeLimitExceeded
RENDER_TASK_SOFT_TIME_LIMIT = 600 # pragma: no cover # 10 分钟
# 硬超时:超过此时间进程会被强制 kill
RENDER_TASK_TIME_LIMIT = 660 # pragma: no cover # 10 分钟 + 1 分钟清理缓冲
def _get_job_service():
"""延迟导入 JobService,避免循环依赖。"""
from apps.api.app.services.job_service import JobService
from packages.adapters.sqlalchemy_impl.job_repository import SQLAlchemyJobRepository
db = SessionLocal()
repo = SQLAlchemyJobRepository(db)
return JobService(repo), db
@celery_app.task(
name="worker.compose_video",
bind=True,
max_retries=3,
default_retry_delay=60,
soft_time_limit=RENDER_TASK_SOFT_TIME_LIMIT,
time_limit=RENDER_TASK_TIME_LIMIT,
)
def compose_video(self, job_id: str, **kwargs): # pragma: no cover
"""视频合成任务。
使用 UnifiedRenderService(图层架构)进行渲染。
Args:
job_id: JobService 中的任务 ID
**kwargs: 来自 Job.payload 的额外参数(plan_id, output_path 等)
"""
job_service, db = _get_job_service()
try:
job = job_service.get_job(job_id)
if job is None:
logger.error("Job not found: %s", job_id)
return {"status": "error", "message": f"Job {job_id} not found"}
plan_id = job.payload.get("plan_id", "")
if not plan_id:
job_service.fail_job(job_id, "Missing plan_id in job payload")
return {"status": "error", "message": "Missing plan_id"}
# 使用 unified 渲染引擎
return _compose_with_unified_engine(self, job_service, job, plan_id, db)
except SoftTimeLimitExceeded:
# Celery 软超时:任务执行超过 soft_time_limit
error_msg = f"渲染任务超时(超过 {RENDER_TASK_SOFT_TIME_LIMIT // 60} 分钟)"
logger.error("视频合成超时: job_id=%s", job_id)
try:
job_service.fail_job(job_id, error_msg)
except Exception:
logger.exception("更新 Job 超时失败状态时出错")
# 超时不重试
return {"status": "error", "message": error_msg, "error_type": "timeout"}
except subprocess.TimeoutExpired as exc:
# FFmpeg 子进程超时
error_msg = f"FFmpeg 渲染超时({exc.timeout}s"
logger.error("视频合成 FFmpeg 超时: job_id=%s timeout=%s", job_id, exc.timeout)
try:
job_service.fail_job(job_id, error_msg)
except Exception:
logger.exception("更新 Job 超时失败状态时出错")
# 超时不重试
return {"status": "error", "message": error_msg, "error_type": "ffmpeg_timeout"}
except self.retry_exc as exc:
logger.warning("视频合成重试中: job_id=%s, exc=%s", job_id, exc)
raise
except subprocess.CalledProcessError as exc:
# FFmpeg 执行失败,提取有意义的错误信息
from video_processing.video_validation import get_exit_code_message
exit_msg = get_exit_code_message(exc.returncode)
stderr_text = (exc.stderr or "").strip()
stderr_tail = stderr_text[-300:] if len(stderr_text) > 300 else stderr_text
error_msg = f"渲染失败: {exit_msg}"
if stderr_tail:
error_msg += f" | {stderr_tail[:200]}"
logger.error("视频合成 FFmpeg 失败: job_id=%s %s", job_id, exit_msg)
try:
job_service.fail_job(job_id, error_msg[:500])
except Exception:
logger.exception("更新 Job 失败状态时出错")
# FFmpeg 错误不重试(通常是素材或配置问题)
return {"status": "error", "message": error_msg, "error_type": "ffmpeg_error", "exit_code": exc.returncode}
except Exception as exc:
logger.exception("视频合成异常: job_id=%s", job_id)
try:
job_service.fail_job(job_id, str(exc)[:500])
except Exception:
logger.exception("更新 Job 失败状态时出错")
raise self.retry(exc=exc, countdown=60) from exc
finally:
db.close()
def _compose_with_unified_engine(task, job_service, job, plan_id: str, db) -> dict: # pragma: no cover
"""新引擎渲染路径(UnifiedRenderService + RenderAdapter)。"""
job_id = job.id
# 标记为 running
job_service.update_progress(job_id, progress=10.0, current_stage="初始化统一渲染引擎")
from video_processing.render_adapter import RenderAdapter
adapter = RenderAdapter(db)
# 校验合成条件
job_service.update_progress(job_id, progress=15.0, current_stage="校验合成条件")
valid, errors, warnings, ready_count, total_count = adapter.validate_plan(plan_id)
if not valid:
error_msg = "; ".join(errors)
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 进度回调
def progress_cb(progress: float, stage: str) -> None:
try:
job_service.update_progress(job_id, progress=progress, current_stage=stage)
except Exception:
logger.exception("更新进度失败")
# 执行渲染
job_service.update_progress(job_id, progress=20.0, current_stage="开始渲染")
logger.info("统一渲染引擎开始: job_id=%s plan_id=%s", job_id, plan_id)
result = adapter.render_plan(
plan_id=plan_id,
job_id=job_id,
progress_cb=progress_cb,
)
if not result.success:
error_msg = f"渲染失败: {result.error_message}"
job_service.fail_job(job_id, error_msg[:500])
raise RuntimeError(result.error_message)
# 更新 Job 状态为完成
result_data = {
"plan_id": plan_id,
"output_path": str(result.output_path) if result.output_path else "",
"storage_key": f"rendered/{plan_id}/{job_id}.mp4",
"output_url": result.output_url,
"estimated_duration": result.duration,
"clip_count": result.clip_count,
"engine": "unified",
"width": result.width,
"height": result.height,
"file_size": result.file_size,
}
job_service.complete_job(job_id, result=result_data)
logger.info(
"视频合成完成(unified): job_id=%s plan_id=%s duration=%.2fs",
job_id,
plan_id,
result.duration,
)
return {"status": "completed", "job_id": job_id, "result": result_data}
def _cleanup_output(job_id: str) -> None:
"""清理临时输出文件。"""
try:
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
if Path(output_path).exists():
Path(output_path).unlink()
except Exception as e:
logger.warning(f"清理输出文件失败: {e}", exc_info=True)
@@ -0,0 +1,452 @@
"""剪辑计划渲染任务 — 使用 UnifiedRenderService 统一渲染引擎.
Celery 任务 worker.render_edit_plan:
1. 加载 EditPlan + EditPlanClips
2. 通过 RenderAdapter 调用 UnifiedRenderService 渲染
3. 下载各片段素材 + 渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan / EditPlanClip 状态
7. 更新 GenerationTask 进度
"""
from __future__ import annotations
import logging
from datetime import datetime, timezone
from pathlib import Path
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
logger = logging.getLogger(__name__)
OUTPUT_WIDTH = 1280
OUTPUT_HEIGHT = 720
OUTPUT_FPS = 25.0
# ── 共享工具模块导入 ──────────────────────────────────────────────────────────
from video_processing.dedup_helpers import create_video_record_and_dedup
# ── Repository imports (延迟导入避免循环依赖) ─────────────────────────────────
def _get_repos():
"""获取数据库仓储实例"""
from packages.adapters.sqlalchemy_impl import SQLAlchemyEditPlanRepository
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import (
SQLAlchemyEditPlanClipRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
db = SessionLocal()
try:
plan_repo = SQLAlchemyEditPlanRepository(db)
clip_repo = SQLAlchemyEditPlanClipRepository(db)
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
yield plan_repo, clip_repo, gen_task_repo, db
finally:
db.close()
# ── Celery Task ───────────────────────────────────────────────────────────────
def _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg: str):
"""统一的计划失败标记工具。"""
plan = plan_repo.get(plan_id)
if plan and plan.status.value == "rendering":
plan.mark_failed()
plan_repo.update(plan)
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task and gen_task.status.value != "failed":
gen_task.status = "failed"
gen_task.error_message = error_msg
gen_task.completed_at = datetime.now(timezone.utc)
try:
gen_task.append_log(
stage="render_failed",
message=error_msg[:500],
level="ERROR",
)
except Exception:
pass
gen_task_repo.update(gen_task)
def _finalize_render_success(
plan,
plan_repo,
clip_repo,
gen_task_repo,
db,
plan_id: str,
output_url: str,
storage_key: str,
duration: float,
file_size: int,
width: int,
height: int,
rendered_clip_ids: list[str],
failed_clip_ids: list[str],
generation_task_id: str,
output_path: Path,
engine: str,
thumbnail_url: str = "",
cover_candidates: list[dict] | None = None,
) -> dict:
"""渲染成功后的统一收尾:查重 + 更新状态 + 返回结果。"""
# 创建 GeneratedVideo 记录 + 查重
project_id = plan.project_id or ""
batch_id = plan.config.get("batch_id", "")
mode = plan.config.get("mode", "edit_plan")
# 从 plan.config.title.text 读取视频名称
plan_config = plan.config or {}
title_cfg = plan_config.get("title", {}) or {}
if not isinstance(title_cfg, dict):
title_cfg = {}
video_name = (title_cfg.get("text") or "").strip() or f"generated-{generation_task_id[:8]}.mp4"
if generation_task_id:
try:
create_video_record_and_dedup(
generation_task_id=generation_task_id,
project_id=project_id,
user_id=plan.created_by_user_id or "",
batch_id=batch_id,
file_url=output_url or "",
file_size=file_size,
duration=duration,
video_path=str(output_path),
mode=mode,
session=db,
width=width,
height=height,
fps=OUTPUT_FPS,
name=video_name,
thumbnail_url=thumbnail_url,
)
except Exception as dedup_err:
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
# 更新片段状态为 rendered
for clip_id in rendered_clip_ids:
clip = clip_repo.get(clip_id)
if clip and clip.status.value == "ready":
clip.mark_rendered()
clip_repo.update(clip)
# 更新 EditPlan 状态为 completed + 回写实际渲染时长 + 结果数
plan.config["rendered_url"] = output_url or ""
plan.config["rendered_storage_key"] = storage_key
if hasattr(plan, "total_duration") and duration > 0:
plan.total_duration = duration
plan.mark_completed()
plan_repo.update(plan)
# 更新 GenerationTask 状态为 completed
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "completed"
gen_task.progress = 100.0
# 剪辑计划是多片段合成 1 个成片,result_count = 1
gen_task.result_count = 1
gen_task.append_log(
stage="render_complete",
message=f"渲染完成,输出时长 {duration:.1f}s",
level="INFO",
engine=engine,
clip_count=len(rendered_clip_ids),
)
gen_task.completed_at = datetime.now(timezone.utc)
# 回写封面 URL 到 GenerationTask,供封面生成接口读取
if cover_candidates:
first_cover = cover_candidates[0].get("image_url") or cover_candidates[0].get("url") or ""
if first_cover:
gen_task.cover_url = first_cover
logger.info(
"预览渲染完成,回写 cover_url: plan_id=%s task_id=%s url=%s",
plan_id,
generation_task_id,
first_cover[:80],
)
gen_task_repo.update(gen_task)
logger.info(
"剪辑计划渲染完成: plan_id=%s engine=%s rendered=%d failed=%d duration=%.1fs",
plan_id,
engine,
len(rendered_clip_ids),
len(failed_clip_ids),
duration,
)
return {
"status": "completed",
"plan_id": plan_id,
"rendered_count": len(rendered_clip_ids),
"failed_count": len(failed_clip_ids),
"output_url": output_url,
"duration": duration,
}
def _render_with_unified(
plan,
clips,
plan_id: str,
generation_task_id: str,
plan_repo,
clip_repo,
gen_task_repo,
db,
) -> dict:
"""统一渲染引擎路径(通过 RenderAdapter 调用 UnifiedRenderService)。
RenderAdapter 内部处理:素材下载、BGM 准备、ASR 自动字幕、渲染执行、OSS 上传。
本函数只负责:业务状态更新、查重、收尾。
"""
from video_processing.render_adapter import RenderAdapter
adapter = RenderAdapter(db)
# 进度回调:更新 GenerationTask 进度
def _progress_cb(progress: float, stage: str):
if not generation_task_id:
return
try:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
# 映射到 30%~90% 区间(素材下载前已到 30%)
mapped_progress = 30.0 + progress * 0.6
gen_task.progress = min(mapped_progress, 95.0)
gen_task.append_log(
stage="render_progress",
message=stage,
level="INFO",
progress=mapped_progress,
)
gen_task_repo.update(gen_task)
except Exception:
pass
try:
result = adapter.render_plan(
plan_id=plan_id,
job_id=generation_task_id or plan_id,
progress_cb=_progress_cb,
)
except Exception as render_err:
logger.error("渲染失败(unified): %s%s", plan_id, render_err)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"渲染失败: {render_err}")
return {"status": "error", "message": f"渲染失败: {render_err}"}
if not result.success:
full_error = result.error_message or "渲染失败"
if result.error_detail:
full_error = f"{full_error}\n--- stderr ---\n{result.error_detail}"
logger.error("渲染失败(unified): %s%s", plan_id, result.error_message)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, full_error)
return {"status": "error", "message": result.error_message or "渲染失败"}
output_path = result.output_path or Path("")
output_url = result.output_url or ""
thumbnail_url = result.thumbnail_url or ""
# adapter 上传到 rendered/{plan_id}/{job_id}.mp4,从 URL 提取实际 key
# 不能用 output.mp4 硬编码,否则 cover 等下游通过 key 构造的 URL 指向不存在的文件
if output_url:
from packages.shared.storage import get_shared_storage_service
storage_key = get_shared_storage_service().normalize_storage_key(output_url)
else:
storage_key = f"rendered/{plan_id}/{generation_task_id or plan_id}.mp4"
# 用 adapter 返回的 clip 明细(以 adapter 的结果为准)
rendered_clip_ids = result.rendered_clip_ids or []
failed_clip_ids = result.failed_clip_ids or []
# 将封面候选帧写入 plan.config(供封面 API 直接使用,跳过 MediaKit 抽帧)
if result.cover_candidates:
plan_config = plan.config or {}
plan_config["cover_candidates"] = result.cover_candidates
plan.config = plan_config
logger.info(
"封面候选帧已写入 plan.config: plan_id=%s count=%d",
plan_id,
len(result.cover_candidates),
)
return _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id=plan_id,
output_url=output_url or "",
storage_key=storage_key,
duration=result.duration,
file_size=result.file_size,
width=result.width,
height=result.height,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
generation_task_id=generation_task_id,
output_path=output_path,
engine="unified",
thumbnail_url=thumbnail_url,
cover_candidates=result.cover_candidates,
)
@celery_app.task(
name="worker.render_edit_plan",
bind=True,
max_retries=2,
soft_time_limit=600, # 10 分钟软超时
time_limit=660, # 11 分钟硬超时
)
def render_edit_plan(self, plan_id: str) -> dict: # pragma: no cover
"""渲染剪辑计划
流程:
1. 加载 EditPlan + EditPlanClips
2. 通过 RenderAdapter 调用 UnifiedRenderService 渲染
3. 下载素材 + 渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan → completed, EditPlanClips → rendered
7. 更新 GenerationTask 进度
"""
logger.info("开始渲染剪辑计划: plan_id=%s", plan_id)
generation_task_id = ""
for repos in _get_repos():
plan_repo, clip_repo, gen_task_repo, db = repos
try:
# 1. 加载 EditPlan
plan = plan_repo.get(plan_id)
if plan is None:
logger.error("剪辑计划不存在: %s", plan_id)
return {"status": "error", "message": f"计划不存在: {plan_id}"}
# 获取 generation_task_id(提前读取,确保 except 块可用)
generation_task_id = plan.config.get("generation_task_id", "")
# 2. 准备渲染(使用 unified 渲染引擎)
# 3. 加载片段列表(按 order 排序)
clips = clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
if not clips:
logger.warning("剪辑计划没有片段: %s", plan_id)
plan.mark_failed()
plan_repo.update(plan)
return {"status": "error", "message": "没有可渲染的片段"}
# 更新 GenerationTask 状态为 running
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "running"
gen_task.started_at = datetime.now(timezone.utc)
gen_task.append_log(
stage="render_start",
message=f"开始渲染,片段数 {len(clips)}",
level="INFO",
engine="unified",
clip_count=len(clips),
)
gen_task_repo.update(gen_task)
# 3. 渲染前取消检查
if generation_task_id:
current_task = gen_task_repo.get(generation_task_id)
if current_task:
task_status = (
current_task.status.value if hasattr(current_task.status, "value") else str(current_task.status)
)
if task_status == "cancelled":
logger.info("任务已被取消,中止渲染: plan_id=%s task_id=%s", plan_id, generation_task_id)
if plan.status.value == "rendering":
try:
plan.resume_editing()
plan_repo.update(plan)
except ValueError:
pass
return {"status": "cancelled", "plan_id": plan_id, "message": "任务已取消"}
# 4. 渲染(unified 引擎:RenderAdapter 统一处理下载 + BGM + ASR + 渲染 + 上传)
result = _render_with_unified(
plan=plan,
clips=clips,
plan_id=plan_id,
generation_task_id=generation_task_id,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
)
result["engine"] = "unified"
return result
except Exception as exc:
# 超时异常不重试,直接标记失败
from celery.exceptions import SoftTimeLimitExceeded
is_timeout = isinstance(exc, SoftTimeLimitExceeded)
if is_timeout:
logger.error("渲染剪辑计划超时: plan_id=%s", plan_id)
else:
logger.exception("渲染剪辑计划异常: %s", plan_id)
# 尝试标记计划和 GenerationTask 为失败
error_msg = "渲染任务超时(超过10分钟)" if is_timeout else f"渲染异常: {type(exc).__name__}: {exc}"
try:
plan = plan_repo.get(plan_id)
if plan and plan.status.value == "rendering":
plan.mark_failed()
plan_repo.update(plan)
except Exception as e:
logger.warning("标记计划失败时异常: plan_id=%s error=%s", plan_id, e, exc_info=True)
# 更新 GenerationTask 状态为 failed,前端轮询能看到失败状态
try:
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task and gen_task.status.value != "failed":
gen_task.status = "failed"
gen_task.error_message = error_msg
gen_task.completed_at = datetime.now(timezone.utc)
try:
log_stage = "render_timeout" if is_timeout else "render_failed"
gen_task.append_log(
stage=log_stage,
message=error_msg[:500],
level="ERROR",
exception_type=type(exc).__name__,
)
except Exception:
pass
gen_task_repo.update(gen_task)
logger.info(
"GenerationTask 已标记为 failed: task_id=%s plan_id=%s",
generation_task_id,
plan_id,
)
except Exception as e:
logger.warning(
"更新 GenerationTask 失败状态时异常: task_id=%s error=%s", generation_task_id, e, exc_info=True
)
raise self.retry(exc=exc, countdown=60) from exc
return {"status": "error", "message": "数据库连接失败"}
+61 -372
View File
@@ -223,7 +223,6 @@ def _load_template_segment_durations(template_id: str) -> list[float]:
return []
# DEPRECATED: 仅兼容无 source_edit_plan_id 的旧调用,后续移除
def _build_plan_and_clips_from_task(
task_id: str,
downloaded_paths: list[Path],
@@ -888,9 +887,7 @@ def _load_task_info(task_id: str) -> dict | None:
"output_height": getattr(gen_task, "output_height", OUTPUT_HEIGHT) or OUTPUT_HEIGHT,
"cover_url": getattr(gen_task, "cover_url", "") or "",
"custom_title": getattr(gen_task, "custom_title", "") or "",
"title_config": dict(getattr(gen_task, "title_config", {}) or {}),
"voice_ids": list(getattr(gen_task, "voice_ids", []) or []),
"source_edit_plan_id": getattr(gen_task, "source_edit_plan_id", "") or "",
}
finally:
session.close()
@@ -970,15 +967,14 @@ def _render_video(
bgm_config: dict | None = None,
voice_ids: list[str] | None = None,
custom_title: str = "",
title_config: dict | None = None,
) -> tuple[Path, float, list[dict] | None]:
) -> tuple[Path, float]:
"""渲染视频(含配音混音)。
使用 RenderAdapter 统一渲染入口,复用 BGM/ASR/分辨率/缩略图逻辑。
Args:
Returns:
(output_path, render_duration, cover_candidates)
(output_path, render_duration)
"""
if not downloaded_videos:
raise RuntimeError(f"素材下载结果为空: task_id={task_id}")
@@ -1003,34 +999,27 @@ def _render_video(
list(template_config.keys()),
)
# ── 用户自定义标题title_config 优先,custom_title 兜底 ─────────────
effective_title_cfg: dict | None = None
if title_config and isinstance(title_config, dict) and title_config.get("text", "").strip():
effective_title_cfg = dict(title_config)
elif custom_title:
# ── 用户自定义标题覆盖模板标题配置 ──────────────────────────────────
if custom_title:
try:
parsed = json.loads(custom_title) if isinstance(custom_title, str) else custom_title
if isinstance(parsed, dict) and parsed.get("text", "").strip():
effective_title_cfg = parsed
user_title_cfg = json.loads(custom_title) if isinstance(custom_title, str) else custom_title
if isinstance(user_title_cfg, dict) and user_title_cfg.get("text", "").strip():
# 字段名归一化: 前端 font_size/font_color → 后端 size/color
if "font_size" in user_title_cfg and "size" not in user_title_cfg:
user_title_cfg["size"] = user_title_cfg["font_size"]
if "font_color" in user_title_cfg and "color" not in user_title_cfg:
user_title_cfg["color"] = user_title_cfg["font_color"]
plan_cfg = dict(virtual_plan.config or {})
plan_cfg["title"] = user_title_cfg
virtual_plan.config = plan_cfg
logger.info(
"[task_id=%s] [渲染] 用户自定义标题已注入: text=%s",
task_id,
user_title_cfg.get("text", "")[:30],
)
except (json.JSONDecodeError, TypeError):
logger.warning("[task_id=%s] custom_title JSON解析失败: %s", task_id, custom_title[:100])
if effective_title_cfg:
# 字段名归一化: 前端 font_size/font_color → 后端 size/color
if "font_size" in effective_title_cfg and "size" not in effective_title_cfg:
effective_title_cfg["size"] = effective_title_cfg["font_size"]
if "font_color" in effective_title_cfg and "color" not in effective_title_cfg:
effective_title_cfg["color"] = effective_title_cfg["font_color"]
plan_cfg = dict(virtual_plan.config or {})
plan_cfg["title"] = effective_title_cfg
virtual_plan.config = plan_cfg
logger.info(
"[task_id=%s] [渲染] 标题配置已注入(source=%s): text=%s",
task_id,
"title_config" if title_config else "custom_title",
effective_title_cfg.get("text", "")[:30],
)
# 用户自定义 BGM 覆盖模板 BGM(用户指定优先级最高)
if bgm_config:
plan_cfg = virtual_plan.config or {}
@@ -1124,10 +1113,8 @@ def _render_video(
# 配音素材库音频已在统一渲染引擎内部通过 audio 图层混音处理
output_path = render_output_path
# RenderAdapter 在渲染完成后用本地 ffmpeg 抽取的封面候选帧(已上传 OSS)
cover_candidates = getattr(render_result, "cover_candidates", None)
return output_path, render_duration, cover_candidates
return output_path, render_duration
def _upload_and_record(
@@ -1213,135 +1200,6 @@ def _upload_and_record(
soft_time_limit=600, # 10 分钟软超时
time_limit=660, # 11 分钟硬超时
)
def _sync_task_config_to_plan(source_edit_plan_id: str, task_info: dict, db) -> str | None:
"""将 GenerationTask 的配置同步到 EditPlan.config,返回配音本地路径(如果有)。
包括:title_config、BGM、输出分辨率。配音单独处理(需下载到本地)。
"""
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
)
plan_repo = SQLAlchemyEditPlanRepository(db)
plan = plan_repo.get(source_edit_plan_id)
if plan is None:
logger.error("[task] EditPlan not found: %s", source_edit_plan_id)
return None
plan_config = dict(plan.config or {})
changed = False
# 标题配置
title_config = task_info.get("title_config") or {}
if title_config and isinstance(title_config, dict) and title_config.get("text", "").strip():
cfg = dict(title_config)
# 字段名归一化
if "font_size" in cfg and "size" not in cfg:
cfg["size"] = cfg["font_size"]
if "font_color" in cfg and "color" not in cfg:
cfg["color"] = cfg["font_color"]
plan_config["title"] = cfg
changed = True
logger.info("[task] title_config synced to plan: %s", cfg.get("text", "")[:30])
# BGM 配置
bgm_config = task_info.get("bgm_config") or {}
if bgm_config:
from packages.domain.bgm_utils import merge_bgm_config
existing_bgm = plan_config.get("bgm", {}) or {}
plan_config["bgm"] = merge_bgm_config(existing_bgm, bgm_config)
changed = True
# 输出分辨率
ow = task_info.get("output_width") or OUTPUT_WIDTH
oh = task_info.get("output_height") or OUTPUT_HEIGHT
if ow >= 100 and oh >= 100:
export_cfg = dict(plan_config.get("export", {}) or {})
export_cfg["resolution"] = f"{ow}x{oh}"
plan_config["export"] = export_cfg
changed = True
if changed:
plan.config = plan_config
plan_repo.update(plan)
logger.info("[task] plan.config synced: plan_id=%s", source_edit_plan_id)
# 配音下载
voiceover_path: str | None = None
voice_library_id = task_info.get("voice_library_id", "")
voice_ids = task_info.get("voice_ids", []) or []
effective_voice_id = voice_library_id or (voice_ids[0] if voice_ids else "")
if effective_voice_id:
import tempfile
voice_tmp = Path(tempfile.gettempdir()) / f"voice_{source_edit_plan_id}_{id(task_info)}.mp3"
try:
if _download_voice_asset(effective_voice_id, voice_tmp):
voiceover_path = str(voice_tmp)
logger.info("[task] voice downloaded: %s -> %s", effective_voice_id, voiceover_path)
except Exception:
logger.warning("[task] voice download failed: %s", effective_voice_id, exc_info=True)
return voiceover_path
def _render_from_edit_plan(
task_id: str,
source_edit_plan_id: str,
task_info: dict,
) -> tuple[Path, float, list[dict] | None]:
"""从 EditPlan 数据库记录直接渲染(不再内存重建clips)。
Returns:
(output_path, render_duration, cover_candidates)
"""
from video_processing.render_adapter import RenderAdapter
from worker_app.db import SessionLocal
db = SessionLocal()
try:
# 同步配置到 plan.config + 下载配音
voiceover_path = _sync_task_config_to_plan(source_edit_plan_id, task_info, db)
# 进度回调
def _progress_cb(progress: float, stage: str):
mapped = 40.0 + progress * 0.4
_update_task_progress(task_id, min(mapped, 80.0), stage)
adapter = RenderAdapter(db)
render_start = time.monotonic()
logger.info("[task_id=%s] [渲染] RenderAdapter.render_plan 开始 (plan_id=%s)", task_id, source_edit_plan_id)
result = adapter.render_plan(
plan_id=source_edit_plan_id,
job_id=task_id,
progress_cb=_progress_cb,
voiceover_audio_path=voiceover_path,
)
if not result.success:
raise RuntimeError(f"渲染失败: {result.error_message}")
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] RenderAdapter.render_plan 完成: 耗时=%.1fs, 时长=%.2fs",
task_id,
render_elapsed,
result.duration,
)
output_path = result.output_path
cover_candidates = getattr(result, "cover_candidates", None)
return output_path, result.duration, cover_candidates
finally:
db.close()
# 清理临时配音文件
# voiceover_path 在外部作用域,这里不直接引用
def generate_video(self, task_id: str) -> dict:
"""生成视频任务 — 使用 UnifiedRenderService 统一渲染。
@@ -1425,178 +1283,6 @@ def generate_video(self, task_id: str) -> dict:
if template_id:
_validate_template_exists(template_id)
# ── 新路径:有 source_edit_plan_id 时直接从数据库 EditPlan 渲染 ──
source_edit_plan_id = task_info.get("source_edit_plan_id", "")
if source_edit_plan_id:
logger.info(
"[task_id=%s] 使用 EditPlan 数据库路径渲染: plan_id=%s",
task_id,
source_edit_plan_id,
)
_update_task_progress(task_id, 30, "加载草稿数据")
if gen_task:
gen_task.append_log("渲染模式", "从草稿数据渲染(与预览一致)")
_flush_logs(task_id, gen_task)
output_path, render_duration, cover_candidates = _render_from_edit_plan(
task_id=task_id,
source_edit_plan_id=source_edit_plan_id,
task_info=task_info,
)
if gen_task:
gen_task.append_log("渲染", f"渲染完成, 时长={render_duration:.1f}s")
_flush_logs(task_id, gen_task)
_update_task_progress(task_id, 80, "渲染完成")
# ── 4. 上传 OSS + 查重记录 ───────────────────────────────
_update_task_progress(task_id, 85, "开始上传")
file_url, duration, file_size, video_count = _upload_and_record(
task_id=task_id,
output_path=output_path,
project_id=project_id,
batch_id=batch_id,
editing_mode=editing_mode,
user_id=user_id,
video_name=task_info.get("video_title", ""),
)
if gen_task:
gen_task.append_log(
"OSS上传",
f"上传成功, 大小={file_size}",
file_size=file_size,
file_url=file_url,
)
_flush_logs(task_id, gen_task)
_update_task_progress(task_id, 95, "上传完成")
# ── 4.5 封面帧持久化 ────────────────────────────────────────────
try:
if cover_candidates:
first = cover_candidates[0]
cover_frame_url = first.get("image_url") or first.get("url") or ""
if cover_frame_url:
_cover_session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.models import (
GenerationTaskModel,
)
_cover_model = (
_cover_session.query(GenerationTaskModel)
.filter(GenerationTaskModel.id == task_id)
.first()
)
if _cover_model:
_cover_model.cover_url = cover_frame_url
meta = dict(_cover_model.metadata or {})
meta["cover_candidates"] = cover_candidates
_cover_model.metadata = meta
_cover_session.commit()
finally:
_cover_session.close()
except Exception:
logger.warning("[task_id=%s] 封面帧持久化失败", task_id, exc_info=True)
# ── 5. 标记完成 ──────────────────────────────────────────────────
_update_task_status(task_id, "mark_completed", result_count=video_count)
# 5.1 更新标题使用次数
try:
_title_session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.adapters.sqlalchemy_impl.title_library_repository import (
SQLAlchemyTitleLibraryRepository,
)
_task_repo = SQLAlchemyGenerationTaskRepository(_title_session)
_gen_task = _task_repo.get(task_id)
if _gen_task and _gen_task.title_ids and _gen_task.created_by_user_id:
_title_repo = SQLAlchemyTitleLibraryRepository(_title_session)
for _tid in _gen_task.title_ids:
try:
_title_repo.increment_usage_count(_tid, _gen_task.created_by_user_id)
except Exception:
logger.warning(
"[task_id=%s] 更新标题使用次数失败: title_id=%s",
task_id,
_tid,
exc_info=True,
)
finally:
_title_session.close()
except Exception:
logger.warning("[task_id=%s] 更新标题使用次数异常", task_id, exc_info=True)
# 5.2 更新素材使用次数
try:
from worker_app.core.asset_usage import mark_asset_used_for_generation
_asset_session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.asset_repository import (
SQLAlchemyAssetRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
_task_repo = SQLAlchemyGenerationTaskRepository(_asset_session)
_asset_repo = SQLAlchemyAssetRepository(_asset_session)
_gen_task = _task_repo.get(task_id)
if _gen_task and _gen_task.asset_ids:
for _aid in _gen_task.asset_ids:
try:
_asset = _asset_repo.get(_aid)
if _asset:
mark_asset_used_for_generation(_asset)
_asset_repo.update(_asset)
except Exception:
logger.warning(
"[task_id=%s] 更新素材使用次数失败: asset_id=%s",
task_id,
_aid,
exc_info=True,
)
finally:
_asset_session.close()
except Exception:
logger.warning("[task_id=%s] 更新素材使用次数异常", task_id, exc_info=True)
if gen_task:
gen_task.append_log(
"任务完成",
f"视频生成完成: 时长={duration:.2f}s, 大小={file_size}",
duration=round(duration, 2),
file_size=file_size,
video_count=video_count,
)
_flush_logs(task_id, gen_task)
logger.info(
"[task_id=%s] [任务完成] duration=%.2fs file_size=%d (edit_plan path)",
task_id,
duration,
file_size,
)
return {
"status": "completed",
"task_id": task_id,
"output_path": str(output_path),
"file_size": file_size,
"duration": duration,
"mode": editing_mode.value,
}
# DEPRECATED: 以下为旧路径,仅兼容无 source_edit_plan_id 的旧调用,后续移除
with tempfile.TemporaryDirectory(prefix="xiaoxia-generation-") as temp_dir:
temp_path = Path(temp_dir)
@@ -1699,7 +1385,7 @@ def generate_video(self, task_id: str) -> dict:
else:
_resolved_resolution = task_info.get("resolution", "")
output_path, render_duration, cover_candidates = _render_video(
output_path, render_duration = _render_video(
task_id=task_id,
downloaded_videos=downloaded_videos,
voice_path=audio_path,
@@ -1713,7 +1399,6 @@ def generate_video(self, task_id: str) -> dict:
bgm_config=task_info.get("bgm_config", {}),
voice_ids=task_info.get("voice_ids", []),
custom_title=task_info.get("custom_title", ""),
title_config=task_info.get("title_config", {}),
)
if gen_task:
@@ -1745,47 +1430,51 @@ def generate_video(self, task_id: str) -> dict:
_update_task_progress(task_id, 95, "上传完成")
# ── 4.5 封面帧持久化 ────────────────────────────────────────────
# RenderAdapter 在渲染完成后已用本地 ffmpeg 从 output_path 抽帧
# (标题通过 ASS 烧录,帧天然带标题),并上传 OSS 返回 cover_candidates。
# 这里把第一帧写入 gen_task.cover_url,完整列表写入 metadata
# 封面路由(generation_cover.py)的 A/B/C/D 步骤即可直接命中。
# ── 4.5 封面抽帧 ────────────────────────────────────────────────
# 预览视频上传完成后,提取封面帧写入 gen_task.cover_url
# 这样封面路由(generation_cover.py 步骤A)可以通过 generation_task_id 直接找到
try:
if cover_candidates:
first = cover_candidates[0]
# 候选帧字段兼容:RenderAdapter 用 image_urlthumbnail_generator 用 url
cover_frame_url = first.get("image_url") or first.get("url") or ""
if cover_frame_url:
_cover_session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.models import (
GenerationTaskModel,
)
from packages.shared.mediakit_client import get_mediakit_client
_cover_model = (
_cover_session.query(GenerationTaskModel)
.filter(GenerationTaskModel.id == task_id)
.first()
)
if _cover_model:
_cover_model.cover_url = cover_frame_url
# 持久化完整候选列表到 metadata
meta = dict(_cover_model.metadata or {})
meta["cover_candidates"] = cover_candidates
_cover_model.metadata = meta
_cover_session.commit()
logger.info(
"[task_id=%s] 封面帧已持久化(ffmpeg本地抽帧): cover_url=%s candidates=%d",
task_id,
cover_frame_url[:80],
len(cover_candidates),
mk_client = get_mediakit_client()
if mk_client.is_available:
_update_task_progress(task_id, 96, "提取封面帧")
snapshots = mk_client.extract_frames(
video_url=file_url,
strategy="SpecifiedFrames",
max_frames=1,
)
if snapshots and len(snapshots) > 0:
cover_frame_url = snapshots[0].get("image_url", "")
if cover_frame_url and gen_task:
# 通过独立 session 持久化 cover_url
_cover_session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.models import (
GenerationTaskModel,
)
finally:
_cover_session.close()
_cover_model = (
_cover_session.query(GenerationTaskModel)
.filter(GenerationTaskModel.id == task_id)
.first()
)
if _cover_model:
_cover_model.cover_url = cover_frame_url
_cover_session.commit()
logger.info(
"[task_id=%s] 封面帧提取成功: %s",
task_id,
cover_frame_url[:80],
)
finally:
_cover_session.close()
else:
logger.warning("[task_id=%s] 封面帧提取返回空结果", task_id)
else:
logger.warning("[task_id=%s] 渲染未产出 cover_candidates,封面将依赖 API 兜底", task_id)
logger.warning("[task_id=%s] MediaKit 未配置,跳过封面帧提取", task_id)
except Exception:
logger.warning("[task_id=%s] 封面帧持久化失败(不影响主流程)", task_id, exc_info=True)
logger.warning("[task_id=%s] 封面帧提取失败(不影响主流程)", task_id, exc_info=True)
# ── 5. 标记完成 ──────────────────────────────────────────────────
_update_task_status(task_id, "mark_completed", result_count=video_count)
-2
View File
@@ -17,8 +17,6 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
libpq-dev \
libpq5 \
ffmpeg \
fonts-noto-cjk \
fontconfig \
&& rm -rf /var/lib/apt/lists/*
# 创建虚拟环境
-5
View File
@@ -6,13 +6,8 @@ set -e
CONCURRENCY="${WORKER_CONCURRENCY:-2}"
# ⚠️ 部署约束:此 Worker 必须且只能运行单实例(replicas=1)
# -B 标志嵌入 celery beatbeat 负责定期触发 pending 超时清理等定时任务
# 多实例部署会导致每个 Worker 独立运行 Beat,造成定时任务重复执行
# 若需横向扩展 Worker,必须将 Beat 拆分为独立服务(celery beat -A worker_app.celery_app
exec celery \
-A worker_app.celery_app \
worker \
--loglevel=info \
"-B" \
"--concurrency=${CONCURRENCY}"
@@ -42,7 +42,6 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
output_height=getattr(model, "output_height", 720) or 720,
cover_url=getattr(model, "cover_url", "") or "",
custom_title=getattr(model, "custom_title", "") or "",
title_config=dict(getattr(model, "title_config", {}) or {}),
logs=model.logs or "[]",
created_at=model.created_at,
updated_at=model.updated_at,
@@ -87,7 +86,6 @@ class SQLAlchemyGenerationTaskRepository:
output_height=task.output_height,
cover_url=task.cover_url or "",
custom_title=task.custom_title or "",
title_config=dict(task.title_config) if task.title_config else {},
logs=task.logs,
created_at=task.created_at,
updated_at=task.updated_at,
@@ -275,7 +273,6 @@ class SQLAlchemyGenerationTaskRepository:
model.output_height = task.output_height
model.cover_url = task.cover_url or ""
model.custom_title = task.custom_title or ""
model.title_config = dict(task.title_config) if task.title_config else {}
model.logs = task.logs
self.session.commit()
return task
@@ -313,42 +310,3 @@ class SQLAlchemyGenerationTaskRepository:
model.completed_at = datetime.now(timezone.utc)
self.session.commit()
return len(models)
def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int:
"""清理超时的 pending 任务(未被 Worker 拉取的任务)。
全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列,
导致新用户无法创建任务。将超时的 pending 任务标记为 failed。
Args:
timeout_minutes: 超时时间(分钟),默认 30 分钟
Returns:
清理的任务数量
"""
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 = (
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,
)
)
self.session.commit()
return count
@@ -298,7 +298,6 @@ class GenerationTaskModel(Base):
output_height = Column(Integer, nullable=False, default=720)
cover_url = Column(String(1000), nullable=False, default="")
custom_title = Column(String(500), nullable=False, default="")
title_config = Column(JSON, nullable=False, default=dict)
bgm_config = Column(JSON, nullable=False, default=dict)
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
logs = Column(Text, nullable=False, default="[]", server_default="[]")
-1
View File
@@ -69,7 +69,6 @@ class CreateGenerationTaskUseCase:
output_height=command.output_height,
cover_url=command.cover_url,
custom_title=command.custom_title,
title_config=command.title_config,
)
return self.generation_task_repository.create(task)
-3
View File
@@ -121,7 +121,6 @@ class GenerationTask:
output_height: int = 720
cover_url: str = ""
custom_title: str = ""
title_config: dict = field(default_factory=dict)
extra_meta: dict = field(default_factory=dict)
logs: str = "[]"
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@@ -154,7 +153,6 @@ class GenerationTask:
output_height: int = 720,
cover_url: str = "",
custom_title: str = "",
title_config: dict | None = None,
extra_meta: dict | None = None,
) -> "GenerationTask":
if not project_id.strip() and not template_id.strip():
@@ -186,7 +184,6 @@ class GenerationTask:
output_height=output_height,
cover_url=cover_url,
custom_title=custom_title,
title_config=dict(title_config) if title_config else {},
extra_meta=dict(extra_meta) if extra_meta else {},
)
+1 -1
View File
@@ -364,7 +364,7 @@ class MediaKitClient:
data = response.json()
status = data.get("status")
if status in ("completed", "success"):
if status == "success":
result = data.get("result", {})
snapshots = result.get("snapshots", [])
logger.info(
-171
View File
@@ -1,171 +0,0 @@
"""封面标题文字叠加(Pillow)— API / Worker 共用。
在封面帧上绘制白色标题文字 + 黑色描边/阴影,支持 CJK 字体和自动换行。
从已渲染视频抽帧时通常不需要调用(标题已烧录);
从源素材抽帧(API E2 兜底)时调用,保证封面带标题。
"""
from __future__ import annotations
import logging
from pathlib import Path
from typing import Optional
logger = logging.getLogger(__name__)
# 按优先级查找 CJK 字体(Debian/Ubuntu fonts-noto-cjk 安装路径)
_FONT_CANDIDATES = (
"/usr/share/fonts/opentype/noto/NotoSansCJK-Bold.ttc",
"/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc",
"/usr/share/fonts/truetype/noto/NotoSansCJK-Bold.ttc",
"/usr/share/fonts/truetype/noto/NotoSansCJK-Regular.ttc",
"/usr/share/fonts/truetype/wqy/wqy-zenhei.ttc",
)
def find_title_font(size: int):
"""查找可用的 CJK 字体并返回 PIL ImageFont,找不到返回 None。"""
try:
from PIL import ImageFont
except ImportError:
return None
for fp in _FONT_CANDIDATES:
if Path(fp).exists():
try:
return ImageFont.truetype(fp, size=size)
except Exception:
continue
logger.warning("未找到 CJK 字体,标题叠加将使用 PIL 默认字体(中文可能显示为方块)")
return ImageFont.load_default()
def _parse_hex_color(color: str, fallback=(255, 255, 255)) -> tuple[int, int, int]:
"将 #RRGGBB / #RGB 解析为 RGB 元组,失败返回 fallback。"
if not color or not isinstance(color, str):
return fallback
c = color.strip().lstrip("#")
try:
if len(c) == 6:
return (int(c[0:2], 16), int(c[2:4], 16), int(c[4:6], 16))
if len(c) == 3:
return (int(c[0] * 2, 16), int(c[1] * 2, 16), int(c[2] * 2, 16))
except (ValueError, IndexError):
pass
return fallback
def wrap_title_text(text: str, font, max_width: int) -> list[str]:
"""按像素宽度对中英文混合文本自动换行,支持显式 \\n。"""
lines: list[str] = []
current = ""
for ch in text:
if ch == "\n":
if current:
lines.append(current)
current = ""
continue
trial = current + ch
try:
bbox = font.getbbox(trial)
width = bbox[2] - bbox[0]
except Exception:
width = len(trial) * (font.size // 2)
if width <= max_width:
current = trial
else:
if current:
lines.append(current)
current = ch
if current:
lines.append(current)
return lines
def apply_title_to_image(
image_path: str,
title_text: str,
*,
color: str = "#ffffff",
position: str = "bottom",
font_size: Optional[int] = None,
margin_ratio: float = 0.06,
stroke_width_ratio: float = 0.04,
) -> Optional[str]:
"""在图片上绘制标题文字并覆盖保存。
Args:
image_path: 图片路径(处理结果覆盖写回)
title_text: 标题文字;为空直接返回 None 表示跳过
color: 字体颜色(#RRGGBB),默认白色
position: top / center / bottom
font_size: 字号,None 时按图片宽度自动计算
margin_ratio: 边缘留白占短边比例
stroke_width_ratio: 描边宽度占字号比例
Returns:
成功返回 image_path;标题为空或 PIL 不可用返回 None。
"""
if not title_text or not title_text.strip():
return None
try:
from PIL import Image, ImageDraw
except ImportError:
logger.warning("Pillow 未安装,跳过标题叠加: image=%s", image_path)
return None
img = Image.open(image_path).convert("RGB")
draw = ImageDraw.Draw(img)
img_w, img_h = img.size
if font_size is None:
font_size = max(28, min(72, img_w // 16))
font = find_title_font(font_size)
if font is None:
return None
text_rgb = _parse_hex_color(color)
stroke_width = max(2, int(font_size * stroke_width_ratio))
margin = int(min(img_w, img_h) * margin_ratio)
max_text_width = img_w - 2 * margin
lines = wrap_title_text(title_text.strip(), font, max_text_width)
if not lines:
return None
line_heights = []
for ln in lines:
bbox = font.getbbox(ln)
line_heights.append(bbox[3] - bbox[1])
line_height = max(line_heights) if line_heights else font_size
line_gap = int(line_height * 0.3)
total_height = len(lines) * line_height + (len(lines) - 1) * line_gap
if position == "top":
y_start = margin
elif position == "center":
y_start = (img_h - total_height) // 2
else:
y_start = img_h - total_height - margin
for i, ln in enumerate(lines):
bbox = font.getbbox(ln)
line_w = bbox[2] - bbox[0]
x = (img_w - line_w) // 2
y = y_start + i * (line_height + line_gap)
# 阴影
draw.text((x + 2, y + 2), ln, font=font, fill=(0, 0, 0))
# 文字(颜色由 color 参数控制)+ 黑色描边
draw.text(
(x, y),
ln,
font=font,
fill=text_rgb,
stroke_width=stroke_width,
stroke_fill=(0, 0, 0),
)
img.save(image_path, "JPEG", quality=92)
return image_path
-1
View File
@@ -29,4 +29,3 @@ httpx==0.27.2
# Prometheus monitoring
prometheus-client==0.21.1
Pillow==10.4.0
@@ -158,7 +158,7 @@ class TestRenderVideoVoiceInjection:
from packages.domain import EditingMode
output_path, render_duration, cover_candidates = _render_video(
output_path, render_duration = _render_video(
task_id="test_task_123",
downloaded_videos=[Path("/tmp/video1.mp4")],
voice_path=None,
-135
View File
@@ -1,135 +0,0 @@
"""Tests for PUT /templates/{id}/editor/clips batch update endpoint.
Updated for transactional replace_all_clips_transactional method.
"""
from __future__ import annotations
import os
from unittest.mock import MagicMock
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
@pytest.fixture
def mock_services():
plan_svc = MagicMock()
tpl_svc = MagicMock()
plan_svc.get_plan_or_raise.return_value = MagicMock(id="plan-1", template_id="tpl-1")
plan_svc.replace_all_clips_transactional.return_value = 2
return tpl_svc, plan_svc
class TestBatchUpdateClips:
def test_batch_update_calls_transactional_replace(self, mock_services):
"""验证批量更新调用事务性替换方法,传入正确的参数。"""
from app.api.routes.templates_editor.draft import batch_update_clips
from app.api.routes.templates_editor.schemas import (
EditorClipBatchItem,
EditorClipBatchUpdateRequest,
)
_, plan_svc = mock_services
req = EditorClipBatchUpdateRequest(
clips=[
EditorClipBatchItem(asset_id="a1", start_time=0.0, duration=3.0, order=0),
EditorClipBatchItem(asset_id="a2", start_time=3.0, duration=5.0, order=1),
]
)
result = batch_update_clips(
template_id="tpl-1",
req=req,
plan_id="plan-1",
services=mock_services,
_=MagicMock(),
)
assert result.plan_id == "plan-1"
assert result.clip_count == 2
plan_svc.replace_all_clips_transactional.assert_called_once()
call_args = plan_svc.replace_all_clips_transactional.call_args
assert call_args[0][0] == "plan-1"
clips_data = call_args[0][1]
assert len(clips_data) == 2
assert clips_data[0]["asset_id"] == "a1"
assert clips_data[0]["start_time"] == 0.0
assert clips_data[0]["duration"] == 3.0
assert clips_data[1]["asset_id"] == "a2"
def test_batch_update_empty_clips(self, mock_services):
"""空 clips 列表也能正常处理。"""
from app.api.routes.templates_editor.draft import batch_update_clips
from app.api.routes.templates_editor.schemas import EditorClipBatchUpdateRequest
_, plan_svc = mock_services
req = EditorClipBatchUpdateRequest(clips=[])
result = batch_update_clips(
template_id="tpl-1",
req=req,
plan_id="plan-1",
services=mock_services,
_=MagicMock(),
)
assert result.clip_count == 0
plan_svc.replace_all_clips_transactional.assert_called_once()
call_args = plan_svc.replace_all_clips_transactional.call_args
assert call_args[0][1] == []
def test_batch_update_passes_order_correctly(self, mock_services):
"""验证 order 字段正确传递。"""
from app.api.routes.templates_editor.draft import batch_update_clips
from app.api.routes.templates_editor.schemas import (
EditorClipBatchItem,
EditorClipBatchUpdateRequest,
)
_, plan_svc = mock_services
req = EditorClipBatchUpdateRequest(
clips=[
EditorClipBatchItem(asset_id="a1", start_time=0.0, duration=3.0, order=5),
]
)
batch_update_clips(
template_id="tpl-1",
req=req,
plan_id="plan-1",
services=mock_services,
_=MagicMock(),
)
clips_data = plan_svc.replace_all_clips_transactional.call_args[0][1]
assert clips_data[0]["order"] == 5
assert clips_data[0]["asset_id"] == "a1"
assert clips_data[0]["start_time"] == 0.0
assert clips_data[0]["duration"] == 3.0
class TestEditorClipBatchItemValidation:
"""验证 schema 校验规则。"""
def test_asset_id_empty_string_allowed(self):
"""asset_id 空字符串允许通过(占位片段场景)。"""
from app.api.routes.templates_editor.schemas import EditorClipBatchItem
item = EditorClipBatchItem(asset_id="", start_time=0.0, duration=3.0, order=0)
assert item.asset_id == ""
def test_asset_id_valid(self):
"""有效 asset_id 应通过校验。"""
from app.api.routes.templates_editor.schemas import EditorClipBatchItem
item = EditorClipBatchItem(asset_id="abc123", start_time=0.0, duration=3.0, order=0)
assert item.asset_id == "abc123"
def test_order_none_by_default(self):
"""order 默认为 None,表示按数组顺序。"""
from app.api.routes.templates_editor.schemas import EditorClipBatchItem
item = EditorClipBatchItem(asset_id="a1", start_time=0.0, duration=3.0)
assert item.order is None
+132
View File
@@ -0,0 +1,132 @@
"""Tests for cover_url backfill to GenerationTask.
Verifies _finalize_render_success correctly writes cover_url
from cover_candidates to gen_task.cover_url.
"""
from __future__ import annotations
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
# Add worker app to sys.path
_WORKER_ROOT = Path(__file__).resolve().parents[2] / "apps" / "worker"
if str(_WORKER_ROOT) not in sys.path:
sys.path.insert(0, str(_WORKER_ROOT))
class FakeGenTask:
"""Simple stand-in for GenerationTask that tracks attribute assignment."""
def __init__(self):
object.__setattr__(self, "_assigned", {})
self.id = "task-1"
self.status = MagicMock()
self.status.value = "running"
def __setattr__(self, name, value):
if not name.startswith("_"):
self._assigned[name] = value
object.__setattr__(self, name, value)
def append_log(self, **kwargs):
pass
def _make_plan():
plan = MagicMock()
plan.project_id = "proj-1"
plan.created_by_user_id = "user-1"
plan.config = {"batch_id": "batch-1", "mode": "edit_plan", "title": {"text": "test"}}
plan.mark_completed = MagicMock()
return plan
def _call_finalize(cover_candidates=None, gen_task=None, plan=None):
from worker_app.tasks.edit_plan_generation import _finalize_render_success
plan = plan or _make_plan()
gen_task = gen_task or FakeGenTask()
plan_repo = MagicMock()
clip_repo = MagicMock()
gen_task_repo = MagicMock()
gen_task_repo.get.return_value = gen_task
db = MagicMock()
with patch("worker_app.tasks.edit_plan_generation.create_video_record_and_dedup"):
result = _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id="plan-1",
output_url="https://oss.example.com/output.mp4",
storage_key="rendered/plan-1/task-1.mp4",
duration=10.0,
file_size=1024,
width=1280,
height=720,
rendered_clip_ids=["clip-1"],
failed_clip_ids=[],
generation_task_id="task-1",
output_path=Path("/tmp/output.mp4"),
engine="unified",
thumbnail_url="",
cover_candidates=cover_candidates,
)
return result, gen_task, gen_task_repo
class TestFinalizeCoverUrl:
def test_cover_url_set_from_image_url(self):
"""cover_candidates with image_url should set gen_task.cover_url"""
candidates = [
{"image_url": "https://oss.example.com/cover1.jpg", "frame_time": 1.5},
{"image_url": "https://oss.example.com/cover2.jpg", "frame_time": 3.0},
]
_, gen_task, gen_task_repo = _call_finalize(cover_candidates=candidates)
assert gen_task.cover_url == "https://oss.example.com/cover1.jpg"
gen_task_repo.update.assert_called()
def test_cover_url_fallback_to_url_key(self):
"""Should fallback to 'url' key when 'image_url' is absent"""
candidates = [{"url": "https://oss.example.com/cover_url_key.jpg"}]
_, gen_task, _ = _call_finalize(cover_candidates=candidates)
assert gen_task.cover_url == "https://oss.example.com/cover_url_key.jpg"
def test_cover_url_not_set_when_empty_list(self):
"""Empty cover_candidates should not set cover_url"""
_, gen_task, _ = _call_finalize(cover_candidates=[])
assert "cover_url" not in gen_task._assigned
def test_cover_url_not_set_when_none(self):
"""None cover_candidates should not set cover_url"""
_, gen_task, _ = _call_finalize(cover_candidates=None)
assert "cover_url" not in gen_task._assigned
def test_cover_url_not_set_when_url_empty(self):
"""Empty URL strings in candidates should not set cover_url"""
candidates = [{"image_url": "", "url": ""}]
_, gen_task, _ = _call_finalize(cover_candidates=candidates)
assert "cover_url" not in gen_task._assigned
def test_no_generation_task_no_crash(self):
"""Should not crash when gen_task is None"""
candidates = [{"image_url": "https://oss.example.com/cover.jpg"}]
gen_task_repo = MagicMock()
gen_task_repo.get.return_value = None
result, _, _ = _call_finalize(cover_candidates=candidates)
assert result["status"] == "completed"
def test_image_url_priority_over_url(self):
"""image_url should take priority over url key"""
candidates = [{"image_url": "https://a.jpg", "url": "https://b.jpg"}]
_, gen_task, _ = _call_finalize(cover_candidates=candidates)
assert gen_task.cover_url == "https://a.jpg"
+333
View File
@@ -0,0 +1,333 @@
"""P0-2: Celery 任务 render_edit_plan 失败时更新 GenerationTask 状态。
验证:
- 异常发生时 GenerationTask 状态更新为 failed
- error_message 记录了异常类型和描述
- completed_at 被设置
- 即使 generation_task_id 为空也不崩溃
- 即使更新 GenerationTask 本身失败也不影响 retry
"""
from __future__ import annotations
import os
import sys
from dataclasses import dataclass, field
from types import ModuleType
from typing import Any, Optional
from unittest.mock import MagicMock, patch
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
import pytest
# ── Mock worker 模块以避免数据库连接 ──────────────────────────────────────────
# worker_app.db 在 import 时会尝试连接数据库,必须在导入 task 模块前 mock
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", ".."))
# 预注册 mock 模块,阻止真实数据库初始化
_mock_db_mod = ModuleType("worker_app.db")
_mock_db_mod.SessionLocal = MagicMock()
sys.modules.setdefault("worker_app.db", _mock_db_mod)
_mock_celery_mod = ModuleType("worker_app.celery_app")
_mock_celery_app = MagicMock()
# 让 @celery_app.task(...) 装饰器透传原始函数,否则函数变成 MagicMock
_mock_celery_app.task = lambda **kwargs: lambda fn: fn
_mock_celery_mod.celery_app = _mock_celery_app
sys.modules.setdefault("worker_app.celery_app", _mock_celery_mod)
# ── Stub domain objects ───────────────────────────────────────────────────────
@dataclass
class _StubStatus:
value: str
def __eq__(self, other):
if isinstance(other, str):
return self.value == other
if isinstance(other, _StubStatus):
return self.value == other.value
return NotImplemented
@dataclass
class StubEditPlan:
id: str = "plan-001"
template_id: str = "tmpl-001"
status: Any = None
config: dict = field(default_factory=dict)
project_id: str = ""
created_by_user_id: str = "user-001"
def mark_failed(self):
self.status = _StubStatus("failed")
def mark_completed(self):
self.status = _StubStatus("completed")
@dataclass
class StubGenerationTask:
id: str = "gen-task-001"
status: Any = field(default_factory=lambda: _StubStatus("pending"))
error_message: str = ""
progress: float = 0.0
result_count: int = 0
started_at: Any = None
completed_at: Any = None
project_id: str = ""
created_by_user_id: str = "user-001"
@dataclass
class StubClip:
id: str = "clip-001"
plan_id: str = "plan-001"
asset_id: str = "assets/video.mp4"
order: int = 1
status: Any = field(default_factory=lambda: _StubStatus("ready"))
transition_effect: str = ""
text_content: str = ""
clip_type: str = "MAIN"
duration: float = 0.0
def mark_failed(self):
self.status = _StubStatus("failed")
def mark_rendered(self):
self.status = _StubStatus("rendered")
# ── Stub repositories ─────────────────────────────────────────────────────────
class StubPlanRepo:
def __init__(self, plan: StubEditPlan):
self._plan = plan
def get(self, plan_id: str) -> Optional[StubEditPlan]:
if plan_id == self._plan.id:
return self._plan
return None
def update(self, plan: StubEditPlan) -> StubEditPlan:
self._plan = plan
return plan
class StubClipRepo:
def __init__(self, clips: list[StubClip] | None = None):
self._clips = clips or []
def list_by_plan(self, plan_id: str, skip: int = 0, limit: int = 10000) -> list[StubClip]:
return [c for c in self._clips if c.plan_id == plan_id]
def get(self, clip_id: str) -> Optional[StubClip]:
for c in self._clips:
if c.id == clip_id:
return c
return None
def update(self, clip: StubClip) -> StubClip:
return clip
class StubGenTaskRepo:
def __init__(self, task: StubGenerationTask | None = None):
self._store: dict[str, StubGenerationTask] = {}
if task:
self._store[task.id] = task
def get(self, task_id: str) -> Optional[StubGenerationTask]:
return self._store.get(task_id)
def update(self, task: StubGenerationTask) -> StubGenerationTask:
self._store[task.id] = task
return task
# ── Import task module (after mocks are in place) ─────────────────────────────
from worker_app.tasks.edit_plan_generation import render_edit_plan
# ── Tests ─────────────────────────────────────────────────────────────────────
class TestRenderEditPlanFailureUpdatesGenTask:
"""P0-2: render_edit_plan 异常时更新 GenerationTask 状态为 failed"""
def _make_bound_task(self):
"""构建绑定的 Celery task mock"""
task = MagicMock()
task.retry = MagicMock(side_effect=RuntimeError("retry called"))
return task
def test_exception_marks_gen_task_failed(self):
"""异常时 GenerationTask.status 被设为 failed"""
plan = StubEditPlan(status=_StubStatus("rendering"))
plan.config["generation_task_id"] = "gen-task-001"
gen_task = StubGenerationTask(id="gen-task-001", status=_StubStatus("running"))
plan_repo = StubPlanRepo(plan)
gen_task_repo = StubGenTaskRepo(gen_task)
# 让 clip_repo 抛异常以触发 except 路径
clip_repo_bad = MagicMock()
clip_repo_bad.list_by_plan.side_effect = RuntimeError("OSS 连接失败")
def fake_get_repos():
yield plan_repo, clip_repo_bad, gen_task_repo, MagicMock()
bound_task = self._make_bound_task()
with patch(
"worker_app.tasks.edit_plan_generation._get_repos",
side_effect=fake_get_repos,
):
with pytest.raises(RuntimeError, match="retry called"):
render_edit_plan(bound_task, "plan-001")
# 核心断言:GenerationTask 状态为 failed(生产代码赋值为字符串)
assert gen_task.status == "failed"
def test_exception_records_error_message(self):
"""异常时 error_message 包含异常类型和描述"""
plan = StubEditPlan(status=_StubStatus("rendering"))
plan.config["generation_task_id"] = "gen-task-001"
gen_task = StubGenerationTask(id="gen-task-001", status=_StubStatus("running"))
plan_repo = StubPlanRepo(plan)
clip_repo = MagicMock()
clip_repo.list_by_plan.side_effect = RuntimeError("DB 查询超时")
gen_task_repo = StubGenTaskRepo(gen_task)
def fake_get_repos():
yield plan_repo, clip_repo, gen_task_repo, MagicMock()
bound_task = self._make_bound_task()
with patch(
"worker_app.tasks.edit_plan_generation._get_repos",
side_effect=fake_get_repos,
):
with pytest.raises(RuntimeError, match="retry called"):
render_edit_plan(bound_task, "plan-001")
assert gen_task.status == "failed"
assert "DB 查询超时" in gen_task.error_message
assert "RuntimeError" in gen_task.error_message
def test_exception_sets_completed_at(self):
"""异常时 completed_at 被设置"""
plan = StubEditPlan(status=_StubStatus("rendering"))
plan.config["generation_task_id"] = "gen-task-001"
gen_task = StubGenerationTask(id="gen-task-001", status=_StubStatus("running"))
plan_repo = StubPlanRepo(plan)
clip_repo = MagicMock()
clip_repo.list_by_plan.side_effect = RuntimeError("boom")
gen_task_repo = StubGenTaskRepo(gen_task)
def fake_get_repos():
yield plan_repo, clip_repo, gen_task_repo, MagicMock()
bound_task = self._make_bound_task()
with patch(
"worker_app.tasks.edit_plan_generation._get_repos",
side_effect=fake_get_repos,
):
with pytest.raises(RuntimeError):
render_edit_plan(bound_task, "plan-001")
assert gen_task.completed_at is not None
def test_no_generation_task_id_does_not_crash(self):
"""generation_task_id 为空时,异常处理不崩溃"""
plan = StubEditPlan(status=_StubStatus("rendering"))
plan.config = {} # 不设置 generation_task_id
plan_repo = StubPlanRepo(plan)
clip_repo = MagicMock()
clip_repo.list_by_plan.side_effect = RuntimeError("boom")
gen_task_repo = StubGenTaskRepo() # 空 repo
def fake_get_repos():
yield plan_repo, clip_repo, gen_task_repo, MagicMock()
bound_task = self._make_bound_task()
with patch(
"worker_app.tasks.edit_plan_generation._get_repos",
side_effect=fake_get_repos,
):
with pytest.raises(RuntimeError, match="retry called"):
render_edit_plan(bound_task, "plan-001")
# 计划仍被标记为 failed
assert plan.status.value == "failed"
def test_gen_task_update_failure_does_not_block_retry(self):
"""更新 GenerationTask 失败时,不影响 retry 流程"""
plan = StubEditPlan(status=_StubStatus("rendering"))
plan.config["generation_task_id"] = "gen-task-001"
gen_task = StubGenerationTask(id="gen-task-001", status=_StubStatus("running"))
plan_repo = StubPlanRepo(plan)
clip_repo = MagicMock()
clip_repo.list_by_plan.side_effect = RuntimeError("原始错误")
# gen_task_repo.update 也抛异常
gen_task_repo = MagicMock()
gen_task_repo.get.return_value = gen_task
gen_task_repo.update.side_effect = RuntimeError("DB 写入失败")
def fake_get_repos():
yield plan_repo, clip_repo, gen_task_repo, MagicMock()
bound_task = self._make_bound_task()
with patch(
"worker_app.tasks.edit_plan_generation._get_repos",
side_effect=fake_get_repos,
):
with pytest.raises(RuntimeError, match="retry called"):
render_edit_plan(bound_task, "plan-001")
# retry 被调用说明流程正确
bound_task.retry.assert_called_once()
def test_already_failed_gen_task_not_overwritten(self):
"""已经 failed 的 GenerationTask 不会被重复更新"""
plan = StubEditPlan(status=_StubStatus("rendering"))
plan.config["generation_task_id"] = "gen-task-001"
gen_task = StubGenerationTask(
id="gen-task-001",
status=_StubStatus("failed"), # 已经是 failed
error_message="之前的错误",
)
plan_repo = StubPlanRepo(plan)
clip_repo = MagicMock()
clip_repo.list_by_plan.side_effect = RuntimeError("新错误")
gen_task_repo = StubGenTaskRepo(gen_task)
def fake_get_repos():
yield plan_repo, clip_repo, gen_task_repo, MagicMock()
bound_task = self._make_bound_task()
with patch(
"worker_app.tasks.edit_plan_generation._get_repos",
side_effect=fake_get_repos,
):
with pytest.raises(RuntimeError, match="retry called"):
render_edit_plan(bound_task, "plan-001")
# error_message 应保持原值,不被覆盖
assert gen_task.error_message == "之前的错误"
+1 -514
View File
@@ -182,7 +182,6 @@ class TestUnifiedCoverPipelineEndpoint:
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
patch("packages.shared.mediakit_client.get_mediakit_client") as mock_mk_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = mock_task
@@ -195,13 +194,6 @@ class TestUnifiedCoverPipelineEndpoint:
mock_storage_svc.get_url.return_value = "https://oss.example.com/rendered/plan-2/video.mp4"
mock_storage_getter.return_value = mock_storage_svc
# MediaKit 抽帧也返回 None,模拟最终失败
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = None
mock_mk_getter.return_value = mock_mk
# body 不传 asset_ids,步骤 E2 不会进入
from app.api.routes.generation_cover import generate_cover
with pytest.raises(HTTPException) as exc_info:
@@ -215,6 +207,7 @@ class TestUnifiedCoverPipelineEndpoint:
)
assert exc_info.value.status_code == 400
assert "封面尚未生成" in exc_info.value.detail
def test_cover_url_found_via_source_edit_plan(self):
"""步骤B:通过 source_edit_plan_id 找到预览任务的 cover_url。"""
@@ -325,157 +318,6 @@ class TestUnifiedCoverPipelineEndpoint:
template_id="template-y",
)
def test_cover_url_found_via_cover_candidates_image_url(self):
"""步骤Dplan.config.cover_candidates 有 image_url 时,直接使用第一个候选封面。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
# 步骤A/B/C 都找不到,进入步骤D
mock_plan.config = {
"rendered_storage_key": "rendered/plan-z/video.mp4", # 必须有预览视频才能通过前置检查
"cover_candidates": [
{"image_url": "https://oss.example.com/candidates/cover-1.jpg", "score": 0.95},
{"image_url": "https://oss.example.com/candidates/cover-2.jpg", "score": 0.80},
],
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
body = GenerateCoverRequest(cover_type="ai_frame")
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/candidates/cover-1.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="template-z",
plan_id="plan-z",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
# 步骤D从 cover_candidates 第一个元素的 image_url 提取封面
assert result.cover["image_url"] == "https://oss.example.com/candidates/cover-1.jpg"
# 验证 plan.config 被更新(至少调用一次:rendered_storage_key + cover
assert mock_plan_svc.update_plan_config.call_count >= 1
def test_cover_url_found_via_cover_candidates_url_key(self):
"""步骤Dcover_candidates 用 url 键(非 image_url)时,也能正确提取。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {
"rendered_storage_key": "rendered/plan-w/video.mp4", # 必须有预览视频才能通过前置检查
"cover_candidates": [
{"url": "https://oss.example.com/candidates/alt-cover.jpg"},
],
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
body = GenerateCoverRequest(cover_type="ai_frame")
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/candidates/alt-cover.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="template-w",
plan_id="plan-w",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
# 步骤D fallback 到 url 键
assert result.cover["image_url"] == "https://oss.example.com/candidates/alt-cover.jpg"
def test_cover_candidates_skips_non_dict_first_element(self):
"""步骤Dcover_candidates 第一个元素不是 dict 时,安全跳过不崩溃。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
from fastapi import HTTPException
mock_plan = MagicMock()
mock_plan.config = {
"rendered_storage_key": "rendered/plan-skip/video.mp4", # 必须有预览视频才能通过前置检查
"cover_candidates": ["not-a-dict", 42, None],
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
body = GenerateCoverRequest(cover_type="ai_frame")
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
# storage fallback 也找不到封面
mock_storage_svc = MagicMock()
mock_storage_svc.get_url.return_value = ""
mock_storage_getter.return_value = mock_storage_svc
from app.api.routes.generation_cover import generate_cover
# 所有步骤都失败,应返回 400
with pytest.raises(HTTPException) as exc_info:
generate_cover(
body=body,
template_id="template-skip",
plan_id="plan-skip",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert exc_info.value.status_code == 400
class TestSourceEditPlanFallback:
"""测试步骤 2.5:通过 source_edit_plan_id 查找预览视频兜底逻辑。"""
@@ -845,358 +687,3 @@ class TestUploadCoverType:
assert result.cover["image_url"] == "https://oss.example.com/uploaded/my-cover.png"
# 验证没有调用任何预览视频查找逻辑
# (normalize_plan_config 是唯一被调用的外部函数)
def test_cover_extracted_from_source_asset_when_no_preview(self):
"""步骤E2:无后端渲染产物时,直接从用户选择的视频素材抽帧。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {} # 无 rendered_storage_key,无 generation_task_id
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
# 模拟视频素材
mock_asset = MagicMock()
mock_asset.file_type = "video"
mock_asset.storage_key = "uploads/source-clip.mp4"
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = mock_asset
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mediakit.internal/frame-abc.jpg"}]
mock_storage = MagicMock()
mock_storage.get_url.return_value = "https://oss.example.com/uploads/source-clip.mp4"
body = GenerateCoverRequest(
cover_type="ai_frame",
asset_ids=["asset-video-1"],
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"packages.adapters.sqlalchemy_impl.asset_repository.SQLAlchemyAssetRepository",
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service", return_value=mock_storage),
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/final.jpg",
),
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/final.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="tpl-source",
plan_id="plan-source",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert result.cover["image_url"] == "https://oss.example.com/covers/final.jpg"
mock_mk.extract_frames.assert_called_once()
# 确保用的是源素材 URL
call_kwargs = mock_mk.extract_frames.call_args.kwargs
assert "source-clip.mp4" in call_kwargs["video_url"]
def test_e2_passes_plan_title_to_persist_for_overlay(self):
"""步骤E2plan.config.title.text 存在时,作为 title_text 传给 _persist_cover_frame 叠加标题。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {"title": {"enabled": True, "text": "我的视频标题"}}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
mock_asset = MagicMock()
mock_asset.file_type = "video"
mock_asset.storage_key = "uploads/src.mp4"
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = mock_asset
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mk/frame.jpg"}]
mock_storage = MagicMock()
mock_storage.get_url.return_value = "https://oss.example.com/uploads/src.mp4"
body = GenerateCoverRequest(cover_type="ai_frame", asset_ids=["a1"])
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"packages.adapters.sqlalchemy_impl.asset_repository.SQLAlchemyAssetRepository",
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service", return_value=mock_storage),
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/final.jpg",
) as mock_persist,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/final.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="tpl",
plan_id="plan-title",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert result.cover["image_url"] == "https://oss.example.com/covers/final.jpg"
# 标题文字必须透传给持久化函数(用于源素材帧叠加标题)
assert mock_persist.call_args.kwargs.get("title_text") == "我的视频标题"
def test_e2_passes_full_title_style_to_persist(self):
"""步骤E2plan.config.title 包含完整样式时,color/position/font_size 都传给 _persist_cover_frame。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {
"title": {
"enabled": True,
"text": "样式标题",
"color": "#00ff00",
"position": "top",
"font_size": 42,
}
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
mock_asset = MagicMock()
mock_asset.file_type = "video"
mock_asset.storage_key = "uploads/src.mp4"
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = mock_asset
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mk/frame.jpg"}]
mock_storage = MagicMock()
mock_storage.get_url.return_value = "https://oss.example.com/uploads/src.mp4"
body = GenerateCoverRequest(cover_type="ai_frame", asset_ids=["a1"])
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"packages.adapters.sqlalchemy_impl.asset_repository.SQLAlchemyAssetRepository",
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service", return_value=mock_storage),
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/styled.jpg",
) as mock_persist,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/styled.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="tpl",
plan_id="plan-style",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert result.cover["image_url"] == "https://oss.example.com/covers/styled.jpg"
kwargs = mock_persist.call_args.kwargs
assert kwargs["title_text"] == "样式标题"
assert kwargs["title_color"] == "#00ff00"
assert kwargs["title_position"] == "top"
assert kwargs["title_font_size"] == 42
def test_e2_title_style_fallback_font_color(self):
"""步骤E2:前端传 font_color 时能正确兼容读取。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {
"title": {
"enabled": True,
"text": "兼容标题",
"font_color": "#123456",
"position": "center",
}
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
mock_asset = MagicMock()
mock_asset.file_type = "video"
mock_asset.storage_key = "uploads/src.mp4"
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = mock_asset
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mk/frame.jpg"}]
mock_storage = MagicMock()
mock_storage.get_url.return_value = "https://oss.example.com/uploads/src.mp4"
body = GenerateCoverRequest(cover_type="ai_frame", asset_ids=["a1"])
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"packages.adapters.sqlalchemy_impl.asset_repository.SQLAlchemyAssetRepository",
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service", return_value=mock_storage),
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/compat.jpg",
) as mock_persist,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/compat.jpg"}
}
from app.api.routes.generation_cover import generate_cover
generate_cover(
body=body,
template_id="tpl",
plan_id="plan-compat",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
kwargs = mock_persist.call_args.kwargs
assert kwargs["title_color"] == "#123456"
assert kwargs["title_position"] == "center"
assert kwargs["title_font_size"] is None
def test_step_e_skips_non_video_assets(self):
"""步骤E2:asset_ids 里只有图片素材时,不调用 MediaKit 并返回 400。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
from fastapi import HTTPException
mock_plan = MagicMock()
mock_plan.config = {}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
mock_image_asset = MagicMock()
mock_image_asset.file_type = "image"
mock_image_asset.storage_key = "uploads/photo.png"
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = mock_image_asset
mock_mk = MagicMock()
mock_mk.is_available = True
body = GenerateCoverRequest(
cover_type="ai_frame",
asset_ids=["asset-img-1"],
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"packages.adapters.sqlalchemy_impl.asset_repository.SQLAlchemyAssetRepository",
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_storage_getter.return_value = MagicMock()
from app.api.routes.generation_cover import generate_cover
with pytest.raises(HTTPException) as exc_info:
generate_cover(
body=body,
template_id="tpl-img",
plan_id="plan-img",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert exc_info.value.status_code == 400
mock_mk.extract_frames.assert_not_called()
@@ -1,151 +0,0 @@
"""GenerationTaskRepository - cleanup_stale_pending 超时 pending 清理单元测试。"""
import sys
from datetime import datetime, timedelta, timezone
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
from sqlalchemy import create_engine, text
from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.adapters.sqlalchemy_impl.models import Base
from packages.domain import GenerationTask, GenerationTaskStatus
def _repository():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
session = sessionmaker(bind=engine)()
return SQLAlchemyGenerationTaskRepository(session), session, engine
def _make_task(**kwargs) -> GenerationTask:
defaults = dict(
project_id="proj-1",
asset_library_id="lib-1",
created_by_user_id="user-1",
)
defaults.update(kwargs)
return GenerationTask.create(**defaults)
# ---------------------------------------------------------------------------
# cleanup_stale_pending 基本测试
# ---------------------------------------------------------------------------
def test_cleanup_stale_pending_no_tasks_returns_zero():
"""没有任务时返回 0。"""
repo, _, _ = _repository()
count = repo.cleanup_stale_pending(timeout_minutes=30)
assert count == 0
def test_cleanup_stale_pending_recent_pending_not_cleaned():
"""30 分钟内的 pending 任务不被清理。"""
repo, _, _ = _repository()
task = _make_task()
repo.create(task)
# 刚创建的 pending 任务不应被清理
count = repo.cleanup_stale_pending(timeout_minutes=30)
assert count == 0
assert repo.get(task.id).status == GenerationTaskStatus.PENDING
def test_cleanup_stale_pending_old_pending_marked_failed():
"""超过 30 分钟的 pending 任务被标记为 failed。"""
repo, _, engine = _repository()
task = _make_task()
repo.create(task)
# 手动把 created_at 改到 1 小时前
with engine.connect() as conn:
conn.execute(
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
{"ts": datetime.now(timezone.utc) - timedelta(hours=1), "id": task.id},
)
conn.commit()
count = repo.cleanup_stale_pending(timeout_minutes=30)
assert count == 1
saved = repo.get(task.id)
assert saved.status == GenerationTaskStatus.FAILED
assert saved.error_message == "pending timeout: auto cleanup"
assert saved.error_info.get("error_type") == "PendingTimeout"
assert "30" in saved.error_info["message"]
assert "failed_at" in saved.error_info
assert saved.completed_at is not None
def test_cleanup_stale_pending_running_not_touched():
"""running 任务不受影响,只清理 pending。"""
repo, _, engine = _repository()
task = _make_task()
repo.create(task)
task.mark_processing()
repo.update(task)
# 回写 created_at 到 1 小时前
with engine.connect() as conn:
conn.execute(
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
{"ts": datetime.now(timezone.utc) - timedelta(hours=1), "id": task.id},
)
conn.commit()
count = repo.cleanup_stale_pending(timeout_minutes=30)
assert count == 0
assert repo.get(task.id).status == GenerationTaskStatus.RUNNING
def test_cleanup_stale_pending_custom_timeout():
"""自定义超时时间生效。"""
repo, _, engine = _repository()
task = _make_task()
repo.create(task)
# 回写 created_at 到 20 分钟前
with engine.connect() as conn:
conn.execute(
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
{"ts": datetime.now(timezone.utc) - timedelta(minutes=20), "id": task.id},
)
conn.commit()
# 30 分钟超时:不清理
count_30 = repo.cleanup_stale_pending(timeout_minutes=30)
assert count_30 == 0
# 15 分钟超时:清理
count_15 = repo.cleanup_stale_pending(timeout_minutes=15)
assert count_15 == 1
assert repo.get(task.id).status == GenerationTaskStatus.FAILED
def test_cleanup_stale_pending_multiple():
"""批量清理多个超时的 pending 任务。"""
repo, _, engine = _repository()
tasks = []
for i in range(5):
t = _make_task(project_id=f"proj-{i}")
repo.create(t)
tasks.append(t)
# 全部回写 created_at 到 2 小时前
with engine.connect() as conn:
for t in tasks:
conn.execute(
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
{"ts": datetime.now(timezone.utc) - timedelta(hours=2), "id": t.id},
)
conn.commit()
count = repo.cleanup_stale_pending(timeout_minutes=30)
assert count == 5
for t in tasks:
assert repo.get(t.id).status == GenerationTaskStatus.FAILED
+5 -24
View File
@@ -13,18 +13,6 @@ os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
def _patch_session_local(mock_session):
"""Patch worker_app.db.SessionLocal robustly even when other tests
have pre-registered a MagicMock for worker_app.db in sys.modules.
Uses patch.dict to inject a clean module so that
'from worker_app.db import SessionLocal' resolves correctly."""
from types import ModuleType
_fresh_db = ModuleType("worker_app.db")
_fresh_db.SessionLocal = lambda *a, **kw: mock_session
return patch.dict(sys.modules, {"worker_app.db": _fresh_db})
class TestLoadTemplateSegmentDurations:
"""_load_template_segment_durations 单元测试 (covers lines 198-226)."""
@@ -52,7 +40,8 @@ class TestLoadTemplateSegmentDurations:
mock_session = MagicMock()
mock_session.query.return_value = mock_query
with _patch_session_local(mock_session):
# Patch at the source module since it's imported inside the function
with patch("worker_app.db.SessionLocal", return_value=mock_session):
result = _load_template_segment_durations("tpl_123")
assert result == [5.0, 8.0, 3.0]
@@ -74,24 +63,16 @@ class TestLoadTemplateSegmentDurations:
mock_session = MagicMock()
mock_session.query.return_value = mock_query
with _patch_session_local(mock_session):
with patch("worker_app.db.SessionLocal", return_value=mock_session):
result = _load_template_segment_durations("tpl_456")
assert result == [5.0]
def test_db_error_returns_empty(self):
"""数据库异常返回空列表,不抛出。"""
from types import ModuleType
from worker_app.tasks.generation import _load_template_segment_durations
_err_db = ModuleType("worker_app.db")
def _raise(*a, **kw):
raise Exception("DB down")
_err_db.SessionLocal = _raise
with patch.dict(sys.modules, {"worker_app.db": _err_db}):
with patch("worker_app.db.SessionLocal", side_effect=Exception("DB down")):
result = _load_template_segment_durations("tpl_789")
assert result == []
@@ -105,7 +86,7 @@ class TestLoadTemplateSegmentDurations:
mock_session = MagicMock()
mock_session.query.return_value = mock_query
with _patch_session_local(mock_session):
with patch("worker_app.db.SessionLocal", return_value=mock_session):
result = _load_template_segment_durations("tpl_empty")
assert result == []
+92 -78
View File
@@ -1,6 +1,7 @@
"""Unit tests for apps/api/app/api/routes/health.py
覆盖 _check_database() _check_migrations() psycopg3 连接逻辑
确保增量覆盖率 60%目标覆盖 lines 52, 127
"""
from unittest.mock import AsyncMock, MagicMock, patch
@@ -8,70 +9,63 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
def _make_cursor(fetchone_result=None):
"""Create a mock cursor with context manager support."""
cur = MagicMock()
cur.__enter__ = MagicMock(return_value=cur)
cur.__exit__ = MagicMock(return_value=False)
if fetchone_result is not None:
cur.fetchone.return_value = fetchone_result
return cur
def _make_conn(cursor_result=None):
conn = MagicMock()
conn.cursor.return_value = cursor_result or _make_cursor()
return conn
@pytest.mark.asyncio
class TestCheckDatabase:
"""Tests for _check_database() health check function."""
async def test_check_database_success(self):
mock_cur = _make_cursor(fetchone_result=(1,))
mock_conn = _make_conn(mock_cur)
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_database_success(self, mock_connect, mock_settings):
"""PostgreSQL 连接成功时返回 healthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
from apps.api.app.api.routes import health
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_cursor.__enter__ = MagicMock(return_value=mock_cursor)
mock_cursor.__exit__ = MagicMock(return_value=False)
mock_cursor.fetchone.return_value = (1,)
mock_conn.cursor.return_value = mock_cursor
mock_connect.return_value = mock_conn
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.return_value = mock_conn
result = await health._check_database()
from apps.api.app.api.routes.health import _check_database
result = await _check_database()
assert result["status"] == "healthy"
assert result["type"] == "postgresql"
assert result["message"] == "Database connection successful"
mock_psycopg.connect.assert_called_once_with(
mock_connect.assert_called_once_with(
"postgresql+psycopg://test:test@localhost/test", connect_timeout=3
)
mock_cur.execute.assert_called_once_with("SELECT 1")
mock_cursor.execute.assert_called_once_with("SELECT 1")
mock_conn.close.assert_called_once()
async def test_check_database_connection_failure(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_database_connection_failure(self, mock_connect, mock_settings):
"""PostgreSQL 连接失败时返回 unhealthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
mock_connect.side_effect = Exception("connection refused")
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_database
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.side_effect = Exception("connection refused")
result = await health._check_database()
result = await _check_database()
assert result["status"] == "unhealthy"
assert result["type"] == "postgresql"
assert "connection refused" in result["message"]
async def test_check_database_in_memory(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
async def test_check_database_in_memory(self, mock_settings):
"""使用内存数据库时跳过 PostgreSQL 检查。"""
mock_settings.USE_IN_MEMORY_DB = True
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_database
with patch.object(health, "settings", mock_settings):
result = await health._check_database()
result = await _check_database()
assert result["status"] == "healthy"
assert result["type"] == "in_memory"
@@ -79,66 +73,80 @@ class TestCheckDatabase:
@pytest.mark.asyncio
class TestCheckMigrations:
"""Tests for _check_migrations() health check function."""
async def test_check_migrations_success(self):
mock_cur = _make_cursor(fetchone_result=(5,))
mock_conn = _make_conn(mock_cur)
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_migrations_success(self, mock_connect, mock_settings):
"""所有迁移表存在时返回 healthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
from apps.api.app.api.routes import health
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_cursor.__enter__ = MagicMock(return_value=mock_cursor)
mock_cursor.__exit__ = MagicMock(return_value=False)
mock_cursor.fetchone.return_value = (5,) # 5 tables found
mock_conn.cursor.return_value = mock_cursor
mock_connect.return_value = mock_conn
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.return_value = mock_conn
result = await health._check_migrations()
from apps.api.app.api.routes.health import _check_migrations
result = await _check_migrations()
assert result["status"] == "healthy"
assert result["message"] == "Database migrations applied"
mock_psycopg.connect.assert_called_once_with(
mock_connect.assert_called_once_with(
"postgresql+psycopg://test:test@localhost/test", connect_timeout=3
)
mock_conn.close.assert_called_once()
async def test_check_migrations_missing_tables(self):
mock_cur = _make_cursor(fetchone_result=(2,))
mock_conn = _make_conn(mock_cur)
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_migrations_missing_tables(self, mock_connect, mock_settings):
"""迁移表不完整时返回 unhealthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
from apps.api.app.api.routes import health
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_cursor.__enter__ = MagicMock(return_value=mock_cursor)
mock_cursor.__exit__ = MagicMock(return_value=False)
mock_cursor.fetchone.return_value = (2,) # Only 2 of 5 tables
mock_conn.cursor.return_value = mock_cursor
mock_connect.return_value = mock_conn
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.return_value = mock_conn
result = await health._check_migrations()
from apps.api.app.api.routes.health import _check_migrations
result = await _check_migrations()
assert result["status"] == "unhealthy"
assert "Missing tables" in result["message"]
assert "2/5" in result["message"]
async def test_check_migrations_connection_failure(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_migrations_connection_failure(self, mock_connect, mock_settings):
"""数据库连接失败时返回 unhealthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
mock_connect.side_effect = Exception("connection refused")
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_migrations
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.side_effect = Exception("connection refused")
result = await health._check_migrations()
result = await _check_migrations()
assert result["status"] == "unhealthy"
assert "Migration check failed" in result["message"]
async def test_check_migrations_in_memory(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
async def test_check_migrations_in_memory(self, mock_settings):
"""使用内存数据库时跳过迁移检查。"""
mock_settings.USE_IN_MEMORY_DB = True
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_migrations
with patch.object(health, "settings", mock_settings):
result = await health._check_migrations()
result = await _check_migrations()
assert result["status"] == "healthy"
assert "no migrations needed" in result["message"]
@@ -146,27 +154,33 @@ class TestCheckMigrations:
@pytest.mark.asyncio
class TestStartupCheck:
"""Tests for startup_check() endpoint."""
async def test_startup_all_healthy(self):
from apps.api.app.api.routes import health
@patch("apps.api.app.api.routes.health._check_migrations")
@patch("apps.api.app.api.routes.health._check_database")
async def test_startup_all_healthy(self, mock_db, mock_mig):
"""所有检查通过时返回 started。"""
mock_db.return_value = {"status": "healthy"}
mock_mig.return_value = {"status": "healthy"}
with patch.object(health, "_check_migrations", new_callable=AsyncMock) as mock_mig, patch.object(health, "_check_database", new_callable=AsyncMock) as mock_db:
mock_db.return_value = {"status": "healthy"}
mock_mig.return_value = {"status": "healthy"}
result = await health.startup_check()
from apps.api.app.api.routes.health import startup_check
result = await startup_check()
assert result["status"] == "started"
async def test_startup_db_unhealthy(self):
import json
@patch("apps.api.app.api.routes.health._check_migrations")
@patch("apps.api.app.api.routes.health._check_database")
async def test_startup_db_unhealthy(self, mock_db, mock_mig):
"""数据库不健康时返回 starting + 503。"""
mock_db.return_value = {"status": "unhealthy", "message": "fail"}
mock_mig.return_value = {"status": "healthy"}
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import startup_check
with patch.object(health, "_check_migrations", new_callable=AsyncMock) as mock_mig, patch.object(health, "_check_database", new_callable=AsyncMock) as mock_db:
mock_db.return_value = {"status": "unhealthy", "message": "fail"}
mock_mig.return_value = {"status": "healthy"}
result = await health.startup_check()
result = await startup_check()
assert result.status_code == 503
import json
body = json.loads(result.body)
assert body["status"] == "starting"
@@ -1,147 +0,0 @@
"""Tests for EditPlanService.replace_all_clips_transactional."""
from __future__ import annotations
import os
from unittest.mock import MagicMock, PropertyMock, patch
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
class TestReplaceAllClipsTransactional:
"""事务性替换片段方法测试。"""
@patch("packages.adapters.sqlalchemy_impl.models.EditPlanClipModel")
@patch("app.services.edit_plan_service.EditPlanClip")
def test_success_commits_once(self, mock_clip_cls, mock_model_cls):
"""成功时单次 commit,不 rollback。"""
from app.services.edit_plan_service import EditPlanService
db = MagicMock()
# Mock query chain for delete
query_mock = MagicMock()
query_mock.filter.return_value.delete.return_value = 3
db.query.return_value = query_mock
# Mock query chain for mark_ready (pending_with_asset)
# After the create loop, query returns empty list (no pending clips with asset)
ready_query = MagicMock()
ready_query.filter.return_value.filter.return_value.filter.return_value.all.return_value = []
db.query.side_effect = [query_mock, ready_query]
# Mock EditPlanClip.create to return a mock entity
mock_entity = MagicMock()
mock_entity.id = "clip-1"
mock_entity.plan_id = "plan-1"
mock_entity.clip_type = "main"
mock_entity.order = 0
mock_entity.asset_id = "asset-1"
mock_entity.text_content = ""
mock_entity.start_time = 0.0
mock_entity.duration = 3.0
mock_entity.transition_effect = "cut"
mock_entity.transition_duration = 0.0
mock_entity.playback_speed = 1.0
mock_entity.status.value = "pending"
mock_entity.config = {}
mock_clip_cls.create.return_value = mock_entity
# Mock the model constructor
mock_model_instance = MagicMock()
mock_model_cls.return_value = mock_model_instance
# Mock clip_repo
clip_repo = MagicMock()
clip_repo.session = db
svc = EditPlanService.__new__(EditPlanService)
svc._clip_repo = clip_repo
result = svc.replace_all_clips_transactional(
"plan-1",
[{"asset_id": "asset-1", "start_time": 0.0, "duration": 3.0, "order": 0}],
)
assert result == 1
db.commit.assert_called_once()
db.rollback.assert_not_called()
db.add.assert_called_once_with(mock_model_instance)
@patch("packages.adapters.sqlalchemy_impl.models.EditPlanClipModel")
@patch("app.services.edit_plan_service.EditPlanClip")
def test_failure_rolls_back(self, mock_clip_cls, mock_model_cls):
"""异常时自动 rollback。"""
from app.services.edit_plan_service import EditPlanService
db = MagicMock()
query_mock = MagicMock()
query_mock.filter.return_value.delete.return_value = 0
db.query.return_value = query_mock
# Simulate failure during create
mock_clip_cls.create.side_effect = ValueError("模拟异常")
clip_repo = MagicMock()
clip_repo.session = db
svc = EditPlanService.__new__(EditPlanService)
svc._clip_repo = clip_repo
with pytest.raises(ValueError, match="模拟异常"):
svc.replace_all_clips_transactional(
"plan-1",
[{"asset_id": "bad", "start_time": 0.0, "duration": 1.0, "order": 0}],
)
db.rollback.assert_called_once()
db.commit.assert_not_called()
@patch("packages.adapters.sqlalchemy_impl.models.EditPlanClipModel")
@patch("app.services.edit_plan_service.EditPlanClip")
def test_order_defaults_to_index(self, mock_clip_cls, mock_model_cls):
"""order=0 时使用索引值作为 order。"""
from app.services.edit_plan_service import EditPlanService
db = MagicMock()
query_mock = MagicMock()
query_mock.filter.return_value.delete.return_value = 0
db.query.return_value = query_mock
ready_query = MagicMock()
ready_query.filter.return_value.filter.return_value.filter.return_value.all.return_value = []
db.query.side_effect = [query_mock, ready_query]
mock_entity = MagicMock()
mock_entity.id = "clip-1"
mock_entity.plan_id = "plan-1"
mock_entity.clip_type = "main"
mock_entity.order = 0 # order=0 → 使用 i=0
mock_entity.asset_id = "a1"
mock_entity.text_content = ""
mock_entity.start_time = 0.0
mock_entity.duration = 1.0
mock_entity.transition_effect = "cut"
mock_entity.transition_duration = 0.0
mock_entity.playback_speed = 1.0
mock_entity.status.value = "pending"
mock_entity.config = {}
mock_clip_cls.create.return_value = mock_entity
mock_model_cls.return_value = MagicMock()
clip_repo = MagicMock()
clip_repo.session = db
svc = EditPlanService.__new__(EditPlanService)
svc._clip_repo = clip_repo
svc.replace_all_clips_transactional(
"plan-1",
[{"asset_id": "a1", "start_time": 0.0, "duration": 1.0, "order": 0}],
)
# order=0 → falsy → use index i=0
create_call = mock_clip_cls.create.call_args
assert create_call.kwargs["order"] == 0
-89
View File
@@ -1,89 +0,0 @@
"""packages.shared.title_overlay 单元测试。"""
from __future__ import annotations
import tempfile
from pathlib import Path
import pytest
from packages.shared.title_overlay import apply_title_to_image, wrap_title_text
@pytest.fixture
def sample_image():
"""生成一张 640x360 的纯黑测试图片。"""
from PIL import Image
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
tmp.close()
img = Image.new("RGB", (640, 360), color=(0, 0, 0))
img.save(tmp.name, "JPEG")
yield tmp.name
Path(tmp.name).unlink(missing_ok=True)
def test_apply_title_to_image_empty_text_returns_none(sample_image):
assert apply_title_to_image(sample_image, "") is None
assert apply_title_to_image(sample_image, " ") is None
def test_apply_title_to_image_draws_title(sample_image):
result = apply_title_to_image(sample_image, "测试标题")
assert result == sample_image
assert Path(sample_image).exists()
assert Path(sample_image).stat().st_size > 0
def test_apply_title_to_image_respects_position(sample_image):
for pos in ("top", "center", "bottom"):
result = apply_title_to_image(sample_image, "位置测试", position=pos)
assert result == sample_image
def test_wrap_title_text_supports_long_text():
from PIL import ImageFont
font = ImageFont.load_default()
lines = wrap_title_text("这是一个比较长的标题需要自动换行处理ABC", font, max_width=40)
assert isinstance(lines, list)
assert len(lines) >= 1
def test_wrap_title_text_respects_explicit_newline():
from PIL import ImageFont
font = ImageFont.load_default()
lines = wrap_title_text("第一行\n第二行", font, max_width=10000)
assert lines == ["第一行", "第二行"]
def test_apply_title_to_image_custom_color(sample_image):
"""自定义颜色参数能正常生成图片。"""
result = apply_title_to_image(sample_image, "彩色标题", color="#ff0000")
assert result == sample_image
assert Path(sample_image).stat().st_size > 0
def test_apply_title_to_image_short_hex_color(sample_image):
"""3 位缩写 hex 颜色也能正常解析。"""
result = apply_title_to_image(sample_image, "短色", color="#f00")
assert result == sample_image
def test_apply_title_to_image_invalid_color_fallback(sample_image):
"""无效颜色字符串 fallback 到白色,不报错。"""
result = apply_title_to_image(sample_image, "异常色", color="not-a-color")
assert result == sample_image
def test_parse_hex_color():
from packages.shared.title_overlay import _parse_hex_color
assert _parse_hex_color("#ffffff") == (255, 255, 255)
assert _parse_hex_color("#000000") == (0, 0, 0)
assert _parse_hex_color("#ff0000") == (255, 0, 0)
assert _parse_hex_color("#f00") == (255, 0, 0)
assert _parse_hex_color("") == (255, 255, 255)
assert _parse_hex_color("invalid") == (255, 255, 255)
assert _parse_hex_color("#gggggg") == (255, 255, 255)
@@ -1,126 +0,0 @@
"""Tests for _writeback_edit_plan_config in generation_tasks route.
覆盖 CI 增量覆盖率不足的代码
- generation_tasks.py 160-193 (_writeback_edit_plan_config 函数体)
- generation_tasks.py 371-372 (路由中调用该函数)
"""
from __future__ import annotations
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
from app.api.routes.generation_tasks import _writeback_edit_plan_config
@pytest.fixture
def mock_db():
"""Mock SQLAlchemy Session."""
db = MagicMock()
db.query.return_value = db
db.filter.return_value = db
return db
@pytest.fixture
def mock_plan():
"""Mock EditPlanModel instance."""
plan = MagicMock()
plan.config = {"existing_key": "existing_value"}
return plan
class TestWritebackEditPlanConfig:
"""_writeback_edit_plan_config 全分支覆盖"""
# ---- 行 160-161: plan_id 为空直接返回 ----
def test_empty_plan_id_returns_immediately(self, mock_db):
_writeback_edit_plan_config(plan_id="", task_id="task_1", title_config={"text": "hi"}, db=mock_db)
mock_db.query.assert_not_called()
mock_db.commit.assert_not_called()
def test_none_plan_id_returns_immediately(self, mock_db):
_writeback_edit_plan_config(plan_id=None, task_id="task_1", title_config=None, db=mock_db)
mock_db.query.assert_not_called()
# ---- 行 165-168: plan 不存在 → warning + 不 commit ----
def test_plan_not_found_no_commit(self, mock_db):
mock_db.first.return_value = None
_writeback_edit_plan_config(plan_id="plan_999", task_id="task_1", title_config=None, db=mock_db)
mock_db.query.assert_called_once()
mock_db.commit.assert_not_called()
# ---- 行 170-182: 正常写入 + title_config ----
def test_success_with_title_config(self, mock_db, mock_plan):
mock_db.first.return_value = mock_plan
_writeback_edit_plan_config(
plan_id="plan_123",
task_id="task_456",
title_config={"text": "标题", "font_size": 36},
db=mock_db,
)
assert mock_plan.config["generation_task_id"] == "task_456"
assert mock_plan.config["title_config"] == {"text": "标题", "font_size": 36}
assert mock_plan.config["existing_key"] == "existing_value"
mock_db.commit.assert_called_once()
# ---- 行 170-175: 正常写入、无 title_config ----
def test_success_without_title_config(self, mock_db, mock_plan):
mock_db.first.return_value = mock_plan
_writeback_edit_plan_config(plan_id="plan_123", task_id="task_789", title_config=None, db=mock_db)
assert mock_plan.config["generation_task_id"] == "task_789"
assert "title_config" not in mock_plan.config
mock_db.commit.assert_called_once()
# ---- 行 170: config 不是 dict → 兜底空 dict ----
def test_config_not_dict_uses_empty_dict(self, mock_db):
bad_plan = MagicMock()
bad_plan.config = "not_a_dict"
mock_db.first.return_value = bad_plan
_writeback_edit_plan_config(plan_id="plan_123", task_id="task_1", title_config=None, db=mock_db)
assert isinstance(bad_plan.config, dict)
assert bad_plan.config["generation_task_id"] == "task_1"
mock_db.commit.assert_called_once()
# ---- 行 183-189: DB 异常 → warning + rollback ----
def test_db_exception_triggers_rollback(self, mock_db, mock_plan):
mock_db.first.return_value = mock_plan
mock_db.commit.side_effect = RuntimeError("DB connection lost")
# 不应抛异常
_writeback_edit_plan_config(plan_id="plan_123", task_id="task_1", title_config=None, db=mock_db)
mock_db.rollback.assert_called_once()
# ---- 行 190-193: rollback 也失败 → 静默 ----
def test_rollback_failure_silent(self, mock_db, mock_plan):
mock_db.first.return_value = mock_plan
mock_db.commit.side_effect = RuntimeError("commit failed")
mock_db.rollback.side_effect = RuntimeError("rollback also failed")
# 两个异常都不应抛出
_writeback_edit_plan_config(plan_id="plan_123", task_id="task_1", title_config=None, db=mock_db)
mock_db.rollback.assert_called_once()
# ---- 行 173: title_config 为空 dict → 不写入 title_config ----
def test_empty_title_config_not_written(self, mock_db, mock_plan):
mock_db.first.return_value = mock_plan
_writeback_edit_plan_config(plan_id="plan_123", task_id="task_1", title_config={}, db=mock_db)
# 空 dict 为 falsy,不写入
assert "title_config" not in mock_plan.config
assert mock_plan.config["generation_task_id"] == "task_1"