Compare commits
19 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c61bead0e0 | |||
| 9dec75c365 | |||
| 7997f083a5 | |||
| cf195685b9 | |||
| ee4636e087 | |||
| 5678779232 | |||
| 44224cfaf6 | |||
| 98cf571ab3 | |||
| 12a9efc65d | |||
| 9c0474ef9c | |||
| 691c811cd4 | |||
| 2114b7e7ae | |||
| 7766ba1479 | |||
| ef686dde8f | |||
| 01026156ae | |||
| 42c0885813 | |||
| cd6d8615e6 | |||
| fcc7863b31 | |||
| 6af2f3c08d |
@@ -0,0 +1,72 @@
|
||||
"""Projects is_default + partial unique index for idempotent default project (Issue #1775)
|
||||
|
||||
Revision ID: 069_project_is_default
|
||||
Revises: 068_user_profile_completed
|
||||
Create Date: 2026-09-08
|
||||
|
||||
背景:
|
||||
小程序端 getOrCreateDefaultProject 在重试/并发/前端重复调用下,
|
||||
仅靠应用层"先查再插"不保证幂等,会给同一用户重复创建默认项目。
|
||||
|
||||
改动:
|
||||
1. projects 表新增 is_default 布尔列(默认 false)
|
||||
2. 部分唯一索引 uq_projects_owner_default:(owner_user_id) WHERE is_default = true
|
||||
—— 保证每个用户至多一个默认项目
|
||||
3. 存量数据回填:把名为"默认项目"的存量项目按创建时间最早者标记为 is_default=true
|
||||
(只标记不删除;存量重复项目的清理另行确认后单独执行)
|
||||
|
||||
注意:部分唯一索引依赖 PostgreSQL,不支持 downgrade 到其他方言。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "069_project_is_default"
|
||||
down_revision = "068_user_profile_completed"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 1. 新增 is_default 列
|
||||
op.add_column(
|
||||
"projects",
|
||||
sa.Column(
|
||||
"is_default",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
),
|
||||
)
|
||||
|
||||
# 2. 存量回填:每个拥有"默认项目"的用户,只把最早创建的那一个标记为默认。
|
||||
# 用 ROW_NUMBER() 取每组第一条;非"默认项目"命名的项目不标记(保守,不动用户自建项目)。
|
||||
op.execute("""
|
||||
UPDATE projects p
|
||||
SET is_default = true
|
||||
WHERE p.id IN (
|
||||
SELECT id FROM (
|
||||
SELECT id,
|
||||
ROW_NUMBER() OVER (
|
||||
PARTITION BY owner_user_id
|
||||
ORDER BY created_at ASC, id ASC
|
||||
) AS rn
|
||||
FROM projects
|
||||
WHERE name = '默认项目'
|
||||
) t
|
||||
WHERE t.rn = 1
|
||||
)
|
||||
""")
|
||||
|
||||
# 3. 部分唯一索引:每用户至多一个默认项目(只约束 is_default = true 的行)
|
||||
op.execute("""
|
||||
CREATE UNIQUE INDEX uq_projects_owner_default
|
||||
ON projects (owner_user_id)
|
||||
WHERE is_default = true
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("DROP INDEX IF EXISTS uq_projects_owner_default")
|
||||
op.drop_column("projects", "is_default")
|
||||
@@ -20,7 +20,7 @@ from packages.application import (
|
||||
GetProjectUseCase,
|
||||
ListAssetLibrariesUseCase,
|
||||
)
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind
|
||||
from packages.domain import AssetLibraryKind
|
||||
|
||||
from ._helpers import check_project_access
|
||||
|
||||
@@ -120,30 +120,11 @@ def ensure_default_library(
|
||||
|
||||
kind = AssetLibraryKind(request.kind)
|
||||
|
||||
# 查找该项目下同 kind 的素材库,返回第一个
|
||||
existing = asset_library_repository.find_by_project(request.project_id)
|
||||
for lib in existing:
|
||||
if lib.kind == kind:
|
||||
return _to_asset_library_response(lib)
|
||||
|
||||
# 不存在 → 自动创建
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
# Issue #1775: 幂等获取/创建——依赖唯一约束 uq_asset_libraries_project_kind,
|
||||
# 并发创建冲突时回滚重查返回已有记录,不再依赖应用层"先查后插",也不会 500。
|
||||
default_name = _DEFAULT_LIBRARY_NAMES.get(request.kind, f"{request.kind}素材库")
|
||||
library = AssetLibrary(
|
||||
id=str(uuid.uuid4()),
|
||||
project_id=request.project_id,
|
||||
name=default_name,
|
||||
kind=kind,
|
||||
asset_count=0,
|
||||
total_size=0,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
created = asset_library_repository.create(library)
|
||||
return _to_asset_library_response(created)
|
||||
library = asset_library_repository.get_or_create_default_library(request.project_id, kind, name=default_name)
|
||||
return _to_asset_library_response(library)
|
||||
|
||||
|
||||
@router.delete("/{library_id}", status_code=status.HTTP_204_NO_CONTENT, response_class=Response)
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_project_repository
|
||||
from app.dependencies import get_asset_library_repository, get_project_repository
|
||||
from app.schemas.project import (
|
||||
CreateProjectRequest,
|
||||
ListProjectsResponse,
|
||||
ProjectResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Response, status
|
||||
from pydantic import BaseModel
|
||||
|
||||
from packages.application import (
|
||||
CreateProjectCommand,
|
||||
@@ -16,10 +17,20 @@ from packages.application import (
|
||||
GetProjectUseCase,
|
||||
ListProjectsUseCase,
|
||||
)
|
||||
from packages.domain import AssetLibraryKind
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class DefaultContextResponse(BaseModel):
|
||||
"""幂等默认上下文响应(Issue #1775):默认项目 + 各类型默认素材库 ID。"""
|
||||
|
||||
project_id: str
|
||||
image_library_id: str
|
||||
video_library_id: str
|
||||
voice_library_id: str
|
||||
|
||||
|
||||
def _to_project_response(item) -> ProjectResponse:
|
||||
return ProjectResponse(
|
||||
id=item.id,
|
||||
@@ -72,6 +83,35 @@ def create_project(
|
||||
return _to_project_response(project)
|
||||
|
||||
|
||||
@router.post("/ensure-default", response_model=DefaultContextResponse)
|
||||
def ensure_default_project_and_libraries(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
asset_library_repository: Any = Depends(get_asset_library_repository),
|
||||
) -> DefaultContextResponse:
|
||||
"""幂等获取/创建当前用户的默认项目和三类默认素材库(Issue #1775)。
|
||||
|
||||
- 同一用户永远只有一个默认项目(部分唯一索引 uq_projects_owner_default)
|
||||
- 同一项目同 kind 永远只有一个默认素材库(唯一约束 uq_asset_libraries_project_kind)
|
||||
- 并发调用/失败重试:唯一约束冲突时返回已存在记录,不报 500
|
||||
- 项目和素材库的创建各自在仓储事务内幂等,冲突回滚后重查返回同一条
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
project = project_repository.get_or_create_default_project(user_id)
|
||||
|
||||
libraries = {}
|
||||
for kind in (AssetLibraryKind.VIDEO, AssetLibraryKind.VOICE, AssetLibraryKind.IMAGE):
|
||||
library = asset_library_repository.get_or_create_default_library(project.id, kind)
|
||||
libraries[kind] = library.id
|
||||
|
||||
return DefaultContextResponse(
|
||||
project_id=project.id,
|
||||
image_library_id=libraries[AssetLibraryKind.IMAGE],
|
||||
video_library_id=libraries[AssetLibraryKind.VIDEO],
|
||||
voice_library_id=libraries[AssetLibraryKind.VOICE],
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/{project_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
|
||||
def delete_project(
|
||||
project_id: str,
|
||||
|
||||
@@ -106,6 +106,10 @@ def list_templates(
|
||||
tag: str | None = Query(None, description="按标签筛选"),
|
||||
keyword: str | None = Query(None, description="按名称关键词搜索"),
|
||||
mode: str | None = Query(None, description="按剪辑模式筛选"),
|
||||
valid_only: bool = Query(
|
||||
False,
|
||||
description="仅返回已配置片段的模板(剪辑页传 true;模板编辑器不传,可查看全部模板含草稿)",
|
||||
),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
|
||||
) -> ListTemplatesResponse:
|
||||
@@ -116,6 +120,7 @@ def list_templates(
|
||||
tag=tag,
|
||||
keyword=keyword,
|
||||
mode=mode,
|
||||
valid_only=valid_only,
|
||||
)
|
||||
use_case = ListTemplatesUseCase(template_repository)
|
||||
templates = use_case.execute(user_id, skip=skip, limit=limit, filter=tpl_filter)
|
||||
|
||||
@@ -36,17 +36,11 @@ from app.services.asset_segment_tracker import (
|
||||
remove_used_segment,
|
||||
)
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from app.services.edit_template_service import EditTemplateService, TemplateNotFoundError
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.adapters.sqlalchemy_impl.template_clip_config_repository import (
|
||||
SQLAlchemyTemplateClipConfigRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.template_repository import (
|
||||
SQLAlchemyTemplateRepository,
|
||||
)
|
||||
from packages.domain.plan_generator_utils import (
|
||||
_calc_random_start_time,
|
||||
build_scene_segments,
|
||||
@@ -399,73 +393,42 @@ def _safe_segment_duration(value, default: float) -> float:
|
||||
|
||||
def _get_template_segments(
|
||||
template_id: str,
|
||||
user_id: str,
|
||||
tpl_svc: EditTemplateService,
|
||||
db: Session,
|
||||
) -> list[tuple[int, float, float]]:
|
||||
"""获取模板的片段配置(顺序、最短时长、最长时长).
|
||||
|
||||
优先从新模板系统(template_clip_configs)查询,
|
||||
若不存在则回退到旧模板系统(template_segments)。
|
||||
单一数据源:模板主表为 ``templates``(用户自建,归属 user_id)/
|
||||
``edit_templates``(全局模板库),片段配置主表为 ``template_clip_configs``
|
||||
(由 ``EditTemplateService.list_clip_configs_for_editor`` 统一读取)。
|
||||
|
||||
不再使用"新表抛异常 → 降级直查配置表 → 再降级查 segments"的异常控制流,
|
||||
也不在正常请求中打印 ``ValueError: 模板不存在`` 堆栈。
|
||||
|
||||
Args:
|
||||
template_id: 模板 ID
|
||||
user_id: 当前登录用户 ID(用于归属校验)
|
||||
tpl_svc: 模板编辑器服务
|
||||
|
||||
Returns:
|
||||
[(segment_order, duration_min, duration_max), ...] 按 order 排序
|
||||
[(segment_order, duration_min, duration_max), ...] 按 order 排序;
|
||||
模板存在但未配置片段时返回空列表。
|
||||
|
||||
Raises:
|
||||
TemplateNotFoundError: 模板不存在、已删除或不归属于当前用户。
|
||||
"""
|
||||
# 优先查新模板系统
|
||||
try:
|
||||
clip_configs = tpl_svc.list_clip_configs(template_id)
|
||||
if clip_configs:
|
||||
result = []
|
||||
for cc in clip_configs:
|
||||
dur_min = _safe_segment_duration(cc.min_duration, _DEFAULT_EDITOR_CLIP_DURATION)
|
||||
dur_max = _safe_segment_duration(
|
||||
cc.max_duration or cc.min_duration,
|
||||
_DEFAULT_EDITOR_CLIP_DURATION,
|
||||
)
|
||||
dur_min, dur_max = min(dur_min, dur_max), max(dur_min, dur_max)
|
||||
result.append((cc.order, dur_min, dur_max))
|
||||
return sorted(result, key=lambda x: x[0])
|
||||
except Exception:
|
||||
logger.warning("新模板系统查询clip_configs失败(主表可能不存在),直接查clip_configs表", exc_info=True)
|
||||
clip_configs = tpl_svc.list_clip_configs_for_editor(template_id, user_id)
|
||||
|
||||
# 兜底:直接查 template_clip_configs 表(片段表有 template_id 外键,不依赖模板主表)
|
||||
try:
|
||||
direct_repo = SQLAlchemyTemplateClipConfigRepository(db)
|
||||
direct_configs = direct_repo.list_by_template(template_id)
|
||||
if direct_configs:
|
||||
result = []
|
||||
for cc in direct_configs:
|
||||
dur_min = _safe_segment_duration(cc.min_duration, _DEFAULT_EDITOR_CLIP_DURATION)
|
||||
dur_max = _safe_segment_duration(
|
||||
cc.max_duration or cc.min_duration,
|
||||
_DEFAULT_EDITOR_CLIP_DURATION,
|
||||
)
|
||||
dur_min, dur_max = min(dur_min, dur_max), max(dur_min, dur_max)
|
||||
result.append((cc.order, dur_min, dur_max))
|
||||
return sorted(result, key=lambda x: x[0])
|
||||
except Exception:
|
||||
logger.warning("直接查clip_configs表也失败,继续回退旧系统", exc_info=True)
|
||||
|
||||
# 回退到旧模板系统(template_segments表)
|
||||
try:
|
||||
old_repo = SQLAlchemyTemplateRepository(db)
|
||||
segments = old_repo.list_segments(template_id)
|
||||
if segments:
|
||||
result = []
|
||||
for s in segments:
|
||||
dur_min = _safe_segment_duration(s.duration_min, _DEFAULT_EDITOR_CLIP_DURATION)
|
||||
dur_max = _safe_segment_duration(s.duration_max, _DEFAULT_EDITOR_CLIP_DURATION)
|
||||
dur_min, dur_max = min(dur_min, dur_max), max(dur_min, dur_max)
|
||||
result.append((s.segment_order, dur_min, dur_max))
|
||||
return sorted(result, key=lambda x: x[0])
|
||||
except Exception:
|
||||
logger.warning("旧模板系统查询segments失败", exc_info=True)
|
||||
|
||||
# 所有途径都失败:模板没有片段配置(可能是无效测试模板)
|
||||
logger.error(
|
||||
"模板无片段配置:template_id=%s(可能是 is_active=false 的无效模板)",
|
||||
template_id,
|
||||
)
|
||||
return []
|
||||
result = []
|
||||
for cc in clip_configs:
|
||||
dur_min = _safe_segment_duration(cc.min_duration, _DEFAULT_EDITOR_CLIP_DURATION)
|
||||
dur_max = _safe_segment_duration(
|
||||
cc.max_duration or cc.min_duration,
|
||||
_DEFAULT_EDITOR_CLIP_DURATION,
|
||||
)
|
||||
dur_min, dur_max = min(dur_min, dur_max), max(dur_min, dur_max)
|
||||
result.append((cc.order, dur_min, dur_max))
|
||||
return sorted(result, key=lambda x: x[0])
|
||||
|
||||
|
||||
def _recommended_time_conflicts(
|
||||
@@ -668,13 +631,21 @@ def create_clips_from_assets_editor(
|
||||
7. 素材时长为 0 或缺失时报 400,不创建无效片段
|
||||
"""
|
||||
tpl_svc, plan_svc = services
|
||||
user_id = str(current_user.user.id)
|
||||
|
||||
# 1. 查询模板 segments
|
||||
segments = _get_template_segments(template_id, tpl_svc, db)
|
||||
# 1. 查询模板片段配置。模板不存在/已删除/无权限 → 404;
|
||||
# 模板存在但确实未配置片段 → 422(配置错误,与 404 区分)。
|
||||
try:
|
||||
segments = _get_template_segments(template_id, user_id, tpl_svc)
|
||||
except TemplateNotFoundError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="模板不存在或无权访问",
|
||||
) from exc
|
||||
if not segments:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="模板没有片段配置,无法创建片段",
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="模板未配置片段",
|
||||
)
|
||||
|
||||
# 防御:schema validator 已过滤 null/空串,这里再归一化一次,
|
||||
|
||||
@@ -41,29 +41,33 @@ def get_draft_plan_id(
|
||||
这是模板编辑器路由的核心依赖——所有编辑器端点都先经过这里,
|
||||
确保 template_id → plan_id 的映射始终存在。
|
||||
|
||||
兼容策略:优先从新模板系统(edit_templates 表)查找,
|
||||
若不存在则回退到旧模板系统(templates 表),确保用户自建模板可用。
|
||||
模板读取遵循单一数据源、显式判定(不使用异常降级):
|
||||
- 用户自建模板在旧表 ``templates``(归属 user_id,is_active=True);
|
||||
- 全局模板在新表 ``edit_templates``(无 user_id,全局可读)。
|
||||
模板不存在、已删除或不归属于当前用户时,一律返回 404。
|
||||
"""
|
||||
tpl_svc, plan_svc = services
|
||||
user_id = str(current_user.user.id)
|
||||
|
||||
# 0. 门禁:校验模板存在且可访问(即使草稿已缓存命中也要校验,
|
||||
# 避免模板被删除/无权访问后仍可通过既有草稿 plan 继续操作)。
|
||||
old_repo = SQLAlchemyTemplateRepository(db)
|
||||
old_template = old_repo.get_active(template_id, user_id)
|
||||
is_global_template = tpl_svc.get_template(template_id) is not None
|
||||
if old_template is None and not is_global_template:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模板不存在")
|
||||
|
||||
# 1. 草稿已存在 → 直接返回
|
||||
draft = tpl_svc.get_template_draft(template_id)
|
||||
if draft is not None:
|
||||
return draft.id
|
||||
|
||||
# 2. 新系统有模板 → 用新服务创建草稿
|
||||
if tpl_svc.get_template(template_id) is not None:
|
||||
# 2. 全局模板(新系统)→ 用新服务创建草稿
|
||||
if is_global_template:
|
||||
draft = tpl_svc.create_template_draft(template_id, user_id=user_id)
|
||||
return draft.id
|
||||
|
||||
# 3. 回退到旧模板系统(templates 表)
|
||||
old_repo = SQLAlchemyTemplateRepository(db)
|
||||
old_template = old_repo.get(template_id, user_id=user_id)
|
||||
if old_template is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模板不存在")
|
||||
|
||||
# 4. 基于旧模板创建草稿计划
|
||||
# 3. 旧模板(templates 表)→ 基于旧模板创建草稿计划
|
||||
from app.services.plan_generator_service import PlanGeneratorService
|
||||
|
||||
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
|
||||
|
||||
@@ -150,7 +150,7 @@ def rollback_template(
|
||||
try:
|
||||
tpl = tpl_svc.rollback_to_version(template_id, request.version)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
|
||||
|
||||
clip_configs = tpl_svc.list_clip_configs(template_id)
|
||||
return EditorRollbackResponse(
|
||||
|
||||
@@ -208,6 +208,8 @@ def _create_pending_asset(
|
||||
find-or-create:prepare 阶段已按 file_hash/client_upload_id 预建的占位记录
|
||||
会被 find_by_library_and_file_hash/find_by_library_and_client_upload_id 命中,
|
||||
直接复用并补齐字段(避免 pre-create + complete 重复建两条)。
|
||||
|
||||
Issue #1776: 素材库计数由 asset_repository.create() 自动维护。
|
||||
"""
|
||||
# 1. 按 client_upload_id / file_hash 查找现有记录
|
||||
existing = None
|
||||
@@ -389,6 +391,7 @@ async def prepare_direct_upload(
|
||||
file_size=request.file_size,
|
||||
)
|
||||
pending_asset_id = pending.id
|
||||
# Issue #1776: 计数由 asset_repository.create() 自动维护
|
||||
except Exception as error:
|
||||
# 预建失败不阻塞签名:complete 仍可按 OSS 文件 + hash 兜底去重
|
||||
logger.warning("预建 asset 占位失败,降级走 old flow: %s", error)
|
||||
@@ -474,6 +477,7 @@ async def complete_direct_upload(
|
||||
client_upload_id=request.client_upload_id,
|
||||
file_size=request.file_size,
|
||||
)
|
||||
# Issue #1776: 计数由 asset_repository.create() 自动维护
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=request.project_id,
|
||||
@@ -568,6 +572,7 @@ async def upload_asset(
|
||||
file_hash=file_hash,
|
||||
client_upload_id=client_upload_id,
|
||||
)
|
||||
# Issue #1776: 计数由 asset_repository.create() 自动维护
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=project_id,
|
||||
|
||||
@@ -789,20 +789,22 @@ class EditPlanService:
|
||||
if plan and hasattr(plan, "config") and plan.config:
|
||||
rhythm_template = plan.config.get("rhythm_template")
|
||||
|
||||
# #1768:先获取素材时长,传入 plan_clip_durations 用于最大片段钳制
|
||||
asset_ids = [c.asset_id for c in clips if c.asset_id]
|
||||
durations = self.get_asset_durations(asset_ids)
|
||||
asset_durations_for_plan = [durations.get(c.asset_id, 0.0) for c in clips]
|
||||
|
||||
target = plan_clip_durations(
|
||||
len(clips),
|
||||
voice,
|
||||
transition_effects=[c.transition_effect for c in clips],
|
||||
transition_durations=[float(c.transition_duration or 0.0) for c in clips],
|
||||
rhythm_template=rhythm_template,
|
||||
asset_durations=asset_durations_for_plan,
|
||||
)
|
||||
if not target:
|
||||
return None
|
||||
|
||||
# 素材时长(短素材起点钳 0)
|
||||
asset_ids = [c.asset_id for c in clips if c.asset_id]
|
||||
durations = self.get_asset_durations(asset_ids)
|
||||
|
||||
clips_data: list[dict] = []
|
||||
for i, c in enumerate(clips):
|
||||
dur = float(target[i])
|
||||
@@ -932,6 +934,17 @@ class EditPlanService:
|
||||
rhythm_templates_for_variants.append(template)
|
||||
logger.info("变体 %d 节奏模板: plan=%s template=%s", idx, plan_ids[idx], template)
|
||||
|
||||
# #1767:BGM 池差异化分配(让批量变体使用不同 BGM / 段落 / 音量)
|
||||
from packages.domain.bgm_pool import allocate_bgm_pool_for_variants
|
||||
|
||||
source_bgm_config = {}
|
||||
source_plan = self.get_plan(source_plan_id)
|
||||
if source_plan and source_plan.config:
|
||||
source_bgm_config = source_plan.config.get("bgm", {}) or {}
|
||||
|
||||
variant_seeds_for_bgm = [rng.randint(0, 999999) for _ in plan_ids]
|
||||
bgm_pool_assignments = allocate_bgm_pool_for_variants(source_bgm_config, variant_seeds_for_bgm)
|
||||
|
||||
# 为每个变体生成独立视觉扰动参数(让批量视频画面本身更不同)
|
||||
from packages.domain.variant_plan_selector import generate_visual_perturbation
|
||||
|
||||
@@ -950,8 +963,20 @@ class EditPlanService:
|
||||
|
||||
pixel_pert = generate_pixel_perturbation(rng)
|
||||
config_update["pixel_perturbation"] = pixel_pert
|
||||
# #1767:写入 BGM 池分配(覆盖 bgm 配置中的 preset_id / audio_offset / volume_adjust_db)
|
||||
if idx < len(bgm_pool_assignments):
|
||||
existing_bgm = dict((source_plan.config or {}).get("bgm", {}) or {})
|
||||
existing_bgm.update(bgm_pool_assignments[idx])
|
||||
config_update["bgm"] = existing_bgm
|
||||
self.update_plan_config(pid, config_update)
|
||||
logger.info("变体 %d 视觉扰动+像素扰动: plan=%s vis=%s pix=%s", idx, pid, perturbation, pixel_pert)
|
||||
logger.info(
|
||||
"变体 %d 视觉扰动+像素扰动+BGM池: plan=%s vis=%s pix=%s bgm=%s",
|
||||
idx,
|
||||
pid,
|
||||
perturbation,
|
||||
pixel_pert,
|
||||
bgm_pool_assignments[idx] if idx < len(bgm_pool_assignments) else None,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("变体 %d 视觉扰动生成失败(不阻断): plan=%s", idx, pid)
|
||||
|
||||
@@ -1187,7 +1212,7 @@ class EditPlanService:
|
||||
config_asset_ids_count = len((plan.config or {}).get("asset_ids", []))
|
||||
clips_with_asset_count = sum(1 for c in clips if c.asset_id)
|
||||
logger.info(
|
||||
"can_generate 诊断: plan=%s status=%s total_clips=%d " "clips_with_asset=%d config_asset_ids_count=%d",
|
||||
"can_generate 诊断: plan=%s status=%s total_clips=%d clips_with_asset=%d config_asset_ids_count=%d",
|
||||
plan_id,
|
||||
plan.status,
|
||||
len(clips),
|
||||
@@ -1199,7 +1224,7 @@ class EditPlanService:
|
||||
config_asset_ids = (plan.config or {}).get("asset_ids", [])
|
||||
if config_asset_ids:
|
||||
logger.warning(
|
||||
"can_generate 最后防线触发: plan=%s clips=%d 均无素材," "从 config.asset_ids(%d个) 自动分配",
|
||||
"can_generate 最后防线触发: plan=%s clips=%d 均无素材,从 config.asset_ids(%d个) 自动分配",
|
||||
plan_id,
|
||||
len(clips),
|
||||
len(config_asset_ids),
|
||||
@@ -1231,7 +1256,7 @@ class EditPlanService:
|
||||
return False, "没有可渲染的就绪片段,自动修复后仍未分配素材"
|
||||
else:
|
||||
logger.warning(
|
||||
"can_generate 失败: plan=%s clips=%d 均无素材," "且 config.asset_ids 为空,无法自动修复",
|
||||
"can_generate 失败: plan=%s clips=%d 均无素材,且 config.asset_ids 为空,无法自动修复",
|
||||
plan_id,
|
||||
len(clips),
|
||||
)
|
||||
|
||||
@@ -34,6 +34,17 @@ from packages.domain.template_clip_converter import (
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TemplateNotFoundError(Exception):
|
||||
"""模板不存在、已删除或当前用户无权访问.
|
||||
|
||||
与"模板存在但无片段配置"区分:路由层应映射为 HTTP 404。
|
||||
"""
|
||||
|
||||
def __init__(self, template_id: str) -> None:
|
||||
self.template_id = template_id
|
||||
super().__init__(f"模板不存在: {template_id}")
|
||||
|
||||
|
||||
class EditTemplateService:
|
||||
"""模板管理服务
|
||||
|
||||
@@ -217,7 +228,14 @@ class EditTemplateService:
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> List[TemplateClipConfig]:
|
||||
"""列出模板的片段配置"""
|
||||
"""列出模板的片段配置
|
||||
|
||||
注意:本方法要求模板存在于新表 ``edit_templates``(全局模板库),
|
||||
主要服务于新模板系统的写入/发布路径。用户自建模板存放在旧表
|
||||
``templates``,不在 ``edit_templates`` 中,读取其片段配置请改用
|
||||
:meth:`list_clip_configs_for_editor`,后者直接读取片段配置主表
|
||||
``template_clip_configs``,不依赖新模板主表、也不靠异常降级。
|
||||
"""
|
||||
# 确保模板存在
|
||||
self.get_template_or_raise(template_id)
|
||||
return self._clip_config_repo.list_by_template(
|
||||
@@ -227,6 +245,52 @@ class EditTemplateService:
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
def list_clip_configs_for_editor(
|
||||
self,
|
||||
template_id: str,
|
||||
user_id: str,
|
||||
*,
|
||||
clip_type: Optional[ClipType] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> List[TemplateClipConfig]:
|
||||
"""编辑器读取模板片段配置的单一数据源入口.
|
||||
|
||||
片段配置主表是 ``template_clip_configs``(直接读取,不抛异常、不降级)。
|
||||
模板主表按双表现状显式判定,不使用 try/except 控制流:
|
||||
|
||||
1. 用户自建模板在旧表 ``templates``(归属 user_id)→ 校验归属与未删除后直接读;
|
||||
2. 全局模板在新表 ``edit_templates``(无 user_id,全局可读)→ 直接读;
|
||||
3. 两者都没有 → 模板不存在/无权限,抛 :class:`TemplateNotFoundError`。
|
||||
|
||||
Args:
|
||||
template_id: 模板 ID
|
||||
user_id: 当前登录用户 ID(用于旧表模板归属校验)
|
||||
|
||||
Raises:
|
||||
TemplateNotFoundError: 模板不存在、已删除或不归属于当前用户。
|
||||
"""
|
||||
# 1) 用户自建模板(旧表 templates,归属 user_id)
|
||||
if self._clip_config_repo.template_owned_by(template_id, user_id):
|
||||
return self._clip_config_repo.list_by_template(
|
||||
template_id,
|
||||
clip_type=clip_type,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
# 2) 全局模板(新表 edit_templates,无 user_id,全局可读)
|
||||
if self._template_repo.get(template_id) is not None:
|
||||
return self._clip_config_repo.list_by_template(
|
||||
template_id,
|
||||
clip_type=clip_type,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
# 3) 两表都没有:不存在 / 已删除 / 无权限
|
||||
raise TemplateNotFoundError(template_id)
|
||||
|
||||
def get_clip_config(self, config_id: str) -> Optional[TemplateClipConfig]:
|
||||
"""获取片段配置详情"""
|
||||
return self._clip_config_repo.get(config_id)
|
||||
|
||||
@@ -39,6 +39,9 @@ from packages.domain.video_filter_builder import (
|
||||
)
|
||||
from packages.domain.video_filter_builder import build_concat_filter as _build_concat_filter_func
|
||||
from packages.domain.video_filter_builder import build_filter_complex as _build_filter_complex
|
||||
from packages.domain.video_filter_builder import (
|
||||
build_title_drawtext_filter,
|
||||
)
|
||||
from packages.domain.video_filter_builder import build_xfade_filter as _build_xfade_filter_func
|
||||
from packages.domain.video_filter_builder import chain_filters as _chain_filters_func
|
||||
from packages.domain.video_filter_builder import has_audio as _has_audio_func
|
||||
@@ -248,6 +251,27 @@ class VideoComposeService:
|
||||
transitions=[c.transition_effect for c in ready_clips],
|
||||
)
|
||||
|
||||
# ── #1789 标题 drawtext 滤镜叠加 ──
|
||||
# 从 plan.config 读取 title_config,生成 drawtext 滤镜链入 filter_complex
|
||||
title_cfg = (plan.config or {}).get("title", {}) or {}
|
||||
if not isinstance(title_cfg, dict):
|
||||
title_cfg = {}
|
||||
# 同时兼容 plan.config["title_config"](API 回写路径)
|
||||
if not title_cfg.get("text") and not title_cfg.get("content"):
|
||||
title_cfg_alt = (plan.config or {}).get("title_config", {}) or {}
|
||||
if isinstance(title_cfg_alt, dict) and (title_cfg_alt.get("text") or title_cfg_alt.get("content")):
|
||||
title_cfg = title_cfg_alt
|
||||
drawtext_filter = build_title_drawtext_filter(title_cfg, output_width, output_height)
|
||||
if drawtext_filter:
|
||||
# 将最终输出标签从 [outv] 改为 [composed],再链入 drawtext → [outv]
|
||||
filter_complex = filter_complex.replace("[outv]", "[composed]")
|
||||
filter_complex += f";[composed]{drawtext_filter}[outv]"
|
||||
logger.info(
|
||||
"[#1789] 标题 drawtext 滤镜已注入: plan_id=%s text=%s",
|
||||
plan_id,
|
||||
(title_cfg.get("text") or title_cfg.get("content") or "")[:30],
|
||||
)
|
||||
|
||||
# 构建完整命令
|
||||
command: list[str] = ["ffmpeg", "-y"]
|
||||
|
||||
|
||||
@@ -5,10 +5,22 @@ import apiClient from "../client"
|
||||
import { getOrCreateDefaultProject } from "../projects"
|
||||
import type { AssetLibraryItem } from "./types"
|
||||
|
||||
/** 获取当前用户的所有素材库 */
|
||||
export const getAssetLibraries = async (): Promise<AssetLibraryItem[]> => {
|
||||
const response = await apiClient.get("/asset-libraries")
|
||||
return response.data.items || []
|
||||
/**
|
||||
* 获取当前用户的素材库
|
||||
*
|
||||
* @param kind 可选,按素材库类型过滤(video/voice/image)。
|
||||
* 后端 GET /asset-libraries 支持 kind 查询参数;这里同时在前端再按返回数据的
|
||||
* kind 字段兜底过滤一次,保证旧后端(忽略未知 query 参数)也不会把其他类型的库
|
||||
* 混进来(#1777:视频选择器只展示视频库)。
|
||||
*/
|
||||
export const getAssetLibraries = async (
|
||||
kind?: AssetLibraryItem["kind"],
|
||||
): Promise<AssetLibraryItem[]> => {
|
||||
const response = await apiClient.get<{ items?: AssetLibraryItem[] }>("/asset-libraries", {
|
||||
params: kind ? { kind } : undefined,
|
||||
})
|
||||
const items = response.data.items || []
|
||||
return kind ? items.filter((lib) => lib.kind === kind) : items
|
||||
}
|
||||
|
||||
/** 创建素材库(自动获取或创建默认项目以提供 project_id) */
|
||||
|
||||
@@ -55,6 +55,19 @@ apiClient.interceptors.response.use(
|
||||
async (error: AxiosError<{ detail?: string; message?: string; msg?: string }>) => {
|
||||
const originalRequest = error.config as InternalAxiosRequestConfig & {
|
||||
_retry?: boolean
|
||||
/**
|
||||
* 调用方自行处理错误提示时置 true:拦截器跳过全局 message 弹窗(#1777)。
|
||||
* 例如失效模板自动回退时,调用方会弹「原模板已失效,已自动切换」,
|
||||
* 不再叠加后端原始错误文案。错误仍会 reject,不影响 catch 逻辑。
|
||||
*/
|
||||
_silentErrorToast?: boolean
|
||||
}
|
||||
|
||||
// 调用方声明自行处理提示:标记为已展示,跳过下面所有全局 message 弹窗
|
||||
if (originalRequest?._silentErrorToast) {
|
||||
// eslint-disable-next-line @typescript-eslint/no-explicit-any
|
||||
;(error as any).__msgShown = true
|
||||
return Promise.reject(error)
|
||||
}
|
||||
|
||||
// 401 → 尝试刷新 Token
|
||||
|
||||
@@ -12,17 +12,25 @@ import type {
|
||||
ListCategoriesResponse,
|
||||
} from "./types"
|
||||
|
||||
/** 获取模板列表 */
|
||||
/** 获取模板列表
|
||||
*
|
||||
* valid_only=true 时请求后端仅返回已配置片段的模板(剪辑页选模板使用,
|
||||
* 避免选中无片段配置的模板导致 from-assets 400,#1769/#1772);
|
||||
* 后端尚未支持该参数时会忽略未知 query 字段,前端再按 segments/is_active 兜底过滤。
|
||||
* 模板编辑器/我的模板不传,可查看全部模板(含未配置片段的草稿)。
|
||||
*/
|
||||
export const getEditingTemplates = async (params?: {
|
||||
category?: string
|
||||
tag?: string
|
||||
skip?: number
|
||||
limit?: number
|
||||
validOnly?: boolean
|
||||
}): Promise<EditingTemplate[]> => {
|
||||
const response = await apiClient.get<ListTemplatesResponse>("/templates", {
|
||||
params: {
|
||||
skip: params?.skip ?? 0,
|
||||
limit: params?.limit ?? 50,
|
||||
...(params?.validOnly ? { valid_only: true } : {}),
|
||||
},
|
||||
})
|
||||
let list = response.data.items
|
||||
|
||||
@@ -91,7 +91,7 @@ export async function createClipsFromAssets(
|
||||
assetIds: string[],
|
||||
clipType = "main",
|
||||
requiredClipsCount?: number,
|
||||
opts?: { signal?: AbortSignal },
|
||||
opts?: { signal?: AbortSignal; silentErrorToast?: boolean },
|
||||
): Promise<ClipsFromAssetsResponse> {
|
||||
const body: Record<string, unknown> = {
|
||||
asset_ids: assetIds,
|
||||
@@ -104,7 +104,12 @@ export async function createClipsFromAssets(
|
||||
const response = await apiClient.post<ClipsFromAssetsResponse>(
|
||||
`/templates/${templateId}/editor/clips/from-assets`,
|
||||
body,
|
||||
{ timeout: 60000, signal: opts?.signal },
|
||||
{
|
||||
timeout: 60000,
|
||||
signal: opts?.signal,
|
||||
// _silentErrorToast 由 api/client.ts 响应拦截器读取(抑制全局错误 toast,#1777)
|
||||
...(opts?.silentErrorToast ? ({ _silentErrorToast: true } as Record<string, unknown>) : {}),
|
||||
},
|
||||
)
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -43,11 +43,17 @@ export async function updateEditPlanClips(
|
||||
templateId: string,
|
||||
clips: EditPlanClipInput[],
|
||||
signal?: AbortSignal,
|
||||
/** 为 true 时抑制全局错误 toast(调用方自行提示,如失效模板回退 #1777) */
|
||||
silentErrorToast?: boolean,
|
||||
): Promise<{ count: number }> {
|
||||
const response = await apiClient.put(
|
||||
`/templates/${templateId}/editor/clips`,
|
||||
{ clips },
|
||||
{ signal },
|
||||
{
|
||||
signal,
|
||||
// _silentErrorToast 由 api/client.ts 响应拦截器读取(抑制全局错误 toast)
|
||||
...(silentErrorToast ? ({ _silentErrorToast: true } as Record<string, unknown>) : {}),
|
||||
},
|
||||
)
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -6440,3 +6440,47 @@
|
||||
opacity: 0.6;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
|
||||
/* ═══ 标题设置 — 颜色预设 ═══ */
|
||||
.ep-color-presets {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
|
||||
.ep-color-swatch {
|
||||
width: 24px;
|
||||
height: 24px;
|
||||
border-radius: 4px;
|
||||
border: 2px solid transparent;
|
||||
cursor: pointer;
|
||||
transition: border-color 0.15s, transform 0.1s;
|
||||
}
|
||||
|
||||
.ep-color-swatch:hover {
|
||||
transform: scale(1.1);
|
||||
}
|
||||
|
||||
.ep-color-swatch.active {
|
||||
border-color: var(--ep-primary, #4f8cff);
|
||||
}
|
||||
|
||||
.ep-color-picker {
|
||||
width: 28px;
|
||||
height: 28px;
|
||||
border: none;
|
||||
border-radius: 4px;
|
||||
cursor: pointer;
|
||||
padding: 0;
|
||||
background: none;
|
||||
}
|
||||
|
||||
.ep-color-picker::-webkit-color-swatch-wrapper {
|
||||
padding: 0;
|
||||
}
|
||||
|
||||
.ep-color-picker::-webkit-color-swatch {
|
||||
border: 1px solid rgba(255, 255, 255, 0.2);
|
||||
border-radius: 4px;
|
||||
}
|
||||
|
||||
@@ -196,6 +196,8 @@ const EditingPlanner: React.FC = () => {
|
||||
|
||||
{/* 右栏 260px:设置面板 */}
|
||||
<RightPanel
|
||||
titleConfig={titleConfig}
|
||||
onTitleConfigChange={setTitleConfig}
|
||||
rightTab={rightTab}
|
||||
onTabChange={setRightTab}
|
||||
selectedClip={clipOps.selectedClip}
|
||||
|
||||
@@ -5,12 +5,15 @@
|
||||
import React from "react"
|
||||
import type { ClipPropertiesPanelProps } from "@/pages/editing-planner/types/clipProperties"
|
||||
import SubtitleSettingsSection from "./clip-properties/SubtitleSettingsSection"
|
||||
import TitleSettingsSection from "./clip-properties/TitleSettingsSection"
|
||||
import BgmSettingsSection from "./clip-properties/BgmSettingsSection"
|
||||
import ClipDetailSection from "./clip-properties/ClipDetailSection"
|
||||
import StatsSection from "./clip-properties/StatsSection"
|
||||
import { useVoicePreview } from "@/pages/editing-planner/hooks/useVoicePreview"
|
||||
|
||||
const ClipPropertiesPanel: React.FC<ClipPropertiesPanelProps> = ({
|
||||
titleConfig,
|
||||
onTitleConfigChange,
|
||||
selectedClip,
|
||||
subtitleSettings,
|
||||
bgmSettings,
|
||||
@@ -40,6 +43,11 @@ const ClipPropertiesPanel: React.FC<ClipPropertiesPanelProps> = ({
|
||||
|
||||
return (
|
||||
<div className="ep-right-panel">
|
||||
{/* ═══ 标题设置 — #1789 ═══ */}
|
||||
{titleConfig && onTitleConfigChange && (
|
||||
<TitleSettingsSection config={titleConfig} onChange={onTitleConfigChange} />
|
||||
)}
|
||||
|
||||
{/* ═══ 字幕设置 ═══ */}
|
||||
<SubtitleSettingsSection
|
||||
settings={subtitleSettings}
|
||||
|
||||
@@ -5,9 +5,12 @@ import type { ClipData } from "../types"
|
||||
import type { SubtitleStyleConfig } from "../types/subtitle"
|
||||
import type { BgmMixConfig } from "@/api/bgm"
|
||||
import type { TemplateMode } from "@/api/editing-planner"
|
||||
import type { TitleConfig } from "@/api/template-editor"
|
||||
import type { AssetItem } from "@/api/assets"
|
||||
|
||||
interface RightPanelProps {
|
||||
titleConfig?: TitleConfig
|
||||
onTitleConfigChange?: (config: TitleConfig | ((prev: TitleConfig) => TitleConfig)) => void
|
||||
rightTab: "properties" | "clips"
|
||||
onTabChange: (tab: "properties" | "clips") => void
|
||||
// 属性 tab
|
||||
@@ -46,6 +49,8 @@ interface RightPanelProps {
|
||||
}
|
||||
|
||||
const RightPanel: React.FC<RightPanelProps> = ({
|
||||
titleConfig,
|
||||
onTitleConfigChange,
|
||||
rightTab,
|
||||
onTabChange,
|
||||
selectedClip,
|
||||
@@ -118,6 +123,8 @@ const RightPanel: React.FC<RightPanelProps> = ({
|
||||
>["onBgmSettingsChange"]
|
||||
return (
|
||||
<ClipPropertiesPanel
|
||||
titleConfig={titleConfig}
|
||||
onTitleConfigChange={onTitleConfigChange}
|
||||
selectedClip={selectedClip}
|
||||
subtitleSettings={sub}
|
||||
bgmSettings={bgm}
|
||||
|
||||
+136
@@ -0,0 +1,136 @@
|
||||
/**
|
||||
* 标题设置区块 — #1789
|
||||
* 提供字号滑块、字体预设、位置、颜色等控制入口
|
||||
*/
|
||||
import React from "react"
|
||||
import type { TitleConfig } from "@/api/template-editor"
|
||||
import { POSITION_OPTIONS, FONT_OPTIONS } from "@/pages/editing-planner/constants/clipProperties"
|
||||
|
||||
interface TitleSettingsSectionProps {
|
||||
config: TitleConfig
|
||||
onChange: (config: TitleConfig | ((prev: TitleConfig) => TitleConfig)) => void
|
||||
}
|
||||
|
||||
const TITLE_COLOR_PRESETS = [
|
||||
"#ffffff",
|
||||
"#000000",
|
||||
"#ff4444",
|
||||
"#ffaa00",
|
||||
"#44ff44",
|
||||
"#4488ff",
|
||||
"#ff44ff",
|
||||
"#ffff44",
|
||||
]
|
||||
|
||||
const TitleSettingsSection: React.FC<TitleSettingsSectionProps> = ({ config, onChange }) => {
|
||||
const update = (partial: Partial<TitleConfig>) => {
|
||||
onChange((prev: TitleConfig) => ({ ...prev, ...partial }))
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="ep-settings-section">
|
||||
<div className="ep-section-title">
|
||||
<span className="ep-section-icon">📝</span>
|
||||
标题设置
|
||||
</div>
|
||||
|
||||
{/* AI 自动选择开关 */}
|
||||
<div className="ep-toggle-row">
|
||||
<span className="ep-toggle-label">AI 自动选择</span>
|
||||
<div
|
||||
className={`ep-toggle ${config.ai_auto_select ? "active" : ""}`}
|
||||
onClick={() => update({ ai_auto_select: !config.ai_auto_select })}
|
||||
>
|
||||
<div className="ep-toggle-knob" />
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{!config.ai_auto_select && (
|
||||
<>
|
||||
{/* 标题文本 */}
|
||||
<div className="ep-field">
|
||||
<label className="ep-field-label">标题文本</label>
|
||||
<input
|
||||
className="ep-form-select"
|
||||
type="text"
|
||||
placeholder="输入标题内容"
|
||||
value={config.content}
|
||||
onChange={(e) => update({ content: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 位置 */}
|
||||
<div className="ep-field">
|
||||
<label className="ep-field-label">位置</label>
|
||||
<select
|
||||
className="ep-form-select"
|
||||
value={config.position}
|
||||
onChange={(e) => update({ position: e.target.value })}
|
||||
>
|
||||
{POSITION_OPTIONS.map((opt) => (
|
||||
<option key={opt.value} value={opt.value}>
|
||||
{opt.label}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</div>
|
||||
|
||||
{/* 字体预设 */}
|
||||
<div className="ep-field">
|
||||
<label className="ep-field-label">字体</label>
|
||||
<select
|
||||
className="ep-form-select"
|
||||
value={config.font_preset}
|
||||
onChange={(e) => update({ font_preset: e.target.value })}
|
||||
>
|
||||
{FONT_OPTIONS.map((f) => (
|
||||
<option key={f} value={f}>
|
||||
{f}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</div>
|
||||
|
||||
{/* 字号滑块 */}
|
||||
<div className="ep-field">
|
||||
<label className="ep-field-label">字号</label>
|
||||
<div className="ep-slider-row">
|
||||
<input
|
||||
className="ep-slider"
|
||||
type="range"
|
||||
min={12}
|
||||
max={72}
|
||||
value={config.font_size}
|
||||
onChange={(e) => update({ font_size: Number(e.target.value) })}
|
||||
/>
|
||||
<span className="ep-slider-value">{config.font_size}px</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 颜色 */}
|
||||
<div className="ep-field">
|
||||
<label className="ep-field-label">颜色</label>
|
||||
<div className="ep-color-presets">
|
||||
{TITLE_COLOR_PRESETS.map((color) => (
|
||||
<div
|
||||
key={color}
|
||||
className={`ep-color-swatch${config.font_color === color ? " active" : ""}`}
|
||||
style={{ backgroundColor: color }}
|
||||
onClick={() => update({ font_color: color })}
|
||||
/>
|
||||
))}
|
||||
<input
|
||||
type="color"
|
||||
className="ep-color-picker"
|
||||
value={config.font_color}
|
||||
onChange={(e) => update({ font_color: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default TitleSettingsSection
|
||||
@@ -4,6 +4,7 @@
|
||||
import type { ClipData } from "./clip"
|
||||
import type { TemplateMode } from "@/api/editing-planner"
|
||||
import type { AssetItem } from "@/api/assets"
|
||||
import type { TitleConfig } from "@/api/template-editor"
|
||||
|
||||
export interface SubtitleSettings {
|
||||
enabled: boolean
|
||||
@@ -28,6 +29,10 @@ export interface BgmSettings {
|
||||
}
|
||||
|
||||
export interface ClipPropertiesPanelProps {
|
||||
/** 标题配置 — #1789 */
|
||||
titleConfig?: TitleConfig
|
||||
/** 标题配置变更 */
|
||||
onTitleConfigChange?: (config: TitleConfig | ((prev: TitleConfig) => TitleConfig)) => void
|
||||
selectedClip: ClipData | null
|
||||
subtitleSettings: SubtitleSettings
|
||||
bgmSettings: BgmSettings
|
||||
|
||||
@@ -13,6 +13,8 @@ export const formatTrimTime = (sec: number): string => {
|
||||
/** 生成时间标尺刻度 */
|
||||
export const generateRulerMarks = (totalDuration: number, step: number): number[] => {
|
||||
const marks: number[] = []
|
||||
// #1790: 无片段时不显示时间刻度
|
||||
if (totalDuration <= 0) return marks
|
||||
for (let t = 0; t <= totalDuration + step; t += step) {
|
||||
marks.push(t)
|
||||
}
|
||||
|
||||
@@ -45,6 +45,7 @@ const GeneratePage: React.FC = () => {
|
||||
selectedTemplate,
|
||||
setSelectedTemplate,
|
||||
userTemplates,
|
||||
handleInvalidTemplate,
|
||||
selectedMaterials,
|
||||
setSelectedMaterials,
|
||||
materialMode,
|
||||
@@ -416,15 +417,11 @@ const GeneratePage: React.FC = () => {
|
||||
/* ── 最终成片(单视频右侧播放) ── */
|
||||
const finalVideo = generatedVideos[0]
|
||||
|
||||
/* ── 布局 class:步骤4标题页=预览+标题侧栏;步骤5/6批量=整行宽;步骤1~3=整行宽 ── */
|
||||
/* ── 布局 class:步骤4标题页=预览+标题侧栏两栏;其余步骤(含步骤5确认生成、步骤6封面)=整行宽 ── */
|
||||
const layoutClassName = useMemo(() => {
|
||||
if (currentStep < 4) return "xx-generate-layout full-width"
|
||||
if (currentStep === 4) return "xx-generate-layout step4-layout"
|
||||
// 步骤5:全宽+内容居中(单视频视频播放器居中,批量网格居中)
|
||||
if (currentStep === 5) return "xx-generate-layout full-width"
|
||||
// 步骤6:封面选择保持两栏布局
|
||||
return isBatch ? "xx-generate-layout full-width" : "xx-generate-layout"
|
||||
}, [currentStep, isBatch])
|
||||
return "xx-generate-layout full-width"
|
||||
}, [currentStep])
|
||||
|
||||
/* ================================================================
|
||||
渲染
|
||||
@@ -524,6 +521,7 @@ const GeneratePage: React.FC = () => {
|
||||
selectedVoice={selectedVoice}
|
||||
onSelectedVoiceChange={setSelectedVoice}
|
||||
onServerClipsChange={setServerClips}
|
||||
onTemplateInvalid={handleInvalidTemplate}
|
||||
generating={generating}
|
||||
generated={generated}
|
||||
generateError={generateError}
|
||||
@@ -545,18 +543,7 @@ const GeneratePage: React.FC = () => {
|
||||
selectedVariantIds={selectedVariantIds}
|
||||
/>
|
||||
|
||||
<GenerateStepActions
|
||||
currentStep={currentStep}
|
||||
onPrev={goPrev}
|
||||
onNext={goNext}
|
||||
onConfirmGenerate={handleConfirmGenerate}
|
||||
generating={generating}
|
||||
generated={generated}
|
||||
generateError={generateError}
|
||||
selectedCount={isBatch ? selectedVariantIds.length : 1}
|
||||
/>
|
||||
|
||||
{/* ════ 步骤5(单视频):成片播放器内联居中(#1761) ════ */}
|
||||
{/* ════ 步骤5(单视频):成片播放器置于按钮上方、居中展示 ════ */}
|
||||
{currentStep === 5 && !isBatch && generated && finalVideo && (
|
||||
<div
|
||||
style={{
|
||||
@@ -580,6 +567,7 @@ const GeneratePage: React.FC = () => {
|
||||
<video
|
||||
src={finalVideo.download_url || finalVideo.file_url}
|
||||
controls
|
||||
autoPlay
|
||||
style={{
|
||||
width: "auto",
|
||||
maxWidth: "100%",
|
||||
@@ -607,55 +595,18 @@ const GeneratePage: React.FC = () => {
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* ════ 步骤6(单视频):右侧成片播放器 ════ */}
|
||||
{currentStep >= 6 && !isBatch && generated && finalVideo && (
|
||||
<div className="xx-generate-right-col">
|
||||
<div
|
||||
className="xx-inline-video-player"
|
||||
style={{
|
||||
display: "flex",
|
||||
justifyContent: "center",
|
||||
alignItems: "center",
|
||||
background: "#000",
|
||||
borderRadius: 12,
|
||||
padding: 8,
|
||||
}}
|
||||
>
|
||||
<video
|
||||
src={finalVideo.download_url || finalVideo.file_url}
|
||||
controls
|
||||
autoPlay={currentStep === 5}
|
||||
// 竖屏自适应(#1750):成片固定 1080×1920(9:16),元数据到达前按 9:16 占位,
|
||||
// 到达后浏览器按真实宽高比 contain;黑底居中杜绝左右大黑边
|
||||
style={{
|
||||
width: "auto",
|
||||
maxWidth: "100%",
|
||||
maxHeight: "70vh",
|
||||
aspectRatio: "9 / 16",
|
||||
objectFit: "contain",
|
||||
borderRadius: 8,
|
||||
}}
|
||||
poster={finalVideo.thumbnail_url || undefined}
|
||||
/>
|
||||
<div style={{ display: "flex", gap: 8, marginTop: 12, justifyContent: "center" }}>
|
||||
<button className="xx-btn xx-btn-ghost xx-btn-sm" onClick={handleDownload}>
|
||||
⬇️ 下载
|
||||
</button>
|
||||
<button className="xx-btn xx-btn-ghost xx-btn-sm" onClick={handleShare}>
|
||||
🔗 分享
|
||||
</button>
|
||||
<button
|
||||
className="xx-btn xx-btn-ghost xx-btn-sm"
|
||||
onClick={() => navigate("/app/products")}
|
||||
>
|
||||
📁 前往成片库
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
<GenerateStepActions
|
||||
currentStep={currentStep}
|
||||
onPrev={goPrev}
|
||||
onNext={goNext}
|
||||
onConfirmGenerate={handleConfirmGenerate}
|
||||
generating={generating}
|
||||
generated={generated}
|
||||
generateError={generateError}
|
||||
selectedCount={isBatch ? selectedVariantIds.length : 1}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 数量选择弹窗 */}
|
||||
|
||||
@@ -50,6 +50,8 @@ export interface GenerateStepContentProps {
|
||||
selectedVoice: string
|
||||
onSelectedVoiceChange: (id: string) => void
|
||||
onServerClipsChange: (clips: EditPlanClip[]) => void
|
||||
/** 当前模板创建片段被判失效(404/400/422)时的自动回退回调(#1777) */
|
||||
onTemplateInvalid?: () => boolean
|
||||
/* 生成 */
|
||||
generating: boolean
|
||||
generated: boolean
|
||||
@@ -108,6 +110,7 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
selectedVoice,
|
||||
onSelectedVoiceChange,
|
||||
onServerClipsChange,
|
||||
onTemplateInvalid,
|
||||
generating,
|
||||
generated,
|
||||
generateError,
|
||||
@@ -153,6 +156,7 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
selectedTemplate={selectedTemplate}
|
||||
templateSegments={templateSegments}
|
||||
onServerClipsChange={onServerClipsChange}
|
||||
onTemplateInvalid={onTemplateInvalid}
|
||||
/>
|
||||
)
|
||||
case 3:
|
||||
@@ -189,7 +193,7 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
/>
|
||||
)
|
||||
case 5:
|
||||
/* 确认生成页:批量=逐任务进度网格;单视频=进度状态卡(成片播放器在左侧大区域) */
|
||||
/* 确认生成页:批量=逐任务进度网格;单视频=仅渲染进度/失败状态(完成后只显示成片播放器,播放器在按钮上方) */
|
||||
if (previewCount > 1) {
|
||||
return (
|
||||
<BatchGenerationGrid
|
||||
@@ -199,10 +203,10 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
/>
|
||||
)
|
||||
}
|
||||
/* 单视频:渲染进度 / 失败重试 / 完成提示(成片播放器在右侧栏) */
|
||||
/* 单视频:生成中显示进度卡、失败显示重试卡;生成完成后不再渲染提示卡,页面只保留成片播放器+操作按钮 */
|
||||
if (generated && !generating && !generateError) return null
|
||||
return (
|
||||
<div className="xx-form-section">
|
||||
<h3>🎬 确认生成</h3>
|
||||
{generating && (
|
||||
<div className="xx-gen-progress-card">
|
||||
<div className="xx-gen-progress-header">
|
||||
@@ -234,14 +238,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
{generated && !generating && (
|
||||
<div className="xx-gen-success-card">
|
||||
<div className="xx-gen-success-info">
|
||||
<div className="xx-gen-success-title">✅ 视频生成完成!</div>
|
||||
<div className="xx-gen-success-sub">右侧可预览成片,点击「下一步」选择封面</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
case 6:
|
||||
|
||||
@@ -23,6 +23,8 @@ interface Step2MaterialSelectProps {
|
||||
templateSegments?: TemplateSegment[]
|
||||
/** 服务端 clips 创建成功后的回调 */
|
||||
onServerClipsChange?: (clips: EditPlanClip[]) => void
|
||||
/** 当前模板创建片段返回 404/400/422(模板失效)时的自动回退回调(#1777) */
|
||||
onTemplateInvalid?: () => boolean
|
||||
}
|
||||
|
||||
const Step2MaterialSelect: React.FC<Step2MaterialSelectProps> = (props) => {
|
||||
@@ -36,16 +38,25 @@ const Step2MaterialSelect: React.FC<Step2MaterialSelectProps> = (props) => {
|
||||
|
||||
<div className="xx-form-field" style={{ marginTop: 12 }}>
|
||||
<label>选择视频库</label>
|
||||
<select
|
||||
value={m.selectedLibraryId}
|
||||
onChange={(e) => m.setSelectedLibraryId(e.target.value)}
|
||||
>
|
||||
{m.libraries.map((lib) => (
|
||||
<option key={lib.id} value={lib.id}>
|
||||
{lib.name}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
{m.libraries.length === 0 && !m.materialsLoading ? (
|
||||
<div className="xx-empty-state">
|
||||
<p>暂无视频素材库</p>
|
||||
<p style={{ fontSize: 13, color: "var(--text-tertiary)" }}>
|
||||
请先在「素材库」中创建视频素材库并上传视频
|
||||
</p>
|
||||
</div>
|
||||
) : (
|
||||
<select
|
||||
value={m.selectedLibraryId}
|
||||
onChange={(e) => m.setSelectedLibraryId(e.target.value)}
|
||||
>
|
||||
{m.libraries.map((lib) => (
|
||||
<option key={lib.id} value={lib.id}>
|
||||
{lib.name}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{m.materialMode === "manual" && (
|
||||
|
||||
@@ -1,117 +0,0 @@
|
||||
import React from "react"
|
||||
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
interface GenerationStatusProps {
|
||||
generating: boolean
|
||||
generated: boolean
|
||||
generateError: string | null
|
||||
progress: number
|
||||
generatedVideos: GeneratedVideo[]
|
||||
getGenerationPhase: (progress: number) => { icon: string; label: string }
|
||||
onScrollToPreview: () => void
|
||||
onRetry: () => void
|
||||
onDismissError: () => void
|
||||
}
|
||||
|
||||
const GenerationStatus: React.FC<GenerationStatusProps> = ({
|
||||
generating,
|
||||
generated,
|
||||
generateError,
|
||||
progress,
|
||||
generatedVideos,
|
||||
getGenerationPhase,
|
||||
onScrollToPreview,
|
||||
onRetry,
|
||||
onDismissError,
|
||||
}) => {
|
||||
return (
|
||||
<div style={{ marginTop: 16 }}>
|
||||
{!generating && !generated && !generateError && (
|
||||
<div className="xx-gen-progress-card" style={{ opacity: 0.85 }}>
|
||||
<div className="xx-gen-progress-header">
|
||||
<div className="xx-gen-progress-icon">🎬</div>
|
||||
<div className="xx-gen-progress-info">
|
||||
<div className="xx-gen-progress-phase">尚未开始生成视频</div>
|
||||
<div className="xx-gen-progress-sub">
|
||||
请返回「选择标题」步骤,点击「确认生成视频」开始渲染最终视频
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
{generating && (
|
||||
<div className="xx-gen-progress-card">
|
||||
<div className="xx-gen-progress-header">
|
||||
<div className="xx-gen-progress-icon">
|
||||
<LoadingOutlined />
|
||||
</div>
|
||||
<div className="xx-gen-progress-info">
|
||||
<div className="xx-gen-progress-phase">
|
||||
{getGenerationPhase(progress).icon} {getGenerationPhase(progress).label}
|
||||
</div>
|
||||
<div className="xx-gen-progress-sub">预计还需 1-2 分钟,请稍候…</div>
|
||||
</div>
|
||||
<div className="xx-gen-progress-percent">{Math.round(progress)}%</div>
|
||||
</div>
|
||||
<div className="xx-gen-progress-bar">
|
||||
<div
|
||||
className="xx-gen-progress-bar-fill"
|
||||
style={{ width: `${Math.min(Math.round(progress), 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
<div className="xx-gen-progress-tip">
|
||||
💡 生成过程中可以切换到其他页面操作,完成后会自动通知
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
{generated && !generating && (
|
||||
<div className="xx-gen-success-card">
|
||||
<div className="xx-gen-success-icon">
|
||||
<CheckCircleFilled style={{ fontSize: 32, color: "#52c41a" }} />
|
||||
</div>
|
||||
<div className="xx-gen-success-info">
|
||||
<div className="xx-gen-success-title">视频生成完成!</div>
|
||||
<div className="xx-gen-success-sub">
|
||||
共生成 {generatedVideos.length} 条视频,可在右侧预览或前往成片库查看
|
||||
</div>
|
||||
</div>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-btn xx-btn-primary xx-btn-sm"
|
||||
onClick={onScrollToPreview}
|
||||
>
|
||||
查看结果
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
{generateError && !generating && (
|
||||
<div className="xx-gen-error-card">
|
||||
<div className="xx-gen-error-icon">
|
||||
<CloseCircleOutlined style={{ fontSize: 28, color: "#ef4444" }} />
|
||||
</div>
|
||||
<div className="xx-gen-error-info">
|
||||
<div className="xx-gen-error-title">生成失败</div>
|
||||
<div className="xx-gen-error-msg">
|
||||
{typeof generateError === "string" ? generateError : JSON.stringify(generateError)}
|
||||
</div>
|
||||
</div>
|
||||
<div style={{ display: "flex", gap: 8 }}>
|
||||
<button type="button" className="xx-btn xx-btn-primary xx-btn-sm" onClick={onRetry}>
|
||||
🔄 重试
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-btn xx-btn-ghost xx-btn-sm"
|
||||
onClick={onDismissError}
|
||||
>
|
||||
知道了
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default GenerationStatus
|
||||
@@ -117,10 +117,6 @@
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
|
||||
.xx-generate-layout.full-width .xx-generate-right-col {
|
||||
display: none;
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
左侧表单区 generate-form
|
||||
============================================================ */
|
||||
@@ -2323,28 +2319,6 @@
|
||||
生成结果(右侧)
|
||||
================================================================ */
|
||||
|
||||
.xx-generate-right-col {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
gap: 16px;
|
||||
}
|
||||
|
||||
/* ── 内联视频播放器(右侧) ── */
|
||||
.xx-inline-video-player {
|
||||
width: 100%;
|
||||
max-width: 320px;
|
||||
background: var(--bg-surface, #fff);
|
||||
border: 1px solid var(--border-primary, #e2e8f0);
|
||||
border-radius: 16px;
|
||||
padding: 16px;
|
||||
box-shadow: 0 2px 8px rgba(0, 0, 0, 0.04);
|
||||
}
|
||||
|
||||
.xx-inline-video-player video {
|
||||
background: #000;
|
||||
}
|
||||
|
||||
.xx-preview-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
@@ -2710,11 +2684,12 @@
|
||||
|
||||
/* ── 封面设置区域改造样式 ── */
|
||||
|
||||
/* 封面操作按钮区 */
|
||||
/* 封面操作按钮区(单视频全宽页居中) */
|
||||
.xx-cover-actions {
|
||||
display: flex;
|
||||
gap: 12px;
|
||||
margin-bottom: 12px;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
/* 已选模板文字 */
|
||||
@@ -3202,10 +3177,11 @@
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
/* ── 批量封面网格 ── */
|
||||
/* ── 批量封面网格(单卡/少卡时居中排列,卡片限宽不拉伸) ── */
|
||||
.xx-cover-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fill, minmax(180px, 1fr));
|
||||
grid-template-columns: repeat(auto-fill, minmax(180px, 220px));
|
||||
justify-content: center;
|
||||
gap: 16px;
|
||||
}
|
||||
|
||||
|
||||
@@ -8,11 +8,22 @@ import type { AssetItem } from "@/api/assets"
|
||||
* 管理素材库列表、当前选中库、素材列表加载
|
||||
*/
|
||||
export function useMaterialLibrary() {
|
||||
/* ── 素材库数据 API ── */
|
||||
const { data: libraries = [] } = useQuery({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
/* ── 素材库数据 API ──
|
||||
* Step2 是视频选片,只拉取 kind=video 的素材库(#1777):
|
||||
* 后端按 kind 查询参数过滤,前端 getAssetLibraries("video") 再兜底过滤一次,
|
||||
* 避免配音库(voice)/图片库(image) 混进「选择视频库」下拉。
|
||||
* queryKey 带 kind,与素材管理页/配音页的 ["asset-libraries"] 全量缓存隔离。
|
||||
*/
|
||||
const { data: allLibraries = [] } = useQuery({
|
||||
queryKey: ["asset-libraries", "video"],
|
||||
queryFn: () => getAssetLibraries("video"),
|
||||
staleTime: 60_000,
|
||||
})
|
||||
// 前端兜底过滤:仅保留 kind=video 的素材库(后端按 kind 查询参数过滤)
|
||||
const libraries = useMemo(
|
||||
() => allLibraries.filter((lib) => lib.kind === "video"),
|
||||
[allLibraries],
|
||||
)
|
||||
const [selectedLibraryId, setSelectedLibraryId] = useState<string>("")
|
||||
|
||||
// 自动选中第一个视频库
|
||||
|
||||
@@ -40,6 +40,8 @@ export interface GenerateFormState {
|
||||
selectedTemplate: string
|
||||
setSelectedTemplate: (id: string) => void
|
||||
userTemplates: EditingTemplate[]
|
||||
/** 当前选中模板在创建片段时被判失效(404/400/422)后的运行时自动回退 */
|
||||
handleInvalidTemplate: () => boolean
|
||||
|
||||
/* 素材 */
|
||||
selectedMaterials: string[]
|
||||
@@ -131,7 +133,8 @@ export const useGenerateFormState = (): GenerateFormState => {
|
||||
const [currentStep, setCurrentStep] = useState(1)
|
||||
|
||||
/* ── 模板选择 ── */
|
||||
const { selectedTemplate, setSelectedTemplate, userTemplates } = useTemplateSelection()
|
||||
const { selectedTemplate, setSelectedTemplate, userTemplates, handleInvalidTemplate } =
|
||||
useTemplateSelection()
|
||||
|
||||
/* ── source_edit_plan_id:仅取 URL 参数,无则 null 让后端兜底 ── */
|
||||
// selectedTemplate 是模板 ID 而非 edit_plan_id,不能混淆;
|
||||
@@ -228,6 +231,7 @@ export const useGenerateFormState = (): GenerateFormState => {
|
||||
selectedTemplate,
|
||||
setSelectedTemplate,
|
||||
userTemplates,
|
||||
handleInvalidTemplate,
|
||||
selectedMaterials,
|
||||
setSelectedMaterials,
|
||||
materialMode,
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
/**
|
||||
* 失效模板判定与自动回退工具(#1777)
|
||||
*
|
||||
* 背景:用户进入生成页后,之前选中的模板可能已被删除、或从未配置片段。
|
||||
* 调用片段相关接口(PUT/POST /templates/{id}/editor/clips[...]/from-assets)时:
|
||||
* - 模板不存在 → 后端返回 404(并行工单 #1774 把「模板不存在」统一为该状态码)
|
||||
* - 模板无片段配置 → 当前部分场景返回 400(detail 含「片段配置」),
|
||||
* 参数校验类错误返回 422
|
||||
* 这三类响应都说明「当前选中的模板不可用于生成」,应清除失效选择并自动切换到
|
||||
* 第一个有效模板,同时提示用户,而不是让页面卡死、无任何反馈。
|
||||
*/
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
|
||||
/** 失效模板相关的 HTTP 状态码 */
|
||||
const INVALID_TEMPLATE_STATUSES = new Set([404, 400, 422])
|
||||
|
||||
/**
|
||||
* 从任意抛出值(axios 错误)提取 HTTP 状态码。
|
||||
* 非 axios 错误 / 无响应时返回 null。
|
||||
*/
|
||||
export function getHttpStatus(err: unknown): number | null {
|
||||
if (!err || typeof err !== "object") return null
|
||||
const status = (err as { response?: { status?: number }; status?: number })?.response?.status
|
||||
return typeof status === "number" ? status : null
|
||||
}
|
||||
|
||||
/** 安全提取后端错误文本(detail/message/msg,422 数组也兜底拼一下) */
|
||||
function extractErrorText(err: unknown): string {
|
||||
if (!err || typeof err !== "object") return ""
|
||||
const data = (err as { response?: { data?: unknown } })?.response?.data
|
||||
if (!data) return ""
|
||||
try {
|
||||
const text = JSON.stringify(data)
|
||||
return typeof text === "string" ? text : ""
|
||||
} catch {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断一次 clips/from-assets 请求失败是否因为「模板失效」。
|
||||
*
|
||||
* 严格判定,避免把无关的 400/422(例如素材参数问题)误判为模板失效:
|
||||
* - 404:模板/编辑计划不存在,一定是模板失效
|
||||
* - 400:仅当后端文本明确提到「片段配置」(无片段配置无法创建片段)才判定
|
||||
* - 422:参数校验类,from-assets 场景下命中「片段/segments」相关字段才判定
|
||||
*/
|
||||
export function isInvalidTemplateError(err: unknown): boolean {
|
||||
const status = getHttpStatus(err)
|
||||
if (status === null || !INVALID_TEMPLATE_STATUSES.has(status)) return false
|
||||
if (status === 404) return true
|
||||
|
||||
const text = extractErrorText(err)
|
||||
if (status === 400) {
|
||||
// 后端当前返回:「模板没有片段配置,无法创建片段」
|
||||
return /片段配置|没有片段|无片段|segments?|clip.*config/i.test(text)
|
||||
}
|
||||
// 422:FastAPI 校验错误,命中模板片段相关字段
|
||||
return /segment|clip|片段|模板/i.test(text)
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断模板是否可用于生成(有效模板)。
|
||||
*
|
||||
* 有效 = 处于激活态(is_active !== false,字段缺失视为 true 兼容旧后端)
|
||||
* 且至少配置了一个片段。
|
||||
* 与后端 valid_only 过滤口径保持一致(#1769/#1772),这里是前端双保险。
|
||||
*/
|
||||
export function isValidTemplate(template: EditingTemplate | null | undefined): boolean {
|
||||
if (!template) return false
|
||||
if (template.is_active === false) return false
|
||||
return (template.segments?.length ?? 0) > 0
|
||||
}
|
||||
|
||||
/** 从模板列表中取出第一个有效模板,没有则返回 null */
|
||||
export function findFirstValidTemplate(
|
||||
templates: EditingTemplate[] | null | undefined,
|
||||
): EditingTemplate | null {
|
||||
if (!Array.isArray(templates)) return null
|
||||
return templates.find(isValidTemplate) ?? null
|
||||
}
|
||||
@@ -1,22 +1,89 @@
|
||||
import { useState, useEffect } from "react"
|
||||
import { useState, useEffect, useRef, useCallback, useMemo } from "react"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { getEditingTemplates } from "@/api/editing-planner"
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
import { findFirstValidTemplate, isValidTemplate } from "./templateFallback"
|
||||
|
||||
/** 失效模板自动切换的提示文案 */
|
||||
export const INVALID_TEMPLATE_FALLBACK_TOAST = "原模板已失效,已自动切换"
|
||||
|
||||
export function useTemplateSelection() {
|
||||
// selectedTemplate 纯内存状态,绝不写入 localStorage/sessionStorage/URL,
|
||||
// 因此失效模板 ID 不会被持久化、刷新后也不会恢复(#1777 要求 4)
|
||||
const [selectedTemplate, setSelectedTemplate] = useState("")
|
||||
const { data: userTemplates = [] } = useQuery<EditingTemplate[]>({
|
||||
|
||||
const { data: allTemplates = [] } = useQuery<EditingTemplate[]>({
|
||||
queryKey: ["generate-templates"],
|
||||
queryFn: () => getEditingTemplates(),
|
||||
// valid_only:后端过滤掉没有片段配置的无效模板(#1769/#1772)。
|
||||
// 旧后端忽略该 query 参数时,下方 isValidTemplate 前端兜底再过滤一次。
|
||||
queryFn: () => getEditingTemplates({ validOnly: true }),
|
||||
staleTime: 60_000,
|
||||
})
|
||||
|
||||
/* 模板加载完成后自动选中第一个 */
|
||||
useEffect(() => {
|
||||
if (userTemplates.length > 0 && !selectedTemplate) {
|
||||
setSelectedTemplate(userTemplates[0].id)
|
||||
}
|
||||
}, [userTemplates, selectedTemplate])
|
||||
// 双保险:后端 valid_only 已过滤,前端再按 is_active + segments 兜底,
|
||||
// 保证下拉/自动选择只包含可用于生成的有效模板。
|
||||
// 用 useMemo 缓存引用,避免每次渲染都 .filter 创建新数组,
|
||||
// 导致下游 useTitleCoverSync effect 无限触发、覆盖用户手动修改(#1789)
|
||||
const validTemplates = useMemo(() => allTemplates.filter(isValidTemplate), [allTemplates])
|
||||
const userTemplates = validTemplates
|
||||
|
||||
return { selectedTemplate, setSelectedTemplate, userTemplates }
|
||||
// 用 ref 持有最新值,供稳定回调 handleInvalidTemplate 使用(避免闭包拿到旧值)
|
||||
const templatesRef = useRef(validTemplates)
|
||||
templatesRef.current = validTemplates
|
||||
const selectedRef = useRef(selectedTemplate)
|
||||
selectedRef.current = selectedTemplate
|
||||
// 已提示过失效的模板 ID,避免用户停留在失效模板上时 clips 防抖请求反复弹 toast;
|
||||
// 用户手动切换/成功切换后重置,保证下一个失效模板仍能提示
|
||||
const fallbackNotifiedRef = useRef<string>("")
|
||||
|
||||
/* 自动选择:模板加载完成且当前未选中时,自动选中第一个有效模板。
|
||||
* 用户手动选择(setSelectedTemplate 被显式调用)后 selectedTemplate 非空,
|
||||
* 本 effect 直接 return,绝不覆盖用户的手动选择(#1777 要求 4:手动优先)。 */
|
||||
useEffect(() => {
|
||||
if (selectedTemplate) return
|
||||
const firstValid = validTemplates[0]
|
||||
if (firstValid) {
|
||||
setSelectedTemplate(firstValid.id)
|
||||
}
|
||||
}, [validTemplates, selectedTemplate])
|
||||
|
||||
/** 用户手动选择模板:优先级最高,重置失效提示标记 */
|
||||
const handleSelectTemplate = useCallback((id: string) => {
|
||||
fallbackNotifiedRef.current = ""
|
||||
setSelectedTemplate(id)
|
||||
}, [])
|
||||
|
||||
/**
|
||||
* 运行时失效回退(#1777 要求 3):
|
||||
* 创建片段接口返回 404(模板不存在)/ 400/422(模板无片段配置)时调用。
|
||||
* - 清除失效选择,自动切换到第一个有效模板,并 toast 提示;
|
||||
* - 没有有效模板时清空选择,Step1 展示明确的「暂无可用模板」空状态引导,
|
||||
* 不让用户卡在失效模板上。
|
||||
* 返回 true 表示已按「模板失效」处理(调用方可据此静默原始错误提示)。
|
||||
*/
|
||||
const handleInvalidTemplate = useCallback((): boolean => {
|
||||
const current = selectedRef.current
|
||||
// 同一个失效模板只提示一次(clips 防抖 effect 在素材/模板变化时会反复触发)
|
||||
if (current && fallbackNotifiedRef.current === current) return true
|
||||
|
||||
const fallback = findFirstValidTemplate(templatesRef.current)
|
||||
fallbackNotifiedRef.current = current || "__empty__"
|
||||
if (fallback) {
|
||||
setSelectedTemplate(fallback.id)
|
||||
message.warning(INVALID_TEMPLATE_FALLBACK_TOAST)
|
||||
} else {
|
||||
// 没有任何有效模板:清空选择,交由 Step1 空状态引导用户去模板编辑器创建
|
||||
setSelectedTemplate("")
|
||||
message.warning("当前没有可用模板,请先在「模板编辑器」中创建并配置片段")
|
||||
}
|
||||
return true
|
||||
}, [])
|
||||
|
||||
return {
|
||||
selectedTemplate,
|
||||
setSelectedTemplate: handleSelectTemplate,
|
||||
userTemplates,
|
||||
handleInvalidTemplate,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { useEffect } from "react"
|
||||
import { useEffect, useRef } from "react"
|
||||
import type { TitleSettings } from "../../types"
|
||||
import type { CoverConfig } from "../../types/cover"
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
@@ -11,7 +11,11 @@ interface UseTitleCoverSyncOptions {
|
||||
}
|
||||
|
||||
/**
|
||||
* 当选中模板变化时,自动同步标题和封面配置
|
||||
* 当选中模板变化时,自动同步标题和封面配置。
|
||||
*
|
||||
* #1789 修复:userTemplates 用 ref 持有最新值,不放入依赖数组。
|
||||
* 否则每次渲染 .filter() 创建的新数组引用都会触发 effect,
|
||||
* 从模板 title_config 覆盖用户手动修改(如字号滑块拖动),导致回弹。
|
||||
*/
|
||||
export function useTitleCoverSync({
|
||||
selectedTemplate,
|
||||
@@ -19,8 +23,12 @@ export function useTitleCoverSync({
|
||||
setTitleSettings,
|
||||
setCoverSettings,
|
||||
}: UseTitleCoverSyncOptions) {
|
||||
// 用 ref 持有最新 userTemplates,避免数组引用变化导致 effect 反复触发
|
||||
const templatesRef = useRef(userTemplates)
|
||||
templatesRef.current = userTemplates
|
||||
|
||||
useEffect(() => {
|
||||
const tpl = userTemplates.find((t) => t.id === selectedTemplate)
|
||||
const tpl = templatesRef.current.find((t) => t.id === selectedTemplate)
|
||||
if (tpl?.title_config) {
|
||||
setTitleSettings((prev: TitleSettings) => ({
|
||||
...prev,
|
||||
@@ -43,5 +51,7 @@ export function useTitleCoverSync({
|
||||
thumbnail_url: tpl.cover_config!.thumbnail_url || prev.thumbnail_url,
|
||||
}))
|
||||
}
|
||||
}, [selectedTemplate, userTemplates, setTitleSettings, setCoverSettings])
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [selectedTemplate, setTitleSettings, setCoverSettings])
|
||||
// ↑ 移除 userTemplates,只在 selectedTemplate 真正变化时触发
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import { updateEditPlanClips, createClipsFromAssets, getEditPlanClips } from "@/
|
||||
import { useMaterialLibrary } from "./step2-materials/useMaterialLibrary"
|
||||
import { useSmartMatch } from "./step2-materials/useSmartMatch"
|
||||
import { useDraftAutoSave } from "./useDraftAutoSave"
|
||||
import { isInvalidTemplateError } from "./useGenerateFormState/templateFallback"
|
||||
|
||||
interface UseStep2MaterialsProps {
|
||||
materialMode: "manual" | "auto"
|
||||
@@ -24,6 +25,8 @@ interface UseStep2MaterialsProps {
|
||||
templateSegments?: TemplateSegment[]
|
||||
/** 服务端 clips 创建成功后的回调,用于通知预览播放器 */
|
||||
onServerClipsChange?: (clips: EditPlanClip[]) => void
|
||||
/** 当前模板创建片段返回 404/400/422(模板失效)时的自动回退回调(#1777) */
|
||||
onTemplateInvalid?: () => boolean
|
||||
}
|
||||
|
||||
export function useStep2Materials({
|
||||
@@ -36,6 +39,7 @@ export function useStep2Materials({
|
||||
selectedTemplate,
|
||||
templateSegments,
|
||||
onServerClipsChange,
|
||||
onTemplateInvalid,
|
||||
}: UseStep2MaterialsProps) {
|
||||
const {
|
||||
libraries,
|
||||
@@ -102,6 +106,8 @@ export function useStep2Materials({
|
||||
selectedTemplateRef.current = selectedTemplate
|
||||
const onServerClipsChangeRef = useRef(onServerClipsChange)
|
||||
onServerClipsChangeRef.current = onServerClipsChange
|
||||
const onTemplateInvalidRef = useRef(onTemplateInvalid)
|
||||
onTemplateInvalidRef.current = onTemplateInvalid
|
||||
|
||||
useEffect(() => {
|
||||
const tid = selectedTemplateRef.current
|
||||
@@ -123,11 +129,12 @@ export function useStep2Materials({
|
||||
const requiredClipsCount = segs.length > 0 ? segs.length : undefined
|
||||
|
||||
try {
|
||||
// 1. 清空旧片段
|
||||
await updateEditPlanClips(tid, [], controller.signal)
|
||||
// 1. 清空旧片段(静默全局 toast:模板失效时由下方回退统一提示)
|
||||
await updateEditPlanClips(tid, [], controller.signal, true)
|
||||
// 2. 调用后端 from-assets 接口创建片段(异步秒级返回,60s 超时仅为兜底)
|
||||
await createClipsFromAssets(tid, ids, "main", requiredClipsCount, {
|
||||
signal: controller.signal,
|
||||
silentErrorToast: true,
|
||||
})
|
||||
// 3. 获取服务端生成的 clips(含 start_time/duration),供预览播放器使用
|
||||
const clipList = await getEditPlanClips(tid, { limit: 500 })
|
||||
@@ -146,6 +153,14 @@ export function useStep2Materials({
|
||||
message.error("智能选片失败,请重试")
|
||||
return
|
||||
}
|
||||
// 模板失效(404 模板不存在 / 400/422 无片段配置):
|
||||
// 清空失效选择并自动切到第一个有效模板 + toast,避免页面卡死无提示(#1777)
|
||||
if (isInvalidTemplateError(err)) {
|
||||
console.warn("[useStep2Materials] 当前模板已失效,触发自动回退:", err)
|
||||
onServerClipsChangeRef.current?.([])
|
||||
onTemplateInvalidRef.current?.()
|
||||
return
|
||||
}
|
||||
console.warn("[useStep2Materials] 写入 clips 失败:", err)
|
||||
}
|
||||
}, 800)
|
||||
|
||||
+1
-1
@@ -41,7 +41,7 @@ export function useVoiceUpload({ voiceLibrary, createLibMutation }: UseVoiceUplo
|
||||
}
|
||||
const libs = await queryClient.fetchQuery({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
queryFn: () => getAssetLibraries(),
|
||||
})
|
||||
lib = libs.find((l: AssetLibraryItem) => l.kind === "voice")
|
||||
if (!lib) throw new Error("无法创建配音库")
|
||||
|
||||
@@ -24,7 +24,7 @@ export function useVoiceMaterialData({ keyword, gender, tagIds }: UseVoiceMateri
|
||||
// ── 获取 voice 类型素材库 ─────────────────────────────────
|
||||
const { data: libraries = [] } = useQuery({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
queryFn: () => getAssetLibraries(),
|
||||
staleTime: 60_000,
|
||||
})
|
||||
|
||||
|
||||
@@ -26,7 +26,7 @@ export function useVoiceUpload({ showToast }: UseVoiceUploadProps) {
|
||||
/* 获取或创建默认配音库 */
|
||||
const libs = await queryClient.fetchQuery({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
queryFn: () => getAssetLibraries(),
|
||||
})
|
||||
const lib = libs.find((l) => l.kind === "voice")
|
||||
if (!lib) throw new Error("配音库不存在,请先在配音库页面创建")
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
/**
|
||||
* 失效模板判定/回退纯函数单测(#1777)
|
||||
*/
|
||||
import { describe, it, expect } from "vitest"
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
import {
|
||||
getHttpStatus,
|
||||
isInvalidTemplateError,
|
||||
isValidTemplate,
|
||||
findFirstValidTemplate,
|
||||
} from "@/pages/generate/hooks/useGenerateFormState/templateFallback"
|
||||
|
||||
function makeTemplate(partial: Partial<EditingTemplate> & { id: string }): EditingTemplate {
|
||||
return {
|
||||
name: partial.id,
|
||||
mode: "pip",
|
||||
category: "默认",
|
||||
tags: [],
|
||||
title_config: {
|
||||
ai_auto_select: false,
|
||||
content: "",
|
||||
font_preset: "",
|
||||
font_color: "",
|
||||
font_size: 28,
|
||||
position: "top",
|
||||
},
|
||||
subtitle_config: {
|
||||
enabled: true,
|
||||
position: "bottom",
|
||||
font: "",
|
||||
color: "",
|
||||
size: 20,
|
||||
animation: "",
|
||||
},
|
||||
bgm_config: { enabled: false, music_id: "" },
|
||||
segments: [{ segment_order: 0, material_type: null }],
|
||||
is_active: true,
|
||||
created_at: "",
|
||||
updated_at: "",
|
||||
...partial,
|
||||
} as EditingTemplate
|
||||
}
|
||||
|
||||
function axiosError(status: number, data?: unknown) {
|
||||
return { isAxiosError: true, response: { status, data } }
|
||||
}
|
||||
|
||||
describe("getHttpStatus", () => {
|
||||
it("提取 axios 错误的 HTTP 状态码", () => {
|
||||
expect(getHttpStatus(axiosError(404))).toBe(404)
|
||||
expect(getHttpStatus(axiosError(400))).toBe(400)
|
||||
})
|
||||
it("非 axios/无响应错误返回 null", () => {
|
||||
expect(getHttpStatus(new Error("network"))).toBeNull()
|
||||
expect(getHttpStatus(null)).toBeNull()
|
||||
expect(getHttpStatus(undefined)).toBeNull()
|
||||
expect(getHttpStatus({ isAxiosError: true })).toBeNull()
|
||||
})
|
||||
})
|
||||
|
||||
describe("isInvalidTemplateError", () => {
|
||||
it("404 始终判定为模板失效(模板不存在)", () => {
|
||||
expect(isInvalidTemplateError(axiosError(404))).toBe(true)
|
||||
expect(isInvalidTemplateError(axiosError(404, { detail: "Not Found" }))).toBe(true)
|
||||
})
|
||||
|
||||
it("400 且后端文案提到「片段配置」判定为模板无片段配置", () => {
|
||||
expect(
|
||||
isInvalidTemplateError(axiosError(400, { detail: "模板没有片段配置,无法创建片段" })),
|
||||
).toBe(true)
|
||||
})
|
||||
|
||||
it("400 但文案与片段配置无关 → 不误判", () => {
|
||||
expect(isInvalidTemplateError(axiosError(400, { detail: "素材参数错误" }))).toBe(false)
|
||||
})
|
||||
|
||||
it("422 命中片段/模板字段判定为失效", () => {
|
||||
expect(
|
||||
isInvalidTemplateError(
|
||||
axiosError(422, { detail: [{ loc: ["body", "segments"], msg: "field required" }] }),
|
||||
),
|
||||
).toBe(true)
|
||||
})
|
||||
|
||||
it("其他状态码(401/403/500/超时/网络)不判定为模板失效", () => {
|
||||
expect(isInvalidTemplateError(axiosError(401))).toBe(false)
|
||||
expect(isInvalidTemplateError(axiosError(403))).toBe(false)
|
||||
expect(isInvalidTemplateError(axiosError(500))).toBe(false)
|
||||
expect(isInvalidTemplateError({ code: "ECONNABORTED", message: "timeout of 60000ms" })).toBe(
|
||||
false,
|
||||
)
|
||||
expect(isInvalidTemplateError(new Error("Network Error"))).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe("isValidTemplate", () => {
|
||||
it("有片段且未被标记 inactive → 有效", () => {
|
||||
expect(isValidTemplate(makeTemplate({ id: "t1" }))).toBe(true)
|
||||
})
|
||||
it("segments 为空 → 无效(无片段配置)", () => {
|
||||
expect(isValidTemplate(makeTemplate({ id: "t2", segments: [] }))).toBe(false)
|
||||
})
|
||||
it("is_active=false → 无效(已停用/删除)", () => {
|
||||
expect(isValidTemplate(makeTemplate({ id: "t3", is_active: false }))).toBe(false)
|
||||
})
|
||||
it("is_active 字段缺失时视为有效(兼容旧后端)", () => {
|
||||
const t = makeTemplate({ id: "t4" })
|
||||
delete (t as Partial<EditingTemplate>).is_active
|
||||
expect(isValidTemplate(t)).toBe(true)
|
||||
})
|
||||
it("null/undefined → 无效", () => {
|
||||
expect(isValidTemplate(null)).toBe(false)
|
||||
expect(isValidTemplate(undefined)).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe("findFirstValidTemplate", () => {
|
||||
it("跳过无效模板,返回第一个有效模板", () => {
|
||||
const list = [
|
||||
makeTemplate({ id: "empty", segments: [] }),
|
||||
makeTemplate({ id: "inactive", is_active: false }),
|
||||
makeTemplate({ id: "valid1" }),
|
||||
makeTemplate({ id: "valid2" }),
|
||||
]
|
||||
expect(findFirstValidTemplate(list)?.id).toBe("valid1")
|
||||
})
|
||||
it("全部无效 → null(用于空状态引导)", () => {
|
||||
expect(
|
||||
findFirstValidTemplate([
|
||||
makeTemplate({ id: "a", segments: [] }),
|
||||
makeTemplate({ id: "b", is_active: false }),
|
||||
]),
|
||||
).toBeNull()
|
||||
})
|
||||
it("空数组/null → null", () => {
|
||||
expect(findFirstValidTemplate([])).toBeNull()
|
||||
expect(findFirstValidTemplate(null)).toBeNull()
|
||||
expect(findFirstValidTemplate(undefined)).toBeNull()
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,74 @@
|
||||
/**
|
||||
* useMaterialLibrary Hook 单测(#1777)
|
||||
* - Step2 视频库选择器只拉取 kind=video 的素材库,配音库(voice)/图片库(image) 不混入
|
||||
* - 自动选中第一个视频库
|
||||
*/
|
||||
import { describe, it, expect, vi, beforeEach } from "vitest"
|
||||
import { renderHook, waitFor } from "@testing-library/react"
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query"
|
||||
import type { ReactNode } from "react"
|
||||
import type { AssetItem, AssetLibraryItem } from "@/api/assets"
|
||||
|
||||
vi.mock("@/api/assets", () => ({
|
||||
getAssetLibraries: vi.fn(),
|
||||
getAssets: vi.fn(),
|
||||
isAssetUsable: vi.fn(() => true),
|
||||
}))
|
||||
|
||||
import { getAssetLibraries, getAssets } from "@/api/assets"
|
||||
import { useMaterialLibrary } from "@/pages/generate/hooks/step2-materials/useMaterialLibrary"
|
||||
|
||||
const mockGetLibraries = vi.mocked(getAssetLibraries)
|
||||
const mockGetAssets = vi.mocked(getAssets)
|
||||
|
||||
function lib(id: string, kind: AssetLibraryItem["kind"], name = id): AssetLibraryItem {
|
||||
return { id, name, kind }
|
||||
}
|
||||
|
||||
function createWrapper() {
|
||||
const queryClient = new QueryClient({
|
||||
defaultOptions: { queries: { retry: false, gcTime: 0 } },
|
||||
})
|
||||
return ({ children }: { children: ReactNode }) =>
|
||||
(<QueryClientProvider client={queryClient}>{children}</QueryClientProvider>) as ReactNode
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
mockGetAssets.mockResolvedValue({ items: [] as AssetItem[], total: 0 })
|
||||
})
|
||||
|
||||
describe("useMaterialLibrary (#1777 kind=video 过滤)", () => {
|
||||
it("按 kind=video 拉取素材库(后端参数过滤)", async () => {
|
||||
mockGetLibraries.mockResolvedValueOnce([lib("v1", "video")])
|
||||
renderHook(() => useMaterialLibrary(), { wrapper: createWrapper() })
|
||||
|
||||
await waitFor(() => expect(mockGetLibraries).toHaveBeenCalledTimes(1))
|
||||
expect(mockGetLibraries).toHaveBeenCalledWith("video")
|
||||
})
|
||||
|
||||
it("下拉库列表只包含视频库(自动选中第一个视频库)", async () => {
|
||||
mockGetLibraries.mockResolvedValueOnce([
|
||||
lib("voice-1", "voice"),
|
||||
lib("img-1", "image"),
|
||||
lib("video-1", "video"),
|
||||
lib("video-2", "video"),
|
||||
])
|
||||
const { result } = renderHook(() => useMaterialLibrary(), { wrapper: createWrapper() })
|
||||
|
||||
await waitFor(() => expect(result.current.libraries).toHaveLength(2))
|
||||
expect(result.current.libraries.map((l) => l.id)).toEqual(["video-1", "video-2"])
|
||||
expect(result.current.libraries.every((l) => l.kind === "video")).toBe(true)
|
||||
// 自动选中第一个视频库
|
||||
expect(result.current.selectedLibraryId).toBe("video-1")
|
||||
})
|
||||
|
||||
it("没有视频库时库列表为空且不自动选中(UI 展示空状态)", async () => {
|
||||
mockGetLibraries.mockResolvedValueOnce([lib("voice-1", "voice"), lib("img-1", "image")])
|
||||
const { result } = renderHook(() => useMaterialLibrary(), { wrapper: createWrapper() })
|
||||
|
||||
await waitFor(() => expect(mockGetLibraries).toHaveBeenCalled())
|
||||
expect(result.current.libraries).toEqual([])
|
||||
expect(result.current.selectedLibraryId).toBe("")
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,174 @@
|
||||
/**
|
||||
* useTemplateSelection Hook 单测(#1777)
|
||||
* - 自动选择跳过无片段/inactive 模板,只选第一个有效模板
|
||||
* - 传 validOnly=true 给后端
|
||||
* - 用户手动选择优先,自动逻辑不覆盖
|
||||
* - handleInvalidTemplate:失效时自动切到第一个有效模板 + toast;无有效模板时清空
|
||||
* - selectedTemplate 仅内存态,不写入 localStorage/sessionStorage
|
||||
*/
|
||||
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"
|
||||
import { renderHook, waitFor, act } from "@testing-library/react"
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query"
|
||||
import type { ReactNode } from "react"
|
||||
|
||||
// antd message mock(拦截 toast)——vi.hoisted 保证 mock 工厂可引用
|
||||
const { messageMock } = vi.hoisted(() => ({
|
||||
messageMock: {
|
||||
warning: vi.fn(),
|
||||
error: vi.fn(),
|
||||
success: vi.fn(),
|
||||
info: vi.fn(),
|
||||
loading: vi.fn(() => vi.fn()),
|
||||
},
|
||||
}))
|
||||
vi.mock("antd", () => ({ message: messageMock }))
|
||||
|
||||
vi.mock("@/api/editing-planner", () => ({
|
||||
getEditingTemplates: vi.fn(),
|
||||
}))
|
||||
|
||||
import { getEditingTemplates } from "@/api/editing-planner"
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
import { useTemplateSelection } from "@/pages/generate/hooks/useGenerateFormState/useTemplateSelection"
|
||||
|
||||
const mockGetTemplates = vi.mocked(getEditingTemplates)
|
||||
|
||||
function tpl(id: string, partial: Partial<EditingTemplate> = {}): EditingTemplate {
|
||||
return {
|
||||
id,
|
||||
name: id,
|
||||
mode: "pip",
|
||||
category: "默认",
|
||||
tags: [],
|
||||
title_config: {
|
||||
ai_auto_select: false,
|
||||
content: "",
|
||||
font_preset: "",
|
||||
font_color: "",
|
||||
font_size: 28,
|
||||
position: "top",
|
||||
},
|
||||
subtitle_config: {
|
||||
enabled: true,
|
||||
position: "bottom",
|
||||
font: "",
|
||||
color: "",
|
||||
size: 20,
|
||||
animation: "",
|
||||
},
|
||||
bgm_config: { enabled: false, music_id: "" },
|
||||
segments: [{ segment_order: 0, material_type: null }],
|
||||
is_active: true,
|
||||
created_at: "",
|
||||
updated_at: "",
|
||||
...partial,
|
||||
} as EditingTemplate
|
||||
}
|
||||
|
||||
function createWrapper() {
|
||||
const queryClient = new QueryClient({
|
||||
defaultOptions: { queries: { retry: false, gcTime: 0 } },
|
||||
})
|
||||
return ({ children }: { children: ReactNode }) =>
|
||||
(<QueryClientProvider client={queryClient}>{children}</QueryClientProvider>) as ReactNode
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
localStorage.clear()
|
||||
sessionStorage.clear()
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
localStorage.clear()
|
||||
sessionStorage.clear()
|
||||
})
|
||||
|
||||
describe("useTemplateSelection (#1777)", () => {
|
||||
it("请求模板时传 validOnly=true,并自动选中第一个有片段的有效模板", async () => {
|
||||
mockGetTemplates.mockResolvedValueOnce([
|
||||
tpl("empty", { segments: [] }),
|
||||
tpl("inactive", { is_active: false }),
|
||||
tpl("valid-a"),
|
||||
tpl("valid-b"),
|
||||
])
|
||||
const { result } = renderHook(() => useTemplateSelection(), { wrapper: createWrapper() })
|
||||
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe("valid-a"))
|
||||
expect(mockGetTemplates).toHaveBeenCalledWith({ validOnly: true })
|
||||
// 暴露给 UI 的 userTemplates 已过滤掉无效模板
|
||||
expect(result.current.userTemplates.map((t) => t.id)).toEqual(["valid-a", "valid-b"])
|
||||
})
|
||||
|
||||
it("列表全部无效时 selectedTemplate 为空(交空状态引导),不选中失效模板", async () => {
|
||||
mockGetTemplates.mockResolvedValueOnce([
|
||||
tpl("empty", { segments: [] }),
|
||||
tpl("inactive", { is_active: false }),
|
||||
])
|
||||
const { result } = renderHook(() => useTemplateSelection(), { wrapper: createWrapper() })
|
||||
await waitFor(() => expect(mockGetTemplates).toHaveBeenCalled())
|
||||
// 给 effect 一个 tick
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe(""))
|
||||
expect(result.current.userTemplates).toHaveLength(0)
|
||||
})
|
||||
|
||||
it("用户手动选择优先:自动逻辑不会覆盖手动选择", async () => {
|
||||
mockGetTemplates.mockResolvedValueOnce([tpl("a"), tpl("b")])
|
||||
const { result } = renderHook(() => useTemplateSelection(), { wrapper: createWrapper() })
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe("a"))
|
||||
|
||||
act(() => result.current.setSelectedTemplate("b"))
|
||||
expect(result.current.selectedTemplate).toBe("b")
|
||||
|
||||
// 重新渲染 / refetch 后仍保持用户的手动选择
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe("b"))
|
||||
})
|
||||
|
||||
it("handleInvalidTemplate:当前模板失效时自动切到第一个有效模板并 toast", async () => {
|
||||
mockGetTemplates.mockResolvedValueOnce([tpl("bad", { segments: [] }), tpl("good")])
|
||||
const { result } = renderHook(() => useTemplateSelection(), { wrapper: createWrapper() })
|
||||
// 自动选中有效模板 good(bad 无片段不会被自动选中)
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe("good"))
|
||||
messageMock.warning.mockClear()
|
||||
|
||||
// 模拟运行时用户停留在一个已失效的模板 id(外部/草稿态),触发回退
|
||||
act(() => result.current.setSelectedTemplate("stale-id"))
|
||||
expect(result.current.selectedTemplate).toBe("stale-id")
|
||||
|
||||
act(() => {
|
||||
const handled = result.current.handleInvalidTemplate()
|
||||
expect(handled).toBe(true)
|
||||
})
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe("good"))
|
||||
expect(messageMock.warning).toHaveBeenCalledWith("原模板已失效,已自动切换")
|
||||
})
|
||||
|
||||
it("handleInvalidTemplate:无有效模板时清空选择并提示去创建", async () => {
|
||||
mockGetTemplates.mockResolvedValueOnce([tpl("bad", { segments: [] })])
|
||||
const { result } = renderHook(() => useTemplateSelection(), { wrapper: createWrapper() })
|
||||
await waitFor(() => expect(result.current.userTemplates).toHaveLength(0))
|
||||
|
||||
act(() => result.current.setSelectedTemplate("stale-id"))
|
||||
act(() => {
|
||||
result.current.handleInvalidTemplate()
|
||||
})
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe(""))
|
||||
expect(messageMock.warning).toHaveBeenCalledWith(expect.stringContaining("没有可用模板"))
|
||||
})
|
||||
|
||||
it("失效模板 ID 不写入任何持久化存储", async () => {
|
||||
mockGetTemplates.mockResolvedValueOnce([tpl("good")])
|
||||
const { result } = renderHook(() => useTemplateSelection(), { wrapper: createWrapper() })
|
||||
await waitFor(() => expect(result.current.selectedTemplate).toBe("good"))
|
||||
|
||||
act(() => result.current.setSelectedTemplate("stale-invalid-id"))
|
||||
act(() => result.current.handleInvalidTemplate())
|
||||
|
||||
const ls = JSON.stringify(localStorage)
|
||||
const ss = JSON.stringify(sessionStorage)
|
||||
expect(ls).not.toContain("stale-invalid-id")
|
||||
expect(ss).not.toContain("stale-invalid-id")
|
||||
// URL 也不含
|
||||
expect(window.location.href).not.toContain("stale-invalid-id")
|
||||
})
|
||||
})
|
||||
@@ -39,6 +39,8 @@ class BGMConfig:
|
||||
sidechain_attack: float = 0.02 # 攻击时间
|
||||
sidechain_release: float = 0.5 # 释放时间
|
||||
sidechain_threshold: float = -25.0 # 触发阈值(dB)
|
||||
audio_offset: float = 0.0 # BGM 段落起始偏移(秒),#1767 策略二
|
||||
volume_adjust_db: float = 0.0 # 音量微调 dB(-3~+3),#1767 策略三
|
||||
|
||||
@classmethod
|
||||
def from_config_dict(cls, bgm_path: str, config: dict) -> "BGMConfig":
|
||||
@@ -54,6 +56,8 @@ class BGMConfig:
|
||||
sidechain_attack=float(config.get("sidechain_attack", 0.02)),
|
||||
sidechain_release=float(config.get("sidechain_release", 0.5)),
|
||||
sidechain_threshold=float(config.get("sidechain_threshold", -25.0)),
|
||||
audio_offset=float(config.get("audio_offset", 0.0)),
|
||||
volume_adjust_db=float(config.get("volume_adjust_db", 0.0)),
|
||||
)
|
||||
|
||||
|
||||
@@ -84,14 +88,18 @@ def prepare_bgm_track(
|
||||
target_duration = 5.0 # 兜底
|
||||
|
||||
bgm_dur = probe_duration(bgm.bgm_path)
|
||||
needs_loop = bgm.loop_enabled and bgm_dur > 0 and bgm_dur < target_duration * 0.9
|
||||
# #1767:seek 后有效时长 = 总时长 - 偏移
|
||||
effective_dur = (
|
||||
max(1.0, bgm_dur - bgm.audio_offset) if bgm.audio_offset > 0 and bgm_dur > bgm.audio_offset else bgm_dur
|
||||
)
|
||||
needs_loop = bgm.loop_enabled and effective_dur > 0 and effective_dur < target_duration * 0.9
|
||||
|
||||
# 构建滤镜链
|
||||
filter_parts: list[str] = []
|
||||
|
||||
if needs_loop:
|
||||
# 计算需要循环多少次才能铺满
|
||||
loop_count = max(1, int(target_duration / bgm_dur) + 2)
|
||||
# 计算需要循环多少次才能铺满(基于 seek 后有效时长)
|
||||
loop_count = max(1, int(target_duration / effective_dur) + 2)
|
||||
# aloop 滤镜:循环指定次数
|
||||
filter_parts.append(f"aloop=loop={loop_count}:size=0")
|
||||
|
||||
@@ -113,11 +121,30 @@ def prepare_bgm_track(
|
||||
filter_parts.append(f"atrim=0:{target_duration:.3f}")
|
||||
filter_parts.append("asetpts=N/SR/TB") # 重置时间戳
|
||||
|
||||
filter_str = ",".join(filter_parts)
|
||||
# #1767:BGM 段落差异化 — 使用 -ss 从偏移位置开始(seek 效率高,不读跳过部分)
|
||||
seek_args: list[str] = []
|
||||
if bgm.audio_offset > 0 and bgm_dur > bgm.audio_offset:
|
||||
seek_args = ["-ss", f"{bgm.audio_offset:.3f}"]
|
||||
logger.info("[bgm] #1767 audio_offset=%.1fs(段落差异化)", bgm.audio_offset)
|
||||
|
||||
# #1767:音量微调 — dB 转线性系数(10^(dB/20))
|
||||
db_adjust_filter = ""
|
||||
if abs(bgm.volume_adjust_db) > 0.01:
|
||||
linear_factor = 10.0 ** (bgm.volume_adjust_db / 20.0)
|
||||
db_adjust_filter = f",volume={linear_factor:.4f}"
|
||||
logger.info("[bgm] #1767 volume_adjust=%.0fdB → linear=%.4f", bgm.volume_adjust_db, linear_factor)
|
||||
|
||||
# #1767:追加 dB 微调到滤镜链末尾
|
||||
if db_adjust_filter:
|
||||
filter_str_base = ",".join(filter_parts)
|
||||
filter_str = filter_str_base + db_adjust_filter
|
||||
else:
|
||||
filter_str = ",".join(filter_parts)
|
||||
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
*seek_args,
|
||||
"-i",
|
||||
bgm.bgm_path,
|
||||
"-filter:a",
|
||||
|
||||
@@ -59,13 +59,18 @@ def build_xfade_filter_chain(
|
||||
transitions: list[str],
|
||||
*,
|
||||
transition_duration: float = DEFAULT_TRANSITION_DURATION,
|
||||
transition_durations: list[float] | None = None,
|
||||
jitters: list[float] | None = None,
|
||||
output_label: str = "outv",
|
||||
) -> tuple[str, float]:
|
||||
"""构建 xfade 转场滤镜链(re-export,#1766 增加逐转场时长与 jitter 支持)."""
|
||||
return _build_xfade_filter_chain_base(
|
||||
clip_durations,
|
||||
clip_video_labels,
|
||||
transitions,
|
||||
transition_duration=transition_duration,
|
||||
transition_durations=transition_durations,
|
||||
jitters=jitters,
|
||||
output_label=output_label,
|
||||
)
|
||||
|
||||
|
||||
@@ -105,17 +105,24 @@ class TransitionEngine:
|
||||
transitions: list[str],
|
||||
*,
|
||||
transition_duration: float | None = None,
|
||||
transition_durations: list[float] | None = None,
|
||||
jitters: list[float] | None = None,
|
||||
output_label: str = "outv",
|
||||
) -> tuple[str, float]:
|
||||
"""构建 xfade 转场滤镜链.
|
||||
|
||||
对每步转场应用验证和降级,然后调用底层 ffmpeg_utils 构建。
|
||||
|
||||
#1766 增强:支持逐转场独立时长(transition_durations)和位置微调(jitters)。
|
||||
|
||||
Args:
|
||||
clip_durations: 每个片段的时长
|
||||
clip_video_labels: 每个片段的视频流标签
|
||||
transitions: 每个片段对应的转场效果
|
||||
transition_duration: 统一转场时长,None 则使用引擎默认值
|
||||
transition_duration: 全局默认转场时长,None 则使用引擎默认值
|
||||
transition_durations: #1766 逐转场时长列表,与 transitions 等长;
|
||||
None 时使用各 resolved config 的 duration
|
||||
jitters: #1766 逐转场位置偏移列表(秒)
|
||||
output_label: 最终输出标签
|
||||
|
||||
Returns:
|
||||
@@ -127,6 +134,8 @@ class TransitionEngine:
|
||||
clip_video_labels=clip_video_labels,
|
||||
transitions=transitions,
|
||||
transition_duration=transition_duration or self._default_duration,
|
||||
transition_durations=transition_durations,
|
||||
jitters=jitters,
|
||||
output_label=output_label,
|
||||
)
|
||||
|
||||
@@ -134,17 +143,20 @@ class TransitionEngine:
|
||||
resolved = self.resolve_clip_transitions(transitions, clip_durations)
|
||||
resolved_effects = [c.effect for c in resolved]
|
||||
|
||||
# 使用统一的时长(取各转场中最大的时长作为基准,底层会做每步钳制)
|
||||
dur = transition_duration or self._default_duration
|
||||
if not dur:
|
||||
dur = max(c.duration for c in resolved) if resolved else DEFAULT_TRANSITION_DURATION
|
||||
# #1766: 逐转场时长(优先使用传入的 transition_durations,否则用 resolved config)
|
||||
if transition_durations is not None:
|
||||
resolved_durations = list(transition_durations)
|
||||
else:
|
||||
resolved_durations = [c.duration for c in resolved]
|
||||
|
||||
# 调用底层构建
|
||||
return build_xfade_filter_chain(
|
||||
clip_durations=clip_durations,
|
||||
clip_video_labels=clip_video_labels,
|
||||
transitions=resolved_effects,
|
||||
transition_duration=dur,
|
||||
transition_duration=transition_duration or self._default_duration,
|
||||
transition_durations=resolved_durations,
|
||||
jitters=jitters,
|
||||
output_label=output_label,
|
||||
)
|
||||
|
||||
|
||||
@@ -1870,16 +1870,27 @@ class UnifiedRenderService:
|
||||
)
|
||||
else:
|
||||
# 有转场效果:用 TransitionEngine 构建 xfade 链
|
||||
layer_dur = 0.0
|
||||
for d in layer_transition_durations:
|
||||
if d > 0:
|
||||
layer_dur = d
|
||||
break
|
||||
# #1766: 提取逐转场时长(每个 clip 的 transition_duration,跳过第一个)
|
||||
# layer_transition_durations[i] 对应 clip i 的转场,第 0 个忽略
|
||||
per_transition_durations = [d for idx, d in enumerate(layer_transition_durations) if idx > 0]
|
||||
# #1766: 提取每个 clip 的 jitter(存在 config 中),跳过第一个
|
||||
layer_jitters = [
|
||||
(
|
||||
all_clips[layer_clip_indices[idx]].config.get("transition_jitter", 0.0)
|
||||
if isinstance(all_clips[layer_clip_indices[idx]].config, dict)
|
||||
else 0.0
|
||||
)
|
||||
for idx in range(len(layer_clip_indices))
|
||||
]
|
||||
per_transition_jitters = [j for idx, j in enumerate(layer_jitters) if idx > 0]
|
||||
xfade_filter, xfade_estimated_dur = self._transition_engine.build_xfade_chain(
|
||||
clip_durations=layer_durations,
|
||||
clip_video_labels=layer_labels,
|
||||
transitions=layer_transitions,
|
||||
transition_duration=layer_dur if layer_dur > 0 else None,
|
||||
transition_durations=(
|
||||
per_transition_durations if any(d > 0 for d in per_transition_durations) else None
|
||||
),
|
||||
jitters=per_transition_jitters if any(j != 0.0 for j in per_transition_jitters) else None,
|
||||
output_label=out_label,
|
||||
)
|
||||
if xfade_filter:
|
||||
|
||||
@@ -2126,6 +2126,14 @@
|
||||
"type": "JSON",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "is_default",
|
||||
"nullable": false,
|
||||
"primary_key": false,
|
||||
"type": "BOOLEAN",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "created_at",
|
||||
@@ -2142,6 +2150,13 @@
|
||||
],
|
||||
"name": "ix_projects_owner_user_id",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"columns": [
|
||||
"owner_user_id"
|
||||
],
|
||||
"name": "uq_projects_owner_default",
|
||||
"unique": true
|
||||
}
|
||||
],
|
||||
"primary_key": [
|
||||
@@ -3570,4 +3585,4 @@
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -35,14 +35,37 @@ class InMemoryAssetLibraryRepository:
|
||||
return True
|
||||
return False
|
||||
|
||||
def increment_asset_count(self, library_id: str, size_delta: int) -> None:
|
||||
def increment_asset_count(self, library_id: str, count_delta: int = 1, size_delta: int = 0) -> None:
|
||||
library = self._libraries.get(library_id)
|
||||
if library:
|
||||
library.asset_count += 1
|
||||
library.asset_count += count_delta
|
||||
library.total_size += size_delta
|
||||
|
||||
def decrement_asset_count(self, library_id: str, size_delta: int) -> None:
|
||||
def decrement_asset_count(self, library_id: str, count_delta: int = 1, size_delta: int = 0) -> None:
|
||||
library = self._libraries.get(library_id)
|
||||
if library:
|
||||
library.asset_count = max(0, library.asset_count - 1)
|
||||
library.asset_count = max(0, library.asset_count - count_delta)
|
||||
library.total_size = max(0, library.total_size - size_delta)
|
||||
|
||||
def recount_assets(self, library_id: str) -> int:
|
||||
"""InMemory 实现无法真正重算(没有 asset 数据源),返回当前计数。"""
|
||||
library = self._libraries.get(library_id)
|
||||
return library.asset_count if library else 0
|
||||
|
||||
def get_or_create_default_library(
|
||||
self,
|
||||
project_id: str,
|
||||
kind: AssetLibraryKind,
|
||||
*,
|
||||
name: str | None = None,
|
||||
) -> AssetLibrary:
|
||||
"""幂等获取/创建默认素材库(Issue #1775,内存实现,模拟唯一约束语义)。"""
|
||||
for lib in self._libraries.values():
|
||||
if lib.project_id == project_id and lib.kind == kind:
|
||||
return lib
|
||||
# 回退到 find_by_project
|
||||
for lib in self.find_by_project(project_id, kind):
|
||||
return lib
|
||||
library_name = name or f"{kind.value}素材库"
|
||||
library = AssetLibrary.create(project_id=project_id, name=library_name, kind=kind)
|
||||
return self.create(library)
|
||||
|
||||
@@ -29,3 +29,30 @@ class InMemoryProjectRepository:
|
||||
del self._items[project_id]
|
||||
return True
|
||||
return False
|
||||
|
||||
def find_default_by_owner(self, owner_user_id: str) -> Project | None:
|
||||
"""查找用户的默认项目(Issue #1775 幂等接口,内存实现)。"""
|
||||
for p in self._items.values():
|
||||
if p.owner_user_id == owner_user_id and getattr(p, "is_default", False):
|
||||
return p
|
||||
return None
|
||||
|
||||
def get_or_create_default_project(
|
||||
self,
|
||||
owner_user_id: str,
|
||||
*,
|
||||
name: str = "默认项目",
|
||||
description: str = "小程序自动创建的默认项目",
|
||||
) -> Project:
|
||||
"""幂等获取/创建默认项目(内存实现,模拟 DB 部分唯一索引语义)。"""
|
||||
existing = self.find_default_by_owner(owner_user_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
project = Project.create(
|
||||
owner_user_id=owner_user_id,
|
||||
name=name,
|
||||
description=description,
|
||||
is_default=True,
|
||||
)
|
||||
self._items[project.id] = project
|
||||
return project
|
||||
|
||||
@@ -77,16 +77,148 @@ class SQLAlchemyAssetLibraryRepository:
|
||||
return True
|
||||
return False
|
||||
|
||||
async def increment_asset_count(self, library_id: str, size_delta: int) -> None:
|
||||
model = self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).first()
|
||||
if model:
|
||||
model.asset_count = (model.asset_count or 0) + 1
|
||||
model.total_size = (model.total_size or 0) + size_delta
|
||||
self.session.commit()
|
||||
def increment_asset_count(self, library_id: str, count_delta: int = 1, size_delta: int = 0) -> None:
|
||||
"""原子递增素材计数(Issue #1776)。
|
||||
|
||||
async def decrement_asset_count(self, library_id: str, size_delta: int) -> None:
|
||||
model = self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).first()
|
||||
if model:
|
||||
model.asset_count = max(0, (model.asset_count or 0) - 1)
|
||||
model.total_size = max(0, (model.total_size or 0) - size_delta)
|
||||
使用 SQL 级 UPDATE 保证并发安全,不单独 commit(由调用方统一事务提交)。
|
||||
"""
|
||||
from sqlalchemy import func
|
||||
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).update(
|
||||
{
|
||||
AssetLibraryModel.asset_count: func.coalesce(AssetLibraryModel.asset_count, 0) + count_delta,
|
||||
AssetLibraryModel.total_size: func.coalesce(AssetLibraryModel.total_size, 0) + size_delta,
|
||||
}
|
||||
)
|
||||
|
||||
def decrement_asset_count(self, library_id: str, count_delta: int = 1, size_delta: int = 0) -> None:
|
||||
"""原子递减素材计数(Issue #1776),下限为 0 防止负数。
|
||||
|
||||
使用 SQL 级 UPDATE 保证并发安全,不单独 commit(由调用方统一事务提交)。
|
||||
使用 CASE WHEN 兼容 SQLite(测试)和 PostgreSQL(生产)。
|
||||
"""
|
||||
from sqlalchemy import case, func
|
||||
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).update(
|
||||
{
|
||||
AssetLibraryModel.asset_count: case(
|
||||
(func.coalesce(AssetLibraryModel.asset_count, 0) - count_delta < 0, 0),
|
||||
else_=func.coalesce(AssetLibraryModel.asset_count, 0) - count_delta,
|
||||
),
|
||||
AssetLibraryModel.total_size: case(
|
||||
(func.coalesce(AssetLibraryModel.total_size, 0) - size_delta < 0, 0),
|
||||
else_=func.coalesce(AssetLibraryModel.total_size, 0) - size_delta,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
def recount_assets(self, library_id: str) -> int:
|
||||
"""重算素材库计数(Issue #1776)。
|
||||
|
||||
直接查询实际素材数量(排除已删除),更新 asset_count 和 total_size。
|
||||
返回重算后的实际计数。
|
||||
"""
|
||||
from sqlalchemy import func
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetModel
|
||||
|
||||
# 查询实际计数(排除 deleted)
|
||||
actual_count = (
|
||||
self.session.query(func.count(AssetModel.id))
|
||||
.filter(
|
||||
AssetModel.asset_library_id == library_id,
|
||||
AssetModel.status != "deleted",
|
||||
)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
# 查询实际总大小
|
||||
actual_size = (
|
||||
self.session.query(func.coalesce(func.sum(AssetModel.file_size), 0))
|
||||
.filter(
|
||||
AssetModel.asset_library_id == library_id,
|
||||
AssetModel.status != "deleted",
|
||||
)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
# 更新素材库记录
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).update(
|
||||
{
|
||||
AssetLibraryModel.asset_count: actual_count,
|
||||
AssetLibraryModel.total_size: actual_size,
|
||||
}
|
||||
)
|
||||
return actual_count
|
||||
|
||||
def get_or_create_default_library(
|
||||
self,
|
||||
project_id: str,
|
||||
kind: AssetLibraryKind,
|
||||
*,
|
||||
name: str | None = None,
|
||||
) -> AssetLibrary:
|
||||
"""幂等获取/创建项目下指定 kind 的默认素材库(Issue #1775)。
|
||||
|
||||
依赖唯一约束 uq_asset_libraries_project_kind(project_id, kind):
|
||||
并发创建只有一个成功,其余 IntegrityError 后回滚重查,
|
||||
保证同一项目同 kind 永远只有一个素材库。
|
||||
"""
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
default_names = {
|
||||
AssetLibraryKind.VIDEO: "视频素材库",
|
||||
AssetLibraryKind.VOICE: "配音素材库",
|
||||
AssetLibraryKind.IMAGE: "图片素材库",
|
||||
}
|
||||
library_name = name or default_names.get(kind, f"{kind.value}素材库")
|
||||
|
||||
# 快速路径
|
||||
existing = (
|
||||
self.session.query(AssetLibraryModel)
|
||||
.filter(AssetLibraryModel.project_id == project_id, AssetLibraryModel.kind == kind.value)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
return self._to_entity(existing)
|
||||
|
||||
library = AssetLibrary.create(project_id=project_id, name=library_name, kind=kind)
|
||||
model = AssetLibraryModel(
|
||||
id=library.id,
|
||||
project_id=library.project_id,
|
||||
name=library.name,
|
||||
kind=library.kind.value,
|
||||
asset_count=0,
|
||||
total_size=0,
|
||||
created_at=library.created_at,
|
||||
updated_at=library.updated_at,
|
||||
)
|
||||
try:
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
return library
|
||||
except IntegrityError:
|
||||
self.session.rollback()
|
||||
existing = (
|
||||
self.session.query(AssetLibraryModel)
|
||||
.filter(
|
||||
AssetLibraryModel.project_id == project_id,
|
||||
AssetLibraryModel.kind == kind.value,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
return self._to_entity(existing)
|
||||
raise
|
||||
|
||||
def _to_entity(self, model: AssetLibraryModel) -> AssetLibrary:
|
||||
return AssetLibrary(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
name=model.name,
|
||||
kind=AssetLibraryKind(model.kind),
|
||||
asset_count=int(model.asset_count or 0),
|
||||
total_size=int(model.total_size or 0),
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
@@ -3,7 +3,7 @@ from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetModel, AssetTagModel
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetLibraryModel, AssetModel, AssetTagModel
|
||||
from packages.domain import Asset, AssetStatus, ClassificationStatus
|
||||
|
||||
|
||||
@@ -141,6 +141,16 @@ class SQLAlchemyAssetRepository:
|
||||
self.session.add(model)
|
||||
self.session.flush()
|
||||
self._sync_asset_tags(asset.id, asset.tag_ids)
|
||||
# Issue #1776: 自动维护素材库计数(同事务内原子更新)
|
||||
if asset.library_id and asset.status.value != "deleted":
|
||||
from sqlalchemy import func
|
||||
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == asset.library_id).update(
|
||||
{
|
||||
AssetLibraryModel.asset_count: func.coalesce(AssetLibraryModel.asset_count, 0) + 1,
|
||||
AssetLibraryModel.total_size: func.coalesce(AssetLibraryModel.total_size, 0) + asset.file_size,
|
||||
}
|
||||
)
|
||||
self.session.commit()
|
||||
return asset
|
||||
|
||||
@@ -175,7 +185,27 @@ class SQLAlchemyAssetRepository:
|
||||
def delete(self, asset_id: str) -> bool:
|
||||
model = self.session.query(AssetModel).filter(AssetModel.id == asset_id).first()
|
||||
if model:
|
||||
library_id = model.asset_library_id
|
||||
file_size = model.file_size or 0
|
||||
# 只统计非 deleted 状态的素材
|
||||
was_counted = model.status != "deleted"
|
||||
self.session.delete(model)
|
||||
# Issue #1776: 自动维护素材库计数
|
||||
if library_id and was_counted:
|
||||
from sqlalchemy import case, func
|
||||
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).update(
|
||||
{
|
||||
AssetLibraryModel.asset_count: case(
|
||||
(func.coalesce(AssetLibraryModel.asset_count, 0) - 1 < 0, 0),
|
||||
else_=func.coalesce(AssetLibraryModel.asset_count, 0) - 1,
|
||||
),
|
||||
AssetLibraryModel.total_size: case(
|
||||
(func.coalesce(AssetLibraryModel.total_size, 0) - file_size < 0, 0),
|
||||
else_=func.coalesce(AssetLibraryModel.total_size, 0) - file_size,
|
||||
),
|
||||
}
|
||||
)
|
||||
self.session.commit()
|
||||
return True
|
||||
return False
|
||||
@@ -187,11 +217,44 @@ class SQLAlchemyAssetRepository:
|
||||
from datetime import datetime, timezone
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
# 先查询待删除素材的库分布(用于更新计数)
|
||||
to_delete = (
|
||||
self.session.query(AssetModel.asset_library_id, AssetModel.file_size)
|
||||
.filter(AssetModel.id.in_(asset_ids), AssetModel.status != "deleted")
|
||||
.all()
|
||||
)
|
||||
if not to_delete:
|
||||
return 0
|
||||
# 按库分组统计
|
||||
library_deltas: dict[str, tuple[int, int]] = {} # library_id -> (count_delta, size_delta)
|
||||
for lib_id, size in to_delete:
|
||||
if lib_id not in library_deltas:
|
||||
library_deltas[lib_id] = (0, 0)
|
||||
c, s = library_deltas[lib_id]
|
||||
library_deltas[lib_id] = (c + 1, s + (size or 0))
|
||||
# 执行软删除
|
||||
count = (
|
||||
self.session.query(AssetModel)
|
||||
.filter(AssetModel.id.in_(asset_ids), AssetModel.status != "deleted")
|
||||
.update({AssetModel.status: "deleted", AssetModel.updated_at: now}, synchronize_session=False)
|
||||
)
|
||||
# Issue #1776: 自动维护各素材库计数
|
||||
if library_deltas:
|
||||
from sqlalchemy import case, func
|
||||
|
||||
for lib_id, (count_delta, size_delta) in library_deltas.items():
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == lib_id).update(
|
||||
{
|
||||
AssetLibraryModel.asset_count: case(
|
||||
(func.coalesce(AssetLibraryModel.asset_count, 0) - count_delta < 0, 0),
|
||||
else_=func.coalesce(AssetLibraryModel.asset_count, 0) - count_delta,
|
||||
),
|
||||
AssetLibraryModel.total_size: case(
|
||||
(func.coalesce(AssetLibraryModel.total_size, 0) - size_delta < 0, 0),
|
||||
else_=func.coalesce(AssetLibraryModel.total_size, 0) - size_delta,
|
||||
),
|
||||
}
|
||||
)
|
||||
self.session.commit()
|
||||
return count
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Integer, String, Text, UniqueConstraint
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Index, Integer, String, Text, UniqueConstraint, text
|
||||
from sqlalchemy.orm import declarative_base
|
||||
|
||||
Base: Any = declarative_base()
|
||||
@@ -44,12 +44,18 @@ class UserModel(Base):
|
||||
|
||||
class ProjectModel(Base):
|
||||
__tablename__ = "projects"
|
||||
__table_args__ = (
|
||||
# Issue #1775: 每个用户至多一个默认项目(部分唯一索引,只约束 is_default=true 的行)。
|
||||
# 注意:不加 UniqueConstraint(那会要求全表唯一),用部分索引表达"每用户一个默认项目"。
|
||||
Index("uq_projects_owner_default", "owner_user_id", unique=True, postgresql_where=text("is_default = true")),
|
||||
)
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
owner_user_id = Column(String(36), nullable=False, index=True)
|
||||
name = Column(String(100), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
shared_users = Column(JSON, nullable=False, default=list) # 被共享的用户 ID 列表
|
||||
is_default = Column(Boolean, nullable=False, default=False, server_default="false")
|
||||
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ class SQLAlchemyProjectRepository:
|
||||
name=model.name,
|
||||
description=model.description,
|
||||
shared_users=model.shared_users or [],
|
||||
is_default=bool(getattr(model, "is_default", False)),
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
@@ -33,9 +34,12 @@ class SQLAlchemyProjectRepository:
|
||||
name=project.name,
|
||||
description=project.description,
|
||||
shared_users=project.shared_users,
|
||||
is_default=project.is_default,
|
||||
created_at=project.created_at,
|
||||
)
|
||||
self.session.add(model)
|
||||
if existing:
|
||||
existing.is_default = project.is_default
|
||||
self.session.commit()
|
||||
return project
|
||||
|
||||
@@ -76,3 +80,59 @@ class SQLAlchemyProjectRepository:
|
||||
self.session.delete(model)
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
def find_default_by_owner(self, owner_user_id: str) -> Project | None:
|
||||
"""查找用户的默认项目(is_default=true)。"""
|
||||
model = (
|
||||
self.session.query(ProjectModel)
|
||||
.filter(ProjectModel.owner_user_id == owner_user_id, ProjectModel.is_default.is_(True))
|
||||
.first()
|
||||
)
|
||||
return self._to_entity(model) if model else None
|
||||
|
||||
def get_or_create_default_project(
|
||||
self,
|
||||
owner_user_id: str,
|
||||
*,
|
||||
name: str = "默认项目",
|
||||
description: str = "小程序自动创建的默认项目",
|
||||
) -> Project:
|
||||
"""幂等获取/创建用户的默认项目(Issue #1775)。
|
||||
|
||||
依赖部分唯一索引 uq_projects_owner_default(每用户至多一条 is_default=true):
|
||||
并发创建时只有一个 INSERT 成功,其余触发 IntegrityError 后回滚重查,
|
||||
保证同一用户永远只有一个默认项目。
|
||||
"""
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
# 快速路径:已有默认项目
|
||||
existing = self.find_default_by_owner(owner_user_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
|
||||
project = Project.create(
|
||||
owner_user_id=owner_user_id,
|
||||
name=name,
|
||||
description=description,
|
||||
is_default=True,
|
||||
)
|
||||
model = ProjectModel(
|
||||
id=project.id,
|
||||
owner_user_id=project.owner_user_id,
|
||||
name=project.name,
|
||||
description=project.description,
|
||||
shared_users=project.shared_users,
|
||||
is_default=True,
|
||||
created_at=project.created_at,
|
||||
)
|
||||
try:
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
return project
|
||||
except IntegrityError:
|
||||
# 并发:另一个请求已插入默认项目,回滚后重查
|
||||
self.session.rollback()
|
||||
existing = self.find_default_by_owner(owner_user_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
raise
|
||||
|
||||
@@ -6,7 +6,10 @@ from typing import List, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import TemplateClipConfigModel
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
TemplateClipConfigModel,
|
||||
TemplateModel,
|
||||
)
|
||||
from packages.domain.template_clip_config import (
|
||||
ClipType,
|
||||
TemplateClipConfig,
|
||||
@@ -38,6 +41,24 @@ class SQLAlchemyTemplateClipConfigRepository:
|
||||
models = query.offset(skip).limit(limit).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def template_owned_by(self, template_id: str, user_id: str) -> bool:
|
||||
"""校验旧模板主表 ``templates`` 中模板归属当前用户且未删除(is_active=True).
|
||||
|
||||
片段配置主表 ``template_clip_configs`` 本身没有 user_id 列,
|
||||
归属关系通过模板主表 ``templates.user_id`` 确定。
|
||||
新表 ``edit_templates`` 为全局模板库(无 user_id 列),不走此校验。
|
||||
"""
|
||||
return (
|
||||
self.session.query(TemplateModel.id)
|
||||
.filter(
|
||||
TemplateModel.id == template_id,
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
)
|
||||
.first()
|
||||
is not None
|
||||
)
|
||||
|
||||
def get(self, config_id: str) -> Optional[TemplateClipConfig]:
|
||||
"""根据 ID 获取配置"""
|
||||
model = self.session.query(TemplateClipConfigModel).filter(TemplateClipConfigModel.id == config_id).first()
|
||||
|
||||
@@ -10,6 +10,7 @@ from __future__ import annotations
|
||||
import uuid
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy import or_
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
@@ -28,6 +29,20 @@ class SQLAlchemyTemplateRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self.session = session
|
||||
|
||||
def _filter_with_segment_configs(self, query):
|
||||
"""只保留在 template_clip_configs 或 template_segments 中存在片段配置的模板。
|
||||
|
||||
两张表都没有记录的模板无法用于生成(from-assets 会 400),
|
||||
剪辑页选模板时应排除;模板编辑器不传 valid_only,仍可见全部模板。
|
||||
"""
|
||||
has_clip_config = self.session.query(TemplateClipConfigModel.id).filter(
|
||||
TemplateClipConfigModel.template_id == TemplateModel.id,
|
||||
)
|
||||
has_segment = self.session.query(TemplateSegmentModel.id).filter(
|
||||
TemplateSegmentModel.template_id == TemplateModel.id,
|
||||
)
|
||||
return query.filter(or_(has_clip_config.exists(), has_segment.exists()))
|
||||
|
||||
# ── Template CRUD ──
|
||||
|
||||
def list_by_user(
|
||||
@@ -40,11 +55,14 @@ class SQLAlchemyTemplateRepository:
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
valid_only: bool = False,
|
||||
) -> List[Template]:
|
||||
query = self.session.query(TemplateModel).filter(
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
)
|
||||
if valid_only:
|
||||
query = self._filter_with_segment_configs(query)
|
||||
if category:
|
||||
query = query.filter(TemplateModel.category == category)
|
||||
if mode:
|
||||
@@ -102,6 +120,27 @@ class SQLAlchemyTemplateRepository:
|
||||
template.segments = self.list_segments(template.id)
|
||||
return template
|
||||
|
||||
def get_active(self, template_id: str, user_id: str) -> Optional[Template]:
|
||||
"""获取归属当前用户且未删除(is_active=True)的模板,否则返回 None.
|
||||
|
||||
用于编辑器访问门禁:模板不存在、已软删除或不属于当前用户时返回 None,
|
||||
由调用方映射为 404。与 :meth:`get` 的区别是额外过滤 is_active。
|
||||
"""
|
||||
model = (
|
||||
self.session.query(TemplateModel)
|
||||
.filter(
|
||||
TemplateModel.id == template_id,
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
template = self._model_to_entity(model)
|
||||
template.segments = self.list_segments(template.id)
|
||||
return template
|
||||
|
||||
def create(self, template: Template) -> Template:
|
||||
model = TemplateModel(
|
||||
id=template.id,
|
||||
@@ -173,11 +212,14 @@ class SQLAlchemyTemplateRepository:
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
valid_only: bool = False,
|
||||
) -> int:
|
||||
query = self.session.query(TemplateModel).filter(
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
)
|
||||
if valid_only:
|
||||
query = self._filter_with_segment_configs(query)
|
||||
if category:
|
||||
query = query.filter(TemplateModel.category == category)
|
||||
if mode:
|
||||
|
||||
@@ -62,6 +62,7 @@ class ListTemplatesFilter:
|
||||
tag: Optional[str] = None
|
||||
keyword: Optional[str] = None
|
||||
mode: Optional[str] = None
|
||||
valid_only: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -106,6 +106,7 @@ class ListTemplatesUseCase:
|
||||
tag=filter.tag,
|
||||
keyword=filter.keyword,
|
||||
mode=filter.mode,
|
||||
valid_only=filter.valid_only,
|
||||
)
|
||||
|
||||
|
||||
@@ -127,6 +128,7 @@ class CountTemplatesUseCase:
|
||||
tag=filter.tag,
|
||||
keyword=filter.keyword,
|
||||
mode=filter.mode,
|
||||
valid_only=filter.valid_only,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,314 @@
|
||||
"""BGM 池差异化分配 — 打破变体间音频指纹一致性 (Issue #1767).
|
||||
|
||||
三层递进策略:
|
||||
1. **BGM 池分配(核心)**:维护风格匹配的 BGM 池,每个变体基于 variant_seed
|
||||
随机分配一首不同 BGM,保证变体间音频指纹不同。
|
||||
2. **段落差异化(池不够时的补充)**:同一首 BGM 做差异化裁剪,不同变体使用
|
||||
不同起始点/段落,进一步降低音频相似度。
|
||||
3. **音量微调**:不同变体 BGM 音量 ±3dB 微调,混音比例有微小差异。
|
||||
|
||||
约束:
|
||||
- 不破坏现有单视频 BGM 选择逻辑(单视频不走池分配)
|
||||
- BGM 情绪/风格与视频内容匹配(基于源 plan 的 BGM style 做风格筛选)
|
||||
- 分配可复现(同 seed 同结果)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── BGM 池条目 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BGMPoolEntry:
|
||||
"""BGM 池条目"""
|
||||
|
||||
id: str
|
||||
preset_id: str # 关联 PRESET_BGM_LIBRARY 中的 ID(用于渲染侧解析音频路径)
|
||||
mood: str # 情绪/风格:upbeat / relax / tech / commerce / emotional / cinematic
|
||||
duration: float # 时长(秒)
|
||||
audio_url: str = "" # CDN/OSS 直链(优先级高于 preset_id)
|
||||
tags: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
# ── BGM 池(10 首,覆盖 6 种风格) ──────────────────────────────────────────
|
||||
|
||||
BGM_POOL: list[BGMPoolEntry] = [
|
||||
# upbeat (轻快)
|
||||
BGMPoolEntry(
|
||||
id="pool_upbeat_001", preset_id="bgm_upbeat_001", mood="upbeat", duration=120.0, tags=["轻快", "阳光", "vlog"]
|
||||
),
|
||||
BGMPoolEntry(
|
||||
id="pool_upbeat_002", preset_id="bgm_upbeat_002", mood="upbeat", duration=95.0, tags=["轻快", "电子", "运动"]
|
||||
),
|
||||
BGMPoolEntry(
|
||||
id="pool_upbeat_003", preset_id="bgm_upbeat_003", mood="upbeat", duration=110.0, tags=["轻快", "夏日", "旅行"]
|
||||
),
|
||||
# relax (治愈)
|
||||
BGMPoolEntry(
|
||||
id="pool_relax_001", preset_id="bgm_relax_001", mood="relax", duration=180.0, tags=["治愈", "钢琴", "冥想"]
|
||||
),
|
||||
BGMPoolEntry(
|
||||
id="pool_relax_002", preset_id="bgm_relax_002", mood="relax", duration=150.0, tags=["治愈", "自然", "放松"]
|
||||
),
|
||||
BGMPoolEntry(
|
||||
id="pool_relax_003", preset_id="bgm_relax_003", mood="relax", duration=200.0, tags=["治愈", "古典", "钢琴"]
|
||||
),
|
||||
# tech (科技)
|
||||
BGMPoolEntry(
|
||||
id="pool_tech_001", preset_id="bgm_tech_001", mood="tech", duration=85.0, tags=["科技", "电子", "数码"]
|
||||
),
|
||||
BGMPoolEntry(
|
||||
id="pool_tech_002", preset_id="bgm_tech_002", mood="tech", duration=100.0, tags=["科技", "极简", "AI"]
|
||||
),
|
||||
# commerce (电商)
|
||||
BGMPoolEntry(
|
||||
id="pool_commerce_001",
|
||||
preset_id="bgm_commerce_001",
|
||||
mood="commerce",
|
||||
duration=75.0,
|
||||
tags=["电商", "时尚", "带货"],
|
||||
),
|
||||
BGMPoolEntry(
|
||||
id="pool_commerce_002",
|
||||
preset_id="bgm_commerce_002",
|
||||
mood="commerce",
|
||||
duration=90.0,
|
||||
tags=["电商", "品牌", "品质"],
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
# ── 风格 → 情绪映射 ────────────────────────────────────────────────────────
|
||||
# preset_bgm.py 中 style 字段 → bgm_pool.py 中 mood 字段
|
||||
|
||||
STYLE_TO_MOOD: dict[str, str] = {
|
||||
"upbeat": "upbeat",
|
||||
"relax": "relax",
|
||||
"tech": "tech",
|
||||
"commerce": "commerce",
|
||||
"emotional": "emotional",
|
||||
"cinematic": "cinematic",
|
||||
}
|
||||
|
||||
|
||||
# ── 策略一:BGM 池分配 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def get_bgm_pool_candidates(source_mood: str | None = None) -> list[BGMPoolEntry]:
|
||||
"""获取 BGM 池候选列表。
|
||||
|
||||
如果指定了 source_mood,优先返回同 mood 的条目;
|
||||
如果同 mood 条目不足 2 个,降级返回全池(保证有足够候选)。
|
||||
|
||||
Args:
|
||||
source_mood: 源 BGM 的情绪/风格(来自 preset_bgm.py 的 style 字段)
|
||||
|
||||
Returns:
|
||||
候选 BGM 列表(至少 2 个条目)
|
||||
"""
|
||||
if not source_mood:
|
||||
return list(BGM_POOL)
|
||||
|
||||
mood = STYLE_TO_MOOD.get(source_mood, source_mood)
|
||||
matched = [e for e in BGM_POOL if e.mood == mood]
|
||||
|
||||
# 同 mood 至少要有 2 首,否则无法"差异化",降级全池
|
||||
if len(matched) >= 2:
|
||||
return matched
|
||||
return list(BGM_POOL)
|
||||
|
||||
|
||||
def select_bgm_from_pool(
|
||||
variant_seed: int,
|
||||
candidates: list[BGMPoolEntry] | None = None,
|
||||
) -> BGMPoolEntry:
|
||||
"""基于 variant_seed 从候选池中选一首 BGM(可复现)。
|
||||
|
||||
Args:
|
||||
variant_seed: 变体随机种子
|
||||
candidates: 候选池(None 时使用全池)
|
||||
|
||||
Returns:
|
||||
选中的 BGM 条目
|
||||
"""
|
||||
pool = candidates if candidates is not None else list(BGM_POOL)
|
||||
if not pool:
|
||||
pool = list(BGM_POOL)
|
||||
rng = random.Random(variant_seed)
|
||||
return rng.choice(pool)
|
||||
|
||||
|
||||
# ── 策略二:段落差异化 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def generate_bgm_segment_offset(variant_seed: int, bgm_duration: float) -> float:
|
||||
"""为变体生成 BGM 段落起始偏移(策略二)。
|
||||
|
||||
不同变体从同一首 BGM 的不同位置开始播放,进一步降低音频指纹相似度。
|
||||
|
||||
偏移范围 [0, max_offset],max_offset = min(30s, bgm_duration * 0.3)。
|
||||
量化到 5 秒整数倍,便于复现和调试。
|
||||
|
||||
Args:
|
||||
variant_seed: 变体随机种子
|
||||
bgm_duration: BGM 总时长(秒)
|
||||
|
||||
Returns:
|
||||
起始偏移(秒),0 ~ max_offset 之间,5s 步长
|
||||
"""
|
||||
rng = random.Random(variant_seed + 7919) # 加素数偏移,避免与 BGM 选择 seed 序列重合
|
||||
max_offset = min(30.0, bgm_duration * 0.3)
|
||||
steps = int(max_offset // 5.0)
|
||||
if steps <= 0:
|
||||
return 0.0
|
||||
return float(rng.randint(0, steps) * 5)
|
||||
|
||||
|
||||
# ── 策略三:音量微调 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def generate_bgm_volume_adjust(variant_seed: int) -> float:
|
||||
"""为变体生成 BGM 音量微调值(策略三)。
|
||||
|
||||
±3dB 微调,让不同变体的 BGM/配音混音比例有微小差异。
|
||||
离散步长:-3, -2, -1, 0, 1, 2, 3 dB。
|
||||
|
||||
Args:
|
||||
variant_seed: 变体随机种子
|
||||
|
||||
Returns:
|
||||
音量调整值(dB),-3.0 ~ 3.0
|
||||
"""
|
||||
rng = random.Random(variant_seed + 104729) # 另一个素数偏移
|
||||
return float(rng.choice([-3, -2, -1, 0, 1, 2, 3]))
|
||||
|
||||
|
||||
# ── 批量分配入口 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def allocate_bgm_pool_for_variants(
|
||||
source_bgm_config: dict,
|
||||
variant_seeds: list[int],
|
||||
) -> list[dict]:
|
||||
"""为批量变体分配不同的 BGM 池配置。
|
||||
|
||||
整合三层策略:池分配 + 段落偏移 + 音量微调。
|
||||
每个变体得到一个 dict,可直接合并到 plan.config["bgm"] 中。
|
||||
|
||||
Args:
|
||||
source_bgm_config: 源 plan 的 BGM 配置(用于风格匹配)
|
||||
variant_seeds: 每个变体的随机种子列表
|
||||
|
||||
Returns:
|
||||
每个变体的 BGM 池配置 dict 列表(与 variant_seeds 等长),
|
||||
每项包含 preset_id / audio_url / audio_offset / volume_adjust_db。
|
||||
如果源 BGM 未启用,返回空列表。
|
||||
"""
|
||||
if not source_bgm_config or not source_bgm_config.get("enabled", False):
|
||||
return []
|
||||
if not variant_seeds:
|
||||
return []
|
||||
|
||||
# 从源 BGM 配置中推断风格
|
||||
source_mood = _infer_source_mood(source_bgm_config)
|
||||
|
||||
# 策略一:获取候选池
|
||||
candidates = get_bgm_pool_candidates(source_mood)
|
||||
|
||||
# 为每个变体分配不同的 BGM(尽量不重复)
|
||||
assignments = _assign_unique_bgm(candidates, variant_seeds)
|
||||
|
||||
results = []
|
||||
for i, (entry, seed) in enumerate(zip(assignments, variant_seeds, strict=False)):
|
||||
# 策略二:段落偏移
|
||||
offset = generate_bgm_segment_offset(seed, entry.duration)
|
||||
|
||||
# 策略三:音量微调
|
||||
volume_adj = generate_bgm_volume_adjust(seed)
|
||||
|
||||
result = {
|
||||
"preset_id": entry.preset_id,
|
||||
"audio_url": entry.audio_url,
|
||||
"audio_offset": offset,
|
||||
"volume_adjust_db": volume_adj,
|
||||
"bgm_pool_entry_id": entry.id,
|
||||
"bgm_pool_mood": entry.mood,
|
||||
}
|
||||
results.append(result)
|
||||
logger.info(
|
||||
"变体 %d BGM 池分配: seed=%d bgm=%s mood=%s offset=%.1fs vol_adj=%+.0fdB",
|
||||
i,
|
||||
seed,
|
||||
entry.id,
|
||||
entry.mood,
|
||||
offset,
|
||||
volume_adj,
|
||||
)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def _infer_source_mood(source_bgm_config: dict) -> str | None:
|
||||
"""从源 BGM 配置推断风格/情绪。
|
||||
|
||||
优先级:
|
||||
1. preset_id → 查 preset_bgm 库获取 style
|
||||
2. bgm_pool_mood → 上游已设置过(二次分配场景)
|
||||
3. 无法推断 → None(返回全池候选)
|
||||
"""
|
||||
preset_id = source_bgm_config.get("preset_id", "")
|
||||
if preset_id:
|
||||
from packages.domain.preset_bgm import get_preset_bgm
|
||||
|
||||
preset = get_preset_bgm(preset_id)
|
||||
if preset:
|
||||
return STYLE_TO_MOOD.get(preset.style, preset.style)
|
||||
|
||||
# 如果之前已经分配过 BGM 池,直接用 mood
|
||||
pool_mood = source_bgm_config.get("bgm_pool_mood", "")
|
||||
if pool_mood:
|
||||
return pool_mood
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _assign_unique_bgm(
|
||||
candidates: list[BGMPoolEntry],
|
||||
variant_seeds: list[int],
|
||||
) -> list[BGMPoolEntry]:
|
||||
"""尽量让每个变体选到不同的 BGM。
|
||||
|
||||
策略:用 seed 选 BGM,如果与前面变体重复,用递增 seed 重试。
|
||||
如果候选池大小 < 变体数,允许重复但不连续。
|
||||
"""
|
||||
if not candidates or not variant_seeds:
|
||||
return []
|
||||
|
||||
assignments: list[BGMPoolEntry] = []
|
||||
used_ids: set[str] = set()
|
||||
|
||||
for i, seed in enumerate(variant_seeds):
|
||||
rng = random.Random(seed)
|
||||
# 先尝试选一个没用过的
|
||||
chosen = None
|
||||
for _attempt in range(len(candidates)):
|
||||
candidate = rng.choice(candidates)
|
||||
if candidate.id not in used_ids:
|
||||
chosen = candidate
|
||||
break
|
||||
if chosen is None:
|
||||
# 候选池已用完,允许重复但取下一个(循环)
|
||||
idx = i % len(candidates)
|
||||
chosen = candidates[idx]
|
||||
|
||||
assignments.append(chosen)
|
||||
used_ids.add(chosen.id)
|
||||
|
||||
return assignments
|
||||
@@ -70,10 +70,11 @@ class Project:
|
||||
name: str
|
||||
description: str = ""
|
||||
shared_users: list[str] = field(default_factory=list) # 被共享的用户 ID 列表
|
||||
is_default: bool = False # 是否为用户的默认项目(小程序自动创建),DB 部分唯一索引保证每人至多一个
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@classmethod
|
||||
def create(cls, owner_user_id: str, name: str, description: str = "") -> "Project":
|
||||
def create(cls, owner_user_id: str, name: str, description: str = "", is_default: bool = False) -> "Project":
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
raise ValueError("项目名称不能为空")
|
||||
@@ -83,6 +84,7 @@ class Project:
|
||||
name=clean_name,
|
||||
description=description.strip(),
|
||||
shared_users=[],
|
||||
is_default=is_default,
|
||||
)
|
||||
|
||||
def is_owner(self, user_id: str) -> bool:
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
"""转场位置与类型随机化(Issue #1766).
|
||||
|
||||
为不同变体生成不同的转场序列(类型 + 时长 + 位置微调),
|
||||
打破"所有变体转场节奏完全一致"的结构相似性,
|
||||
降低平台查重识别为结构相似视频的风险。
|
||||
|
||||
设计要点:
|
||||
1. 转场类型池:5 种效果(dissolve / zoom / slideleft / wipeleft / fade)
|
||||
2. 硬切概率:保证 30%-50% 的转场是硬切(保持节奏感)
|
||||
3. 转场时长随机:0.3s ~ 0.8s
|
||||
4. 转场位置微调:±0.5s 偏移(通过 xfade jitter 实现)
|
||||
5. 与 #1764 节奏模板协同:长片段之间的转场倾向于更长时长
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 转场类型池 ──────────────────────────────────────────────────────────────
|
||||
|
||||
#: 非硬切转场类型池(5 种效果,均为 FFmpeg xfade 支持的 transition 名称)
|
||||
#: 选择标准:视觉效果差异大、FFmpeg 渲染稳定、肉眼可区分
|
||||
TRANSITION_POOL: list[str] = [
|
||||
"dissolve", # 溶解
|
||||
"zoomin", # 缩放(放大进入)
|
||||
"slideleft", # 左滑入
|
||||
"wipeleft", # 左擦除
|
||||
"fade", # 淡入淡出
|
||||
]
|
||||
|
||||
#: 硬切(无转场效果),由 build_xfade_filter_chain 特殊处理(concat filter)
|
||||
HARD_CUT = "cut"
|
||||
|
||||
# ── 时长约束 ────────────────────────────────────────────────────────────────
|
||||
|
||||
#: 随机转场时长下限(秒)
|
||||
TRANSITION_DURATION_MIN = 0.3
|
||||
|
||||
#: 随机转场时长上限(秒)
|
||||
TRANSITION_DURATION_MAX = 0.8
|
||||
|
||||
# ── 硬切比例 ────────────────────────────────────────────────────────────────
|
||||
|
||||
#: 硬切概率下限(至少 30% 硬切,保持节奏感)
|
||||
CUT_RATIO_MIN = 0.3
|
||||
|
||||
#: 硬切概率上限(最多 50% 硬切,保证足够视觉变化)
|
||||
CUT_RATIO_MAX = 0.5
|
||||
|
||||
# ── 位置微调 ────────────────────────────────────────────────────────────────
|
||||
|
||||
#: 转场位置最大偏移(秒),实际偏移在 [-MAX, +MAX] 均匀分布
|
||||
#: 正 = 转场推迟(多留一点前一片段),负 = 转场提前
|
||||
TIMING_JITTER_MAX = 0.5
|
||||
|
||||
|
||||
# ── 内部辅助 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _cut_probability_for_pair(
|
||||
prev_duration: float,
|
||||
next_duration: float,
|
||||
) -> float:
|
||||
"""根据相邻片段时长计算硬切概率。
|
||||
|
||||
与 #1764 节奏模板协同:
|
||||
- 两片段都较短(<3s,快节奏)→ 硬切概率更高(节奏更紧凑)
|
||||
- 两片段都较长(>6s,慢节奏)→ 硬切概率稍低(留出过渡空间)
|
||||
- 混合场景 → 基准概率(CUT_RATIO_MIN + CUT_RATIO_MAX)/ 2
|
||||
|
||||
返回概率始终在 [CUT_RATIO_MIN, CUT_RATIO_MAX] 范围内。
|
||||
"""
|
||||
base = (CUT_RATIO_MIN + CUT_RATIO_MAX) / 2 # 0.4
|
||||
avg_dur = (prev_duration + next_duration) / 2
|
||||
|
||||
if avg_dur < 3.0:
|
||||
# 快节奏:硬切概率偏高
|
||||
return min(CUT_RATIO_MAX, base + 0.1)
|
||||
elif avg_dur > 6.0:
|
||||
# 慢节奏:硬切概率偏低(更多视觉过渡)
|
||||
return max(CUT_RATIO_MIN, base - 0.1)
|
||||
return base
|
||||
|
||||
|
||||
def _apply_jitter(
|
||||
base_duration: float,
|
||||
jitter: float,
|
||||
clip_duration: float,
|
||||
) -> float:
|
||||
"""给转场时长应用微调偏移,钳制到安全范围。
|
||||
|
||||
Args:
|
||||
base_duration: 基础转场时长
|
||||
jitter: 偏移量(可正可负)
|
||||
clip_duration: 较短的相邻片段时长(转场不能超过此值)
|
||||
|
||||
Returns:
|
||||
钳制后的实际转场时长
|
||||
"""
|
||||
effective = base_duration + jitter
|
||||
# 上界:不超过相邻片段时长的 40%(留足内容时间),也不超过 MAX
|
||||
upper = min(TRANSITION_DURATION_MAX, clip_duration * 0.4)
|
||||
lower = TRANSITION_DURATION_MIN if jitter < 0 else max(TRANSITION_DURATION_MIN, base_duration)
|
||||
return max(lower, min(upper, effective))
|
||||
|
||||
|
||||
# ── 核心函数 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def generate_transition_plan(
|
||||
num_transitions: int,
|
||||
*,
|
||||
clip_durations: list[float] | None = None,
|
||||
rng: random.Random | None = None,
|
||||
) -> list[dict]:
|
||||
"""为变体生成一组随机化的转场计划。
|
||||
|
||||
每个转场点独立随机选择类型和时长,硬切比例保持在 30%-50%。
|
||||
|
||||
Args:
|
||||
num_transitions: 转场点数量(= 主片段数 - 1)
|
||||
clip_durations: 各片段时长(用于协同节奏:长片段间转场更长),
|
||||
长度应 >= num_transitions + 1;不足时用默认值
|
||||
rng: 可选随机数生成器(测试可注入固定种子)
|
||||
|
||||
Returns:
|
||||
转场计划列表,每项:
|
||||
- effect: str — "cut" 或 TRANSITION_POOL 中某一效果
|
||||
- duration: float — 转场时长(cut 为 0.0)
|
||||
- jitter: float — 位置偏移量(秒,-0.5 ~ +0.5)
|
||||
"""
|
||||
if num_transitions <= 0:
|
||||
return []
|
||||
if rng is None:
|
||||
rng = random.Random()
|
||||
|
||||
if clip_durations is None:
|
||||
clip_durations = [5.0] * (num_transitions + 1)
|
||||
|
||||
plan: list[dict] = []
|
||||
for i in range(num_transitions):
|
||||
prev_dur = clip_durations[i] if i < len(clip_durations) else 5.0
|
||||
next_dur = clip_durations[i + 1] if (i + 1) < len(clip_durations) else 5.0
|
||||
|
||||
# 计算硬切概率(协同节奏)
|
||||
cut_prob = _cut_probability_for_pair(prev_dur, next_dur)
|
||||
|
||||
# 随机决定是否硬切
|
||||
if rng.random() < cut_prob:
|
||||
effect = HARD_CUT
|
||||
duration = 0.0
|
||||
else:
|
||||
effect = rng.choice(TRANSITION_POOL)
|
||||
base_dur = rng.uniform(TRANSITION_DURATION_MIN, TRANSITION_DURATION_MAX)
|
||||
# 协同节奏:长片段间转场基础时长更长
|
||||
avg_dur = (prev_dur + next_dur) / 2
|
||||
if avg_dur > 6.0:
|
||||
base_dur = min(TRANSITION_DURATION_MAX, base_dur * 1.15)
|
||||
# 应用位置微调
|
||||
jitter = rng.uniform(-TIMING_JITTER_MAX, TIMING_JITTER_MAX)
|
||||
shorter_clip = min(prev_dur, next_dur)
|
||||
duration = _apply_jitter(base_dur, jitter, shorter_clip)
|
||||
|
||||
jitter_val = rng.uniform(-TIMING_JITTER_MAX, TIMING_JITTER_MAX) if effect != HARD_CUT else 0.0
|
||||
|
||||
plan.append(
|
||||
{
|
||||
"effect": effect,
|
||||
"duration": round(duration, 3),
|
||||
"jitter": round(jitter_val, 3),
|
||||
}
|
||||
)
|
||||
|
||||
return plan
|
||||
@@ -29,6 +29,7 @@ import logging
|
||||
import random
|
||||
|
||||
from packages.domain.plan_generator_utils import _resolve_start_time
|
||||
from packages.domain.transition_randomizer import generate_transition_plan
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -129,7 +130,10 @@ def reselect_clips_for_variant(
|
||||
Raises:
|
||||
ValueError: 源片段为空 / 素材池为空 / 素材时长全为 0(无法差异化选片)。
|
||||
"""
|
||||
rng = rng or random.Random()
|
||||
if rng is None:
|
||||
rng = random.Random()
|
||||
elif isinstance(rng, int):
|
||||
rng = random.Random(rng)
|
||||
if not source_clips:
|
||||
raise ValueError("源 plan 无片段,无法为变体重新选片")
|
||||
if not candidate_asset_ids:
|
||||
@@ -222,6 +226,9 @@ def reselect_clips_for_variant(
|
||||
batch_segments.setdefault(aid, []).append(interval)
|
||||
result[idx] = _base_clip_data(src, asset_id=aid, start=start, duration=target_dur)
|
||||
|
||||
# ── 4. #1766 转场随机化:为相邻 main 片段对生成随机转场序列 ────────
|
||||
_apply_transition_randomization(result, rng)
|
||||
|
||||
return [c for c in result if c is not None]
|
||||
|
||||
|
||||
@@ -312,7 +319,10 @@ def generate_visual_perturbation(rng: random.Random | None = None) -> dict:
|
||||
- speed_factor: 0.95~1.05 速度微调(±5%,肉眼不太敏感但时间轴不同)
|
||||
- brightness_shift: -10~+10 亮度偏移(eq=brightness,画面明暗差异)
|
||||
"""
|
||||
rng = rng or random.Random()
|
||||
if rng is None:
|
||||
rng = random.Random()
|
||||
elif isinstance(rng, int):
|
||||
rng = random.Random(rng)
|
||||
return {
|
||||
"hflip": rng.random() < 0.3,
|
||||
"zoom_ratio": round(1.0 + rng.uniform(0, 0.08), 4),
|
||||
@@ -321,7 +331,7 @@ def generate_visual_perturbation(rng: random.Random | None = None) -> dict:
|
||||
}
|
||||
|
||||
|
||||
def generate_pixel_perturbation(rng: random.Random | None = None) -> dict:
|
||||
def generate_pixel_perturbation(rng: random.Random | int | None = None) -> dict:
|
||||
"""为一个变体生成像素级扰动滤镜参数(Issue #1765)。
|
||||
|
||||
在现有视觉扰动(hflip/zoom/brightness)基础上,额外叠加 2-3 种
|
||||
@@ -336,7 +346,10 @@ def generate_pixel_perturbation(rng: random.Random | None = None) -> dict:
|
||||
返回 dict,可直接存入 plan.config["pixel_perturbation"]。
|
||||
渲染侧读取后追加到 ffmpeg filter chain。
|
||||
"""
|
||||
rng = rng or random.Random()
|
||||
if rng is None:
|
||||
rng = random.Random()
|
||||
elif isinstance(rng, int):
|
||||
rng = random.Random(rng)
|
||||
|
||||
# 可用滤镜池
|
||||
filter_options = ["noise", "unsharp", "curves", "color_balance"]
|
||||
@@ -369,6 +382,52 @@ def generate_pixel_perturbation(rng: random.Random | None = None) -> dict:
|
||||
return result
|
||||
|
||||
|
||||
def _apply_transition_randomization(
|
||||
result: list[dict | None],
|
||||
rng: random.Random,
|
||||
) -> None:
|
||||
"""#1766 对 result 中相邻 main 片段应用转场随机化(就地修改)。
|
||||
|
||||
为每对相邻 main 片段独立选择:
|
||||
- 转场类型(TRANSITION_POOL 中随机,或硬切)
|
||||
- 转场时长(0.3s ~ 0.8s,协同片段时长)
|
||||
- 位置微调 jitter(±0.5s,存入 config["transition_jitter"])
|
||||
|
||||
硬切比例保证在 30%-50%。intro/outro 等非 main 片段的转场保持源值不变。
|
||||
"""
|
||||
# 收集 main 片段的索引(按 order 排序)
|
||||
main_indices = [i for i, c in enumerate(result) if c is not None and c.get("clip_type", "main") == "main"]
|
||||
|
||||
if len(main_indices) < 2:
|
||||
# 不足 2 个 main 片段,无转场点可随机化
|
||||
return
|
||||
|
||||
num_transitions = len(main_indices) - 1
|
||||
# 用 main 片段的 duration 作为协同节奏的输入
|
||||
clip_durations = [result[i]["duration"] for i in main_indices]
|
||||
|
||||
plan = generate_transition_plan(
|
||||
num_transitions,
|
||||
clip_durations=clip_durations,
|
||||
rng=rng,
|
||||
)
|
||||
|
||||
# 将转场计划应用到每对相邻 main 片段
|
||||
# plan[k] 是 main_indices[k] → main_indices[k+1] 之间的转场
|
||||
# 转场信息存储在"目标 clip"(即每对的第二个)的 transition_effect/duration
|
||||
for k, transition_info in enumerate(plan):
|
||||
target_idx = main_indices[k + 1]
|
||||
if result[target_idx] is None:
|
||||
continue
|
||||
clip = result[target_idx]
|
||||
clip["transition_effect"] = transition_info["effect"]
|
||||
clip["transition_duration"] = transition_info["duration"]
|
||||
# jitter 存入 config,供渲染侧 xfade_builder 读取
|
||||
cfg = clip.get("config") or {}
|
||||
cfg["transition_jitter"] = transition_info["jitter"]
|
||||
clip["config"] = cfg
|
||||
|
||||
|
||||
def _base_clip_data(src: dict, *, asset_id: str, start: float, duration: float | None = None) -> dict:
|
||||
"""从源片段构造落库 dict(保留骨架/转场/文案/速度,替换素材与起点)。"""
|
||||
return {
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
|
||||
@@ -373,3 +373,193 @@ def _append_audio_concat(parts: list[str], clip_chains: list[ClipFilterChain]) -
|
||||
# concat 滤镜(使用 audio_label 作为输入)
|
||||
audio_inputs = "".join(f"[{c.audio_label}]" for c in audio_chains)
|
||||
parts.append(f"{audio_inputs}concat=n={len(audio_chains)}:v=0:a=1[outa]")
|
||||
|
||||
|
||||
# ── 标题 drawtext 滤镜构建(#1789)─────────────────────────────────────────────
|
||||
|
||||
# drawtext 字体搜索路径:按优先级列出常见安装位置
|
||||
# 服务器使用 Noto Sans SC(思源黑体)作为默认字体
|
||||
DRAWTEXT_FONT_SEARCH_PATHS: list[str] = [
|
||||
"/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc",
|
||||
"/usr/share/fonts/noto-cjk/NotoSansCJK-Regular.ttc",
|
||||
"/usr/share/fonts/google-noto-cjk/NotoSansCJK-Regular.ttc",
|
||||
"/usr/share/fonts/truetype/noto/NotoSansSC-Regular.ttf",
|
||||
"/usr/share/fonts/noto/NotoSansSC-Regular.ttf",
|
||||
"/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf",
|
||||
]
|
||||
|
||||
# 前端字体名 → drawtext 字体搜索关键字
|
||||
DRAWTEXT_FONT_MAP: dict[str, str] = {
|
||||
"思源黑体": "NotoSansCJK",
|
||||
"思源宋体": "NotoSerifCJK",
|
||||
"苹方": "NotoSansCJK",
|
||||
"PingFang": "NotoSansCJK",
|
||||
"微软雅黑": "NotoSansCJK",
|
||||
"楷体": "NotoSerifCJK",
|
||||
"华康俪金黑": "NotoSansCJK",
|
||||
}
|
||||
|
||||
|
||||
def _escape_drawtext_text(text: str) -> str:
|
||||
"""转义 drawtext 特殊字符。
|
||||
|
||||
FFmpeg drawtext 要求转义:
|
||||
- \\ → \\\\
|
||||
- ' → \\\\'
|
||||
- : → \\\\:
|
||||
- % → %%(drawtext 中 % 是时间码特殊字符)
|
||||
"""
|
||||
result = text.replace("\\", "\\\\\\\\")
|
||||
result = result.replace("'", "\\\\'")
|
||||
result = result.replace(":", "\\\\:")
|
||||
result = result.replace("%", "%%")
|
||||
return result
|
||||
|
||||
|
||||
def _resolve_font_path(font_name: str) -> str:
|
||||
"""解析字体名到服务器实际字体文件路径。
|
||||
|
||||
查找策略:
|
||||
1. 通过 DRAWTEXT_FONT_MAP 映射前端字体名到服务器关键字
|
||||
2. 在 DRAWTEXT_FONT_SEARCH_PATHS 中查找匹配路径
|
||||
3. 未找到则返回空字符串(drawtext 使用内置默认字体)
|
||||
"""
|
||||
keyword = DRAWTEXT_FONT_MAP.get(font_name, font_name)
|
||||
import os
|
||||
|
||||
for path in DRAWTEXT_FONT_SEARCH_PATHS:
|
||||
if keyword.lower() in path.lower() and os.path.isfile(path):
|
||||
return path
|
||||
# fallback:遍历搜索任意可用字体
|
||||
for path in DRAWTEXT_FONT_SEARCH_PATHS:
|
||||
if os.path.isfile(path):
|
||||
return path
|
||||
return ""
|
||||
|
||||
|
||||
def build_title_drawtext_filter(
|
||||
title_config: dict[str, Any],
|
||||
output_width: int = DEFAULT_OUTPUT_WIDTH,
|
||||
output_height: int = DEFAULT_OUTPUT_HEIGHT,
|
||||
) -> str | None:
|
||||
"""从 title_config 生成 FFmpeg drawtext 滤镜字符串。
|
||||
|
||||
支持前端 TitleSettings 的全部参数:
|
||||
- text / 标题文字
|
||||
- font / 字体名
|
||||
- font_size / 字号
|
||||
- font_color / 颜色(#RRGGBB)
|
||||
- position / 位置(top / center / bottom / custom)
|
||||
- bold / 粗体
|
||||
- stroke / 描边
|
||||
- shadow / 阴影
|
||||
- pos_x, pos_y / 自由位置坐标
|
||||
|
||||
Args:
|
||||
title_config: 标题配置 dict(来自 plan.config["title"])
|
||||
output_width: 输出视频宽度
|
||||
output_height: 输出视频高度
|
||||
|
||||
Returns:
|
||||
drawtext 滤镜字符串;标题为空或 disabled 时返回 None
|
||||
"""
|
||||
if not title_config or not isinstance(title_config, dict):
|
||||
return None
|
||||
|
||||
# 字段名归一化:兼容 content/text、font_preset/font 两套命名
|
||||
text = (title_config.get("text") or title_config.get("content") or "").strip()
|
||||
if not text:
|
||||
return None
|
||||
|
||||
enabled = title_config.get("enabled", True)
|
||||
if not enabled:
|
||||
return None
|
||||
|
||||
# ── 样式参数 ──
|
||||
font_name = title_config.get("font") or title_config.get("font_preset") or "思源黑体"
|
||||
font_size = int(title_config.get("font_size") or title_config.get("size") or 36)
|
||||
font_color = title_config.get("font_color") or title_config.get("color") or "#ffffff"
|
||||
# 去掉 # 前缀(drawtext 用纯 hex 或颜色名)
|
||||
if font_color.startswith("#"):
|
||||
font_color = font_color[1:]
|
||||
|
||||
position = title_config.get("position", "top")
|
||||
bold = bool(title_config.get("bold", True))
|
||||
stroke = title_config.get("stroke")
|
||||
shadow = title_config.get("shadow")
|
||||
|
||||
# ── 构建 drawtext 参数 ──
|
||||
params: list[str] = []
|
||||
|
||||
# 字体文件
|
||||
font_path = _resolve_font_path(font_name)
|
||||
if font_path:
|
||||
escaped_path = font_path.replace("\\", "\\\\").replace(":", "\\\\:").replace("'", "\\\\'")
|
||||
params.append(f"fontfile='{escaped_path}'")
|
||||
|
||||
# 文字内容
|
||||
params.append(f"text='{_escape_drawtext_text(text)}'")
|
||||
|
||||
# 字号 & 颜色
|
||||
params.append(f"fontsize={font_size}")
|
||||
params.append(f"fontcolor={font_color}")
|
||||
|
||||
# 粗体:bold 在 drawtext 中通过 font 的 Bold 变体实现
|
||||
# 若字体有 Bold 变体可用 fontfont=bold;否则通过 borderw 模拟
|
||||
if bold:
|
||||
# 使用 font 参数尝试加载 Bold 变体(Noto Sans SC 有 Bold 变体文件)
|
||||
params.append("font=bold")
|
||||
|
||||
# 描边(borderw 需要 libfreetype 支持)
|
||||
if stroke:
|
||||
if isinstance(stroke, bool):
|
||||
border_width = 2
|
||||
border_color = "black"
|
||||
elif isinstance(stroke, dict):
|
||||
border_width = int(stroke.get("width", 2)) if stroke.get("enabled", True) else 0
|
||||
border_color = (stroke.get("color") or "#000000").lstrip("#")
|
||||
else:
|
||||
border_width = 0
|
||||
border_color = "black"
|
||||
if border_width > 0:
|
||||
params.append(f"borderw={border_width}")
|
||||
params.append(f"bordercolor={border_color}")
|
||||
|
||||
# 阴影(shadowcolor + shadowx/y)
|
||||
if shadow:
|
||||
if isinstance(shadow, bool):
|
||||
params.append("shadowcolor=black")
|
||||
params.append("shadowx=2")
|
||||
params.append("shadowy=2")
|
||||
elif isinstance(shadow, dict):
|
||||
if shadow.get("enabled", True):
|
||||
params.append(f"shadowcolor={(shadow.get('color') or '#000000').lstrip('#')}")
|
||||
params.append(f"shadowx={int(shadow.get('offset_x', 2))}")
|
||||
params.append(f"shadowy={int(shadow.get('offset_y', 2))}")
|
||||
|
||||
# ── 位置计算 ──
|
||||
# 优先使用自定义坐标 pos_x / pos_y
|
||||
pos_x = title_config.get("pos_x")
|
||||
pos_y = title_config.get("pos_y")
|
||||
if (
|
||||
position == "custom"
|
||||
and isinstance(pos_x, (int, float))
|
||||
and isinstance(pos_y, (int, float))
|
||||
and not isinstance(pos_x, bool)
|
||||
and not isinstance(pos_y, bool)
|
||||
):
|
||||
params.append(f"x={int(pos_x)}")
|
||||
params.append(f"y={int(pos_y)}")
|
||||
else:
|
||||
# 三档预设位置:top / center / bottom
|
||||
# x 始终水平居中:(w-text_w)/2
|
||||
params.append("x=(w-text_w)/2")
|
||||
if position == "center":
|
||||
params.append("y=(h-text_h)/2")
|
||||
elif position == "bottom":
|
||||
params.append("y=h-text_h-50")
|
||||
else:
|
||||
# top(默认)
|
||||
params.append("y=50")
|
||||
|
||||
return "drawtext=" + ":".join(params)
|
||||
|
||||
@@ -8,11 +8,13 @@
|
||||
禁止慢放、禁止截断配音;
|
||||
4. 任何情况下不得因素材时长/数量报错打断用户。
|
||||
|
||||
#1764 节奏模板:
|
||||
- 预设 6 种权重序列,不同变体用不同节奏模板
|
||||
#1764 节奏模板 + #1768 多样化增强:
|
||||
- 预设 8 种权重序列,不同变体用不同节奏模板
|
||||
- 片段时长 = 配音总时长 × 该片段权重 / 权重总和
|
||||
- 平均分配作为权重全 1 的特例保留
|
||||
- 每个片段 >= MIN_CLIP_DURATION(2秒)
|
||||
- #1768:最大片段时长 <= 素材可用时长 × 90%
|
||||
- #1768:成片总时长与配音时长误差 <= TOTAL_DURATION_TOLERANCE(0.5s)
|
||||
|
||||
本模块为纯函数:输入片段骨架(每段转场效果/时长)与配音总时长,
|
||||
输出每段目标时长(target duration)与成片总时长。不碰 DB、不碰素材。
|
||||
@@ -42,6 +44,8 @@ RHYTHM_TEMPLATES: list[list[int]] = [
|
||||
[3, 1, 1, 1, 3], # 两端长,中间短
|
||||
[1, 1, 3, 2, 1], # 后段渐长
|
||||
[2, 1, 1, 3, 1], # 前段较长 + 第4段最长
|
||||
[2, 2, 1, 1, 2], # #1768 前重后轻
|
||||
[3, 2, 1, 2, 1], # #1768 渐弱节奏
|
||||
]
|
||||
|
||||
|
||||
@@ -87,13 +91,6 @@ def adapt_template_length(template: list[int], clip_count: int) -> list[int]:
|
||||
return result
|
||||
|
||||
|
||||
#: 单段最小时长(秒):低于此值播放器/渲染链路易出问题
|
||||
MIN_CLIP_DURATION = 1.0
|
||||
|
||||
#: 成片总时长与配音时长的可接受误差(秒)
|
||||
TOTAL_DURATION_TOLERANCE = 0.5
|
||||
|
||||
|
||||
def transition_overlap_seconds(transition_effect: Optional[str], transition_duration: float) -> float:
|
||||
"""转场导致的相邻片段重叠时长。
|
||||
|
||||
@@ -112,6 +109,7 @@ def plan_clip_durations(
|
||||
transition_effects: Optional[list[Optional[str]]] = None,
|
||||
transition_durations: Optional[list[float]] = None,
|
||||
rhythm_template: Optional[list[int]] = None,
|
||||
asset_durations: Optional[list[float]] = None,
|
||||
) -> list[float]:
|
||||
"""把配音总时长分配到 clip_count 段,返回每段目标时长(秒)。
|
||||
|
||||
@@ -129,6 +127,9 @@ def plan_clip_durations(
|
||||
transition_effects: 每段转场效果(长度 clip_count,index 0 的转场无效)。
|
||||
transition_durations: 每段转场时长(长度 clip_count)。
|
||||
|
||||
asset_durations: #1768 每段可用素材时长(秒),用于钳制最大片段时长
|
||||
<= 素材可用时长 × 90%。长度 clip_count;None 或空则不钳制上限。
|
||||
|
||||
Returns:
|
||||
每段目标时长列表(长度 clip_count);无配音/非法输入返回 []。
|
||||
"""
|
||||
@@ -187,12 +188,58 @@ def plan_clip_durations(
|
||||
result[min_idx] = MIN_CLIP_DURATION
|
||||
result[max_idx] = round(result[max_idx] - deficit, 3)
|
||||
|
||||
# #1768:最大片段时长钳制(<= 素材可用时长 × 90%)
|
||||
if asset_durations and len(asset_durations) == clip_count:
|
||||
for _ in range(3): # 迭代收敛
|
||||
clamped = False
|
||||
for i in range(len(result)):
|
||||
try:
|
||||
max_dur = float(asset_durations[i]) * 0.9
|
||||
except (TypeError, ValueError, IndexError):
|
||||
continue
|
||||
if result[i] > max_dur and max_dur >= MIN_CLIP_DURATION:
|
||||
excess = result[i] - max_dur
|
||||
result[i] = round(max_dur, 3)
|
||||
# 将多余时长分配给最短的未超限片段
|
||||
candidates = [
|
||||
j
|
||||
for j in range(len(result))
|
||||
if j != i
|
||||
and (
|
||||
not asset_durations
|
||||
or j >= len(asset_durations)
|
||||
or result[j] < float(asset_durations[j]) * 0.9
|
||||
)
|
||||
]
|
||||
if candidates:
|
||||
shortest = min(candidates, key=lambda j: result[j])
|
||||
result[shortest] = round(result[shortest] + excess, 3)
|
||||
clamped = True
|
||||
if not clamped:
|
||||
break
|
||||
|
||||
# 末段吸收舍入误差
|
||||
total_assigned = sum(result[:-1])
|
||||
result[-1] = round(gross - total_assigned, 3)
|
||||
if result[-1] < MIN_CLIP_DURATION:
|
||||
result[-1] = MIN_CLIP_DURATION
|
||||
|
||||
# #1768:时长总和误差校验(成片净时长 ≈ 配音时长)
|
||||
net_total = total_output_duration(result, transition_effects, transition_durations)
|
||||
deviation = abs(net_total - voice)
|
||||
if deviation > TOTAL_DURATION_TOLERANCE:
|
||||
logger.warning(
|
||||
"#1768 时长总和误差 %.3fs 超过阈值 %.1fs(voice=%.2fs, net=%.2fs),末段补偿修正",
|
||||
deviation,
|
||||
TOTAL_DURATION_TOLERANCE,
|
||||
voice,
|
||||
net_total,
|
||||
)
|
||||
# 修正末段使净时长回归配音时长
|
||||
result[-1] = round(result[-1] + (voice - net_total), 3)
|
||||
if result[-1] < MIN_CLIP_DURATION:
|
||||
result[-1] = MIN_CLIP_DURATION
|
||||
|
||||
return result
|
||||
|
||||
|
||||
|
||||
@@ -109,6 +109,8 @@ def build_xfade_filter_chain(
|
||||
transitions: list[str],
|
||||
*,
|
||||
transition_duration: float = DEFAULT_TRANSITION_DURATION,
|
||||
transition_durations: list[float] | None = None,
|
||||
jitters: list[float] | None = None,
|
||||
output_label: str = "outv",
|
||||
) -> tuple[str, float]:
|
||||
"""构建 xfade 转场滤镜链.
|
||||
@@ -116,11 +118,19 @@ def build_xfade_filter_chain(
|
||||
对每步 xfade 自动钳制 transition duration,确保
|
||||
``offset + td ≤ first_input_duration``,避免 FFmpeg exit 234。
|
||||
|
||||
#1766 增强:支持逐转场独立时长(transition_durations)和位置微调(jitters)。
|
||||
传入 transition_durations 时,每个转场点使用各自的时长,而非全局统一值。
|
||||
jitters 用于在 offset 上做 ±N 秒微调,实现转场位置随机化。
|
||||
|
||||
Args:
|
||||
clip_durations: 每个片段的时长(必须与 trim 后的实际时长一致)
|
||||
clip_video_labels: 每个片段的视频流标签(如 "v0", "v1")
|
||||
transitions: 每个片段对应的转场效果(第一个片段的转场被忽略)
|
||||
transition_duration: 转场时长(秒)
|
||||
transition_duration: 全局默认转场时长(秒),transition_durations 缺失时 fallback
|
||||
transition_durations: #1766 逐转场时长列表(与 transitions 等长),
|
||||
第 i 项对应 transitions[i] 的时长;None 时使用 transition_duration
|
||||
jitters: #1766 逐转场位置偏移列表(秒),与 transitions 等长,
|
||||
正值推迟转场、负值提前转场;None 时不做微调
|
||||
output_label: 最终输出标签
|
||||
|
||||
Returns:
|
||||
@@ -150,15 +160,28 @@ def build_xfade_filter_chain(
|
||||
else:
|
||||
first_input_dur = cumulative - total_transition
|
||||
|
||||
# 正确的 offset 计算:offset 应相对于累积输出时长
|
||||
# offset = 累积输出中,转场开始的时间点
|
||||
# = first_input_dur - transition_duration
|
||||
# 这样每个转场之间的"纯内容"时长等于原始 clip 时长
|
||||
offset = max(0.0, first_input_dur - transition_duration)
|
||||
# #1766: 逐转场时长(优先)或全局默认
|
||||
step_duration = (
|
||||
transition_durations[i - 1]
|
||||
if transition_durations and (i - 1) < len(transition_durations)
|
||||
else transition_duration
|
||||
)
|
||||
|
||||
# 安全钳制:offset + td 不能超过第一个输入的时长
|
||||
# #1766: 位置微调 jitter
|
||||
jitter = jitters[i - 1] if jitters and (i - 1) < len(jitters) else 0.0
|
||||
|
||||
# offset = 转场开始点(相对于累积输出起点)
|
||||
# 基础 offset = first_input_dur - step_duration
|
||||
# jitter > 0 推迟转场(offset 增大);jitter < 0 提前转场(offset 减小)
|
||||
offset = max(0.0, first_input_dur - step_duration) + jitter
|
||||
|
||||
# 安全钳制:offset 不能超出可用范围
|
||||
max_offset = max(0.0, first_input_dur - 0.001)
|
||||
offset = max(0.0, min(offset, max_offset))
|
||||
|
||||
# 安全钳制 td:offset + td 不能超过第一个输入的时长
|
||||
available = max(0.0, first_input_dur - offset)
|
||||
safe_td = min(transition_duration, available)
|
||||
safe_td = min(step_duration, available)
|
||||
|
||||
# 同时不能超过剩余总时长
|
||||
remaining = max(0.0, sum(clip_durations) - cumulative)
|
||||
|
||||
@@ -30,9 +30,16 @@ class AssetLibraryRepository(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def increment_asset_count(self, library_id: str, size_delta: int) -> None:
|
||||
def increment_asset_count(self, library_id: str, count_delta: int = 1, size_delta: int = 0) -> None:
|
||||
"""原子递增素材计数(Issue #1776)。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def decrement_asset_count(self, library_id: str, size_delta: int) -> None:
|
||||
def decrement_asset_count(self, library_id: str, count_delta: int = 1, size_delta: int = 0) -> None:
|
||||
"""原子递减素材计数(Issue #1776)。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def recount_assets(self, library_id: str) -> int:
|
||||
"""重算素材库计数(Issue #1776)。"""
|
||||
pass
|
||||
|
||||
@@ -18,6 +18,7 @@ class TemplateRepositoryPort(Protocol):
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
valid_only: bool = False,
|
||||
) -> List[Template]: ...
|
||||
def get(self, template_id: str, user_id: str) -> Optional[Template]: ...
|
||||
def create(self, template: Template) -> Template: ...
|
||||
@@ -31,6 +32,7 @@ class TemplateRepositoryPort(Protocol):
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
valid_only: bool = False,
|
||||
) -> int: ...
|
||||
def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template: ...
|
||||
def list_segments(self, template_id: str) -> List[TemplateSegment]: ...
|
||||
|
||||
Executable
+159
@@ -0,0 +1,159 @@
|
||||
#!/usr/bin/env python3
|
||||
"""素材库计数重算脚本(Issue #1776)。
|
||||
|
||||
用法:
|
||||
# Dry-run: 输出差异清单,不执行修改
|
||||
python scripts/recount_asset_counts.py --dry-run
|
||||
|
||||
# 执行修正
|
||||
python scripts/recount_asset_counts.py
|
||||
|
||||
# 只处理指定项目
|
||||
python scripts/recount_asset_counts.py --project-id <project_id>
|
||||
|
||||
# 只处理指定素材库
|
||||
python scripts/recount_asset_counts.py --library-id <library_id>
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# 添加项目根目录到 path
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from sqlalchemy import create_engine, func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetLibraryModel, AssetModel
|
||||
|
||||
|
||||
def get_db_session() -> Session:
|
||||
"""创建数据库 session。"""
|
||||
import os
|
||||
|
||||
database_url = os.getenv("DATABASE_URL")
|
||||
if not database_url:
|
||||
print("ERROR: DATABASE_URL environment variable not set")
|
||||
sys.exit(1)
|
||||
engine = create_engine(database_url)
|
||||
return Session(engine)
|
||||
|
||||
|
||||
def check_discrepancies(session: Session, project_id: str | None = None, library_id: str | None = None) -> list[dict]:
|
||||
"""检查素材库计数差异。
|
||||
|
||||
返回列表,每项包含:
|
||||
- library_id: 素材库 ID
|
||||
- library_name: 素材库名称
|
||||
- recorded_count: 记录的计数
|
||||
- actual_count: 实际计数
|
||||
- delta: 差异 (actual - recorded)
|
||||
"""
|
||||
query = session.query(AssetLibraryModel)
|
||||
if project_id:
|
||||
query = query.filter(AssetLibraryModel.project_id == project_id)
|
||||
if library_id:
|
||||
query = query.filter(AssetLibraryModel.id == library_id)
|
||||
|
||||
libraries = query.all()
|
||||
discrepancies = []
|
||||
|
||||
for lib in libraries:
|
||||
# 查询实际计数(排除 deleted)
|
||||
actual_count = (
|
||||
session.query(func.count(AssetModel.id))
|
||||
.filter(
|
||||
AssetModel.asset_library_id == lib.id,
|
||||
AssetModel.status != "deleted",
|
||||
)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
actual_size = (
|
||||
session.query(func.coalesce(func.sum(AssetModel.file_size), 0))
|
||||
.filter(
|
||||
AssetModel.asset_library_id == lib.id,
|
||||
AssetModel.status != "deleted",
|
||||
)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
recorded_count = int(lib.asset_count or 0)
|
||||
recorded_size = int(lib.total_size or 0)
|
||||
|
||||
if actual_count != recorded_count or actual_size != recorded_size:
|
||||
discrepancies.append(
|
||||
{
|
||||
"library_id": lib.id,
|
||||
"library_name": lib.name,
|
||||
"project_id": lib.project_id,
|
||||
"kind": lib.kind,
|
||||
"recorded_count": recorded_count,
|
||||
"actual_count": actual_count,
|
||||
"count_delta": actual_count - recorded_count,
|
||||
"recorded_size": recorded_size,
|
||||
"actual_size": actual_size,
|
||||
"size_delta": actual_size - recorded_size,
|
||||
}
|
||||
)
|
||||
|
||||
return discrepancies
|
||||
|
||||
|
||||
def fix_discrepancies(session: Session, discrepancies: list[dict]) -> int:
|
||||
"""修正素材库计数。返回修正数量。"""
|
||||
fixed = 0
|
||||
for d in discrepancies:
|
||||
session.query(AssetLibraryModel).filter(AssetLibraryModel.id == d["library_id"]).update(
|
||||
{
|
||||
AssetLibraryModel.asset_count: d["actual_count"],
|
||||
AssetLibraryModel.total_size: d["actual_size"],
|
||||
}
|
||||
)
|
||||
fixed += 1
|
||||
session.commit()
|
||||
return fixed
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="素材库计数重算脚本(Issue #1776)")
|
||||
parser.add_argument("--dry-run", action="store_true", help="只输出差异清单,不执行修正")
|
||||
parser.add_argument("--project-id", type=str, help="只处理指定项目")
|
||||
parser.add_argument("--library-id", type=str, help="只处理指定素材库")
|
||||
args = parser.parse_args()
|
||||
|
||||
session = get_db_session()
|
||||
|
||||
try:
|
||||
discrepancies = check_discrepancies(session, args.project_id, args.library_id)
|
||||
|
||||
if not discrepancies:
|
||||
print("✅ 所有素材库计数一致,无需修正")
|
||||
return
|
||||
|
||||
# 输出差异清单
|
||||
print(f"发现 {len(discrepancies)} 个素材库计数不一致:\n")
|
||||
print(f"{'Library ID':<40} {'Name':<20} {'Recorded':<10} {'Actual':<10} {'Delta':<10}")
|
||||
print("-" * 90)
|
||||
for d in discrepancies:
|
||||
print(
|
||||
f"{d['library_id']:<40} {d['library_name'][:20]:<20} {d['recorded_count']:<10} {d['actual_count']:<10} {d['count_delta']:+<10}"
|
||||
)
|
||||
|
||||
total_delta = sum(d["count_delta"] for d in discrepancies)
|
||||
print(f"\n总计差异: {total_delta:+d}")
|
||||
|
||||
if args.dry_run:
|
||||
print("\n[DRY-RUN] 未执行修正。移除 --dry-run 参数以执行修正。")
|
||||
else:
|
||||
print("\n正在执行修正...")
|
||||
fixed = fix_discrepancies(session, discrepancies)
|
||||
print(f"✅ 已修正 {fixed} 个素材库计数")
|
||||
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,317 @@
|
||||
"""#1768 节奏模板多样化增强 — 单元测试。
|
||||
|
||||
覆盖:
|
||||
- 8 种预设模板完整性
|
||||
- MIN_CLIP_DURATION = 2.0(修复旧 1.0 覆盖 bug)
|
||||
- 最大片段时长钳制(<= 素材可用时长 × 90%)
|
||||
- 时长总和误差校验(<= 0.5s)
|
||||
- asset_durations 参数向后兼容(None/空 = 不钳制)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.voice_duration_planner import (
|
||||
MIN_CLIP_DURATION,
|
||||
RHYTHM_TEMPLATES,
|
||||
TOTAL_DURATION_TOLERANCE,
|
||||
adapt_template_length,
|
||||
get_rhythm_template,
|
||||
plan_clip_durations,
|
||||
total_output_duration,
|
||||
)
|
||||
|
||||
# ── 模板池 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRhythmTemplatesPool:
|
||||
"""#1768 模板池扩展到 8 种。"""
|
||||
|
||||
def test_template_count_is_8(self):
|
||||
assert len(RHYTHM_TEMPLATES) == 8
|
||||
|
||||
def test_all_templates_have_5_segments(self):
|
||||
for tpl in RHYTHM_TEMPLATES:
|
||||
assert len(tpl) == 5
|
||||
|
||||
def test_new_template_22112_exists(self):
|
||||
assert [2, 2, 1, 1, 2] in RHYTHM_TEMPLATES
|
||||
|
||||
def test_new_template_32121_exists(self):
|
||||
assert [3, 2, 1, 2, 1] in RHYTHM_TEMPLATES
|
||||
|
||||
def test_original_6_templates_preserved(self):
|
||||
originals = [
|
||||
[1, 1, 1, 1, 1],
|
||||
[2, 1, 3, 1, 2],
|
||||
[1, 2, 1, 2, 1],
|
||||
[3, 1, 1, 1, 3],
|
||||
[1, 1, 3, 2, 1],
|
||||
[2, 1, 1, 3, 1],
|
||||
]
|
||||
for orig in originals:
|
||||
assert orig in RHYTHM_TEMPLATES
|
||||
|
||||
def test_all_weights_positive(self):
|
||||
for tpl in RHYTHM_TEMPLATES:
|
||||
assert all(w > 0 for w in tpl)
|
||||
|
||||
def test_weight_sum_variety(self):
|
||||
"""不同模板权重和应不完全相同,确保节奏有差异。"""
|
||||
sums = {sum(t) for t in RHYTHM_TEMPLATES}
|
||||
assert len(sums) >= 3 # 至少有 3 种不同的权重和
|
||||
|
||||
|
||||
# ── 常量修复 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConstantsFixed:
|
||||
"""#1768 修复 MIN_CLIP_DURATION 从 1.0 回到 2.0。"""
|
||||
|
||||
def test_min_clip_duration_is_2(self):
|
||||
assert MIN_CLIP_DURATION == 2.0
|
||||
|
||||
def test_total_duration_tolerance_is_05(self):
|
||||
assert TOTAL_DURATION_TOLERANCE == 0.5
|
||||
|
||||
|
||||
# ── get_rhythm_template ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGetRhythmTemplate:
|
||||
def test_none_seed_returns_average(self):
|
||||
assert get_rhythm_template(None) == [1, 1, 1, 1, 1]
|
||||
|
||||
def test_same_seed_returns_same_template(self):
|
||||
for seed in [0, 42, 999, 123456]:
|
||||
t1 = get_rhythm_template(seed)
|
||||
t2 = get_rhythm_template(seed)
|
||||
assert t1 == t2
|
||||
|
||||
def test_different_seeds_can_yield_different_templates(self):
|
||||
"""大量 seed 应能命中多个不同模板。"""
|
||||
results = {tuple(get_rhythm_template(s)) for s in range(200)}
|
||||
assert len(results) >= 5 # 200 个 seed 至少命中 5 种模板
|
||||
|
||||
|
||||
# ── adapt_template_length ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAdaptTemplateLength:
|
||||
def test_exact_match(self):
|
||||
tpl = [2, 2, 1, 1, 2]
|
||||
assert adapt_template_length(tpl, 5) == tpl
|
||||
|
||||
def test_truncate(self):
|
||||
tpl = [2, 2, 1, 1, 2]
|
||||
assert adapt_template_length(tpl, 3) == [2, 2, 1]
|
||||
|
||||
def test_extend_cycles(self):
|
||||
tpl = [2, 2, 1, 1, 2]
|
||||
result = adapt_template_length(tpl, 8)
|
||||
assert len(result) == 8
|
||||
assert result == [2, 2, 1, 1, 2, 2, 2, 1]
|
||||
|
||||
def test_zero_clips(self):
|
||||
assert adapt_template_length([1, 1, 1], 0) == []
|
||||
|
||||
def test_negative_clips(self):
|
||||
assert adapt_template_length([1, 1, 1], -1) == []
|
||||
|
||||
|
||||
# ── plan_clip_durations 基础行为 ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPlanClipDurationsBasic:
|
||||
def test_invalid_inputs(self):
|
||||
assert plan_clip_durations(0, 30.0) == []
|
||||
assert plan_clip_durations(-1, 30.0) == []
|
||||
assert plan_clip_durations(5, 0.0) == []
|
||||
assert plan_clip_durations(5, -10.0) == []
|
||||
assert plan_clip_durations(5, "abc") == []
|
||||
|
||||
def test_average_distribution_no_transitions(self):
|
||||
result = plan_clip_durations(5, 30.0)
|
||||
assert len(result) == 5
|
||||
assert abs(sum(result) - 30.0) < 0.01
|
||||
|
||||
def test_all_segments_above_min(self):
|
||||
result = plan_clip_durations(5, 30.0, rhythm_template=[3, 1, 1, 1, 3])
|
||||
for dur in result:
|
||||
assert dur >= MIN_CLIP_DURATION
|
||||
|
||||
def test_with_rhythm_template(self):
|
||||
tpl = [2, 2, 1, 1, 2]
|
||||
result = plan_clip_durations(5, 30.0, rhythm_template=tpl)
|
||||
assert len(result) == 5
|
||||
# 权重和 = 8,每段应大致为 7.5, 7.5, 3.75, 3.75, 7.5
|
||||
assert result[0] > result[2] # 权重 2 > 权重 1
|
||||
assert abs(sum(result) - 30.0) < 0.5
|
||||
|
||||
def test_total_duration_matches_voice(self):
|
||||
"""成片净时长 ≈ 配音时长(无转场时完全等于)。"""
|
||||
for voice in [15.0, 30.0, 60.0, 120.0]:
|
||||
result = plan_clip_durations(5, voice)
|
||||
net = total_output_duration(result)
|
||||
assert abs(net - voice) <= TOTAL_DURATION_TOLERANCE
|
||||
|
||||
def test_with_transitions(self):
|
||||
"""有转场时成片净时长也应 ≈ 配音时长。"""
|
||||
effects = [None, "xfade", "xfade", "xfade", "xfade"]
|
||||
durations = [0.0, 1.0, 1.0, 1.0, 1.0]
|
||||
result = plan_clip_durations(5, 30.0, transition_effects=effects, transition_durations=durations)
|
||||
net = total_output_duration(result, effects, durations)
|
||||
assert abs(net - 30.0) <= TOTAL_DURATION_TOLERANCE
|
||||
|
||||
|
||||
# ── #1768 最小片段时长钳制 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestMinClipDurationClamp:
|
||||
def test_min_duration_2s_enforced(self):
|
||||
"""极端权重下,所有片段仍 >= 2.0s。"""
|
||||
tpl = [10, 1, 1, 1, 1]
|
||||
result = plan_clip_durations(5, 20.0, rhythm_template=tpl)
|
||||
for dur in result:
|
||||
assert dur >= 2.0, f"片段时长 {dur} < MIN_CLIP_DURATION(2.0)"
|
||||
|
||||
def test_short_voice_still_meets_minimum(self):
|
||||
"""配音极短时保底每段 MIN_CLIP_DURATION。"""
|
||||
result = plan_clip_durations(5, 3.0)
|
||||
for dur in result:
|
||||
assert dur >= MIN_CLIP_DURATION
|
||||
|
||||
|
||||
# ── #1768 最大片段时长钳制 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestMaxClipDurationClamp:
|
||||
def test_no_clamp_without_asset_durations(self):
|
||||
"""不传 asset_durations 时不做上限钳制(向后兼容)。"""
|
||||
tpl = [5, 1, 1, 1, 1]
|
||||
result = plan_clip_durations(5, 30.0, rhythm_template=tpl)
|
||||
# 第一段权重 5/9 * 30 = 16.67,不应被钳制
|
||||
assert result[0] > 10.0
|
||||
|
||||
def test_no_clamp_with_empty_asset_durations(self):
|
||||
"""asset_durations 为空列表时不做上限钳制。"""
|
||||
tpl = [5, 1, 1, 1, 1]
|
||||
result = plan_clip_durations(5, 30.0, rhythm_template=tpl, asset_durations=[])
|
||||
assert result[0] > 10.0
|
||||
|
||||
def test_clamp_respects_90_percent(self):
|
||||
"""有素材时长时,片段时长 <= 素材可用时长 × 90%。"""
|
||||
tpl = [5, 1, 1, 1, 1]
|
||||
# 素材只有第一段短(12s),90% = 10.8s
|
||||
asset_durs = [12.0, 60.0, 60.0, 60.0, 60.0]
|
||||
result = plan_clip_durations(5, 30.0, rhythm_template=tpl, asset_durations=asset_durs)
|
||||
max_allowed = 12.0 * 0.9
|
||||
assert result[0] <= max_allowed + 0.01, f"第一段 {result[0]} 超过 90% 上限 {max_allowed}"
|
||||
|
||||
def test_clamp_does_not_violate_min(self):
|
||||
"""素材极短时钳制不违反 MIN_CLIP_DURATION。"""
|
||||
# 素材 2.0s,90% = 1.8s < MIN(2.0),不应钳制到 1.8
|
||||
asset_durs = [2.0, 60.0, 60.0, 60.0, 60.0]
|
||||
result = plan_clip_durations(5, 30.0, asset_durations=asset_durs)
|
||||
for dur in result:
|
||||
assert dur >= MIN_CLIP_DURATION
|
||||
|
||||
def test_clamp_preserves_total(self):
|
||||
"""钳制后总时长仍应接近配音时长。"""
|
||||
asset_durs = [10.0, 60.0, 60.0, 60.0, 60.0]
|
||||
voice = 30.0
|
||||
result = plan_clip_durations(5, voice, asset_durations=asset_durs)
|
||||
net = total_output_duration(result)
|
||||
assert abs(net - voice) <= TOTAL_DURATION_TOLERANCE + 0.5 # 允许略多误差
|
||||
|
||||
def test_all_assets_short(self):
|
||||
"""所有素材都短时,钳制全部生效但不违反最小值。"""
|
||||
asset_durs = [8.0, 8.0, 8.0, 8.0, 8.0]
|
||||
result = plan_clip_durations(5, 30.0, asset_durations=asset_durs)
|
||||
for dur in result:
|
||||
assert dur >= MIN_CLIP_DURATION
|
||||
max_allowed = 8.0 * 0.9
|
||||
# 如果 max_allowed >= MIN_CLIP_DURATION 才钳制
|
||||
if max_allowed >= MIN_CLIP_DURATION:
|
||||
assert dur <= max_allowed + 0.1
|
||||
|
||||
|
||||
# ── #1768 时长总和误差校验 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTotalDurationTolerance:
|
||||
def test_no_transition_exact_match(self):
|
||||
"""无转场时总时长精确等于配音。"""
|
||||
result = plan_clip_durations(5, 25.0)
|
||||
assert abs(sum(result) - 25.0) < 0.01
|
||||
|
||||
def test_with_transition_within_tolerance(self):
|
||||
"""有转场时净时长在 0.5s 以内。"""
|
||||
effects = [None, "xfade", "fade", "xfade", "fade"]
|
||||
durations = [0.0, 0.8, 1.2, 0.5, 1.0]
|
||||
result = plan_clip_durations(5, 45.0, transition_effects=effects, transition_durations=durations)
|
||||
net = total_output_duration(result, effects, durations)
|
||||
assert abs(net - 45.0) <= TOTAL_DURATION_TOLERANCE
|
||||
|
||||
@pytest.mark.parametrize("voice", [10.0, 20.0, 30.0, 60.0, 120.0])
|
||||
def test_various_voice_durations(self, voice):
|
||||
result = plan_clip_durations(5, voice)
|
||||
net = total_output_duration(result)
|
||||
assert abs(net - voice) <= TOTAL_DURATION_TOLERANCE
|
||||
|
||||
@pytest.mark.parametrize("tpl", RHYTHM_TEMPLATES)
|
||||
def test_each_template_within_tolerance(self, tpl):
|
||||
"""每种模板分配的总时长都应在误差范围内。"""
|
||||
adapted = adapt_template_length(tpl, 5)
|
||||
result = plan_clip_durations(5, 30.0, rhythm_template=adapted)
|
||||
net = total_output_duration(result)
|
||||
assert (
|
||||
abs(net - 30.0) <= TOTAL_DURATION_TOLERANCE
|
||||
), f"模板 {tpl} 总时长误差 {abs(net - 30.0):.3f}s > {TOTAL_DURATION_TOLERANCE}s"
|
||||
|
||||
|
||||
# ── #1768 组合场景 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCombinedScenarios:
|
||||
def test_rhythm_plus_clamp_plus_tolerance(self):
|
||||
"""节奏模板 + 素材钳制 + 误差校验 同时生效。"""
|
||||
tpl = [3, 2, 1, 2, 1]
|
||||
effects = [None, "xfade", None, "xfade", None]
|
||||
tdurs = [0.0, 1.0, 0.0, 1.0, 0.0]
|
||||
asset_durs = [15.0, 60.0, 60.0, 60.0, 60.0]
|
||||
voice = 30.0
|
||||
|
||||
adapted = adapt_template_length(tpl, 5)
|
||||
result = plan_clip_durations(
|
||||
5,
|
||||
voice,
|
||||
transition_effects=effects,
|
||||
transition_durations=tdurs,
|
||||
rhythm_template=adapted,
|
||||
asset_durations=asset_durs,
|
||||
)
|
||||
|
||||
# 最小值保证
|
||||
for dur in result:
|
||||
assert dur >= MIN_CLIP_DURATION
|
||||
|
||||
# 最大值钳制(第一段 90% = 13.5)
|
||||
assert result[0] <= 15.0 * 0.9 + 0.1
|
||||
|
||||
# 总时长误差
|
||||
net = total_output_duration(result, effects, tdurs)
|
||||
assert abs(net - voice) <= TOTAL_DURATION_TOLERANCE + 0.5
|
||||
|
||||
def test_many_clips_with_cycling_template(self):
|
||||
"""片段数 > 模板长度时循环填充 + 钳制。"""
|
||||
tpl = [2, 2, 1, 1, 2]
|
||||
adapted = adapt_template_length(tpl, 8)
|
||||
assert len(adapted) == 8
|
||||
|
||||
asset_durs = [20.0] * 8
|
||||
result = plan_clip_durations(8, 40.0, rhythm_template=adapted, asset_durations=asset_durs)
|
||||
assert len(result) == 8
|
||||
for dur in result:
|
||||
assert dur >= MIN_CLIP_DURATION
|
||||
@@ -0,0 +1,305 @@
|
||||
"""Issue #1776: asset_libraries.asset_count 同步维护测试。
|
||||
|
||||
覆盖场景:
|
||||
1. 素材创建 → count +1
|
||||
2. 素材硬删除 → count -1
|
||||
3. 素材软删除(batch_delete)→ count -N
|
||||
4. 幂等上传(prepare 占位 + complete 复用)→ 不重复计数
|
||||
5. 重试场景(complete 重试)→ 不重复计数
|
||||
6. recount 方法修正计数
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_library_repository import SQLAlchemyAssetLibraryRepository
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetLibraryModel, AssetModel, Base
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def db_session():
|
||||
"""创建测试用 SQLite 内存数据库。"""
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
yield session
|
||||
session.close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def asset_repo(db_session):
|
||||
return SQLAlchemyAssetRepository(db_session)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def library_repo(db_session):
|
||||
return SQLAlchemyAssetLibraryRepository(db_session)
|
||||
|
||||
|
||||
def _make_library(project_id: str, kind: AssetLibraryKind = AssetLibraryKind.VIDEO) -> AssetLibrary:
|
||||
return AssetLibrary.create(project_id=project_id, name=f"测试{kind.value}库", kind=kind)
|
||||
|
||||
|
||||
def _make_asset(
|
||||
library_id: str,
|
||||
project_id: str,
|
||||
*,
|
||||
status: AssetStatus = AssetStatus.PROCESSING,
|
||||
file_size: int = 1024,
|
||||
file_hash: str = "",
|
||||
client_upload_id: str = "",
|
||||
) -> Asset:
|
||||
return Asset.create(
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
name=f"test_{uuid.uuid4().hex[:8]}.mp4",
|
||||
storage_key=f"uploads/{uuid.uuid4().hex[:8]}/test.mp4",
|
||||
mime_type="video/mp4",
|
||||
status=status,
|
||||
uploaded_by_user_id="test-user",
|
||||
file_hash=file_hash,
|
||||
client_upload_id=client_upload_id,
|
||||
file_size=file_size,
|
||||
)
|
||||
|
||||
|
||||
class TestAssetCountOnCreate:
|
||||
"""素材创建时计数递增。"""
|
||||
|
||||
def test_create_asset_increments_count(self, library_repo, asset_repo, db_session):
|
||||
"""创建一个素材 → count 从 0 变 1。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
assert library.asset_count == 0
|
||||
|
||||
asset = _make_asset(library.id, "project-1")
|
||||
asset_repo.create(asset)
|
||||
|
||||
# 重新查询验证计数
|
||||
updated_library = library_repo.get(library.id)
|
||||
assert updated_library.asset_count == 1
|
||||
assert updated_library.total_size == 1024
|
||||
|
||||
def test_create_multiple_assets_increments_count(self, library_repo, asset_repo, db_session):
|
||||
"""创建多个素材 → count 累加。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
|
||||
for _i in range(3):
|
||||
asset = _make_asset(library.id, "project-1", file_size=100 * (_i + 1))
|
||||
asset_repo.create(asset)
|
||||
|
||||
updated_library = library_repo.get(library.id)
|
||||
assert updated_library.asset_count == 3
|
||||
assert updated_library.total_size == 100 + 200 + 300
|
||||
|
||||
|
||||
class TestAssetCountOnDelete:
|
||||
"""素材删除时计数递减。"""
|
||||
|
||||
def test_hard_delete_decrements_count(self, library_repo, asset_repo, db_session):
|
||||
"""硬删除素材 → count -1。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
asset = asset_repo.create(_make_asset(library.id, "project-1"))
|
||||
assert library_repo.get(library.id).asset_count == 1
|
||||
|
||||
asset_repo.delete(asset.id)
|
||||
|
||||
assert library_repo.get(library.id).asset_count == 0
|
||||
|
||||
def test_hard_delete_already_deleted_no_change(self, library_repo, asset_repo, db_session):
|
||||
"""删除已删除的素材 → count 不变。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
asset = asset_repo.create(_make_asset(library.id, "project-1"))
|
||||
# 先软删除(count 已经 -1)
|
||||
asset_repo.batch_delete([asset.id])
|
||||
assert library_repo.get(library.id).asset_count == 0
|
||||
|
||||
# 再硬删除(不应再 -1)
|
||||
asset_repo.delete(asset.id)
|
||||
assert library_repo.get(library.id).asset_count == 0
|
||||
|
||||
def test_batch_delete_decrements_count(self, library_repo, asset_repo, db_session):
|
||||
"""批量软删除 → count -N。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
asset_ids = []
|
||||
for _i in range(5):
|
||||
asset = asset_repo.create(_make_asset(library.id, "project-1", file_size=200))
|
||||
asset_ids.append(asset.id)
|
||||
assert library_repo.get(library.id).asset_count == 5
|
||||
|
||||
# 删除 3 个
|
||||
deleted_count = asset_repo.batch_delete(asset_ids[:3])
|
||||
assert deleted_count == 3
|
||||
assert library_repo.get(library.id).asset_count == 2
|
||||
assert library_repo.get(library.id).total_size == 200 * 2
|
||||
|
||||
def test_batch_delete_skips_already_deleted(self, library_repo, asset_repo, db_session):
|
||||
"""批量删除已删除的素材 → count 不变。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
asset_ids = []
|
||||
for _i in range(3):
|
||||
asset = asset_repo.create(_make_asset(library.id, "project-1"))
|
||||
asset_ids.append(asset.id)
|
||||
assert library_repo.get(library.id).asset_count == 3
|
||||
|
||||
# 先删除 2 个
|
||||
asset_repo.batch_delete(asset_ids[:2])
|
||||
assert library_repo.get(library.id).asset_count == 1
|
||||
|
||||
# 再删除同样的 2 个(应被跳过)
|
||||
deleted_count = asset_repo.batch_delete(asset_ids[:2])
|
||||
assert deleted_count == 0
|
||||
assert library_repo.get(library.id).asset_count == 1
|
||||
|
||||
def test_count_never_negative(self, library_repo, asset_repo, db_session):
|
||||
"""计数下限为 0,不会出现负数。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
# 手动设置计数为 0
|
||||
library.asset_count = 0
|
||||
library_repo.update(library)
|
||||
|
||||
# 尝试递减(通过直接调用 decrement)
|
||||
library_repo.decrement_asset_count(library.id, count_delta=5)
|
||||
db_session.commit()
|
||||
|
||||
updated = library_repo.get(library.id)
|
||||
assert updated.asset_count == 0
|
||||
|
||||
|
||||
class TestIdempotentUpload:
|
||||
"""幂等上传场景:不重复计数。"""
|
||||
|
||||
def test_prepare_then_complete_no_double_count(self, library_repo, asset_repo, db_session):
|
||||
"""prepare 创建占位 + complete 复用占位 → count 只 +1。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
|
||||
# prepare 阶段:创建占位
|
||||
placeholder = _make_asset(
|
||||
library.id,
|
||||
"project-1",
|
||||
status=AssetStatus.PROCESSING,
|
||||
client_upload_id="upload-token-123",
|
||||
)
|
||||
asset_repo.create(placeholder)
|
||||
assert library_repo.get(library.id).asset_count == 1
|
||||
|
||||
# complete 阶段:查找已有占位并复用(通过 client_upload_id)
|
||||
existing = asset_repo.find_by_library_and_client_upload_id(
|
||||
library_id=library.id,
|
||||
client_upload_id="upload-token-123",
|
||||
)
|
||||
assert existing is not None
|
||||
# 复用占位,不创建新记录 → count 不变
|
||||
assert library_repo.get(library.id).asset_count == 1
|
||||
|
||||
def test_complete_retry_no_double_count(self, library_repo, asset_repo, db_session):
|
||||
"""complete 重试(通过 file_hash 去重)→ count 只 +1。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
|
||||
# 第一次 complete:创建素材
|
||||
asset1 = _make_asset(
|
||||
library.id,
|
||||
"project-1",
|
||||
file_hash="hash-abc-123",
|
||||
)
|
||||
asset_repo.create(asset1)
|
||||
assert library_repo.get(library.id).asset_count == 1
|
||||
|
||||
# 重试 complete:通过 file_hash 查找已有
|
||||
existing = asset_repo.find_by_library_and_file_hash(
|
||||
library_id=library.id,
|
||||
file_hash="hash-abc-123",
|
||||
)
|
||||
assert existing is not None
|
||||
assert existing.id == asset1.id
|
||||
# 不创建新记录 → count 不变
|
||||
assert library_repo.get(library.id).asset_count == 1
|
||||
|
||||
|
||||
class TestRecountAssets:
|
||||
"""recount_assets 方法修正计数。"""
|
||||
|
||||
def test_recount_fixes_drift(self, library_repo, asset_repo, db_session):
|
||||
"""计数漂移后,recount 能修正。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
# 创建 3 个素材
|
||||
for _ in range(3):
|
||||
asset_repo.create(_make_asset(library.id, "project-1"))
|
||||
assert library_repo.get(library.id).asset_count == 3
|
||||
|
||||
# 手动破坏计数(模拟历史数据问题)
|
||||
db_session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library.id).update(
|
||||
{AssetLibraryModel.asset_count: 999, AssetLibraryModel.total_size: 999999}
|
||||
)
|
||||
db_session.commit()
|
||||
assert library_repo.get(library.id).asset_count == 999
|
||||
|
||||
# recount 修正
|
||||
actual = library_repo.recount_assets(library.id)
|
||||
db_session.commit()
|
||||
|
||||
assert actual == 3
|
||||
assert library_repo.get(library.id).asset_count == 3
|
||||
assert library_repo.get(library.id).total_size == 1024 * 3
|
||||
|
||||
def test_recount_excludes_deleted(self, library_repo, asset_repo, db_session):
|
||||
"""recount 排除已删除素材。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
assets = []
|
||||
for _ in range(5):
|
||||
asset = asset_repo.create(_make_asset(library.id, "project-1"))
|
||||
assets.append(asset)
|
||||
assert library_repo.get(library.id).asset_count == 5
|
||||
|
||||
# 软删除 2 个
|
||||
asset_repo.batch_delete([assets[0].id, assets[1].id])
|
||||
assert library_repo.get(library.id).asset_count == 3
|
||||
|
||||
# 破坏计数
|
||||
db_session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library.id).update(
|
||||
{AssetLibraryModel.asset_count: 100}
|
||||
)
|
||||
db_session.commit()
|
||||
|
||||
# recount 应排除 deleted
|
||||
actual = library_repo.recount_assets(library.id)
|
||||
db_session.commit()
|
||||
|
||||
assert actual == 3
|
||||
assert library_repo.get(library.id).asset_count == 3
|
||||
|
||||
|
||||
class TestConcurrentSafety:
|
||||
"""并发安全测试(SQLite 模拟有限并发)。"""
|
||||
|
||||
def test_increment_is_atomic(self, library_repo, db_session):
|
||||
"""increment_asset_count 使用 SQL 级 UPDATE,并发安全。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
|
||||
# 多次递增
|
||||
for _ in range(10):
|
||||
library_repo.increment_asset_count(library.id, count_delta=1, size_delta=100)
|
||||
db_session.commit()
|
||||
|
||||
updated = library_repo.get(library.id)
|
||||
assert updated.asset_count == 10
|
||||
assert updated.total_size == 1000
|
||||
|
||||
def test_decrement_with_floor_zero(self, library_repo, db_session):
|
||||
"""decrement_asset_count 下限为 0。"""
|
||||
library = library_repo.create(_make_library("project-1"))
|
||||
library.asset_count = 3
|
||||
library_repo.update(library)
|
||||
|
||||
# 尝试递减 10 次
|
||||
for _ in range(10):
|
||||
library_repo.decrement_asset_count(library.id, count_delta=1)
|
||||
db_session.commit()
|
||||
|
||||
updated = library_repo.get(library.id)
|
||||
assert updated.asset_count == 0
|
||||
@@ -0,0 +1,281 @@
|
||||
"""BGM 池差异化分配单元测试 (Issue #1767).
|
||||
|
||||
覆盖:
|
||||
- BGM 池定义(10 首,覆盖 4 种风格)
|
||||
- 风格匹配:get_bgm_pool_candidates 按 mood 筛选
|
||||
- 变体分配:select_bgm_from_pool 基于 seed 可复现选择
|
||||
- 批量分配:allocate_bgm_pool_for_variants 确保不同变体不同 BGM
|
||||
- 段落差异化:generate_bgm_segment_offset 生成不同偏移
|
||||
- 音量微调:generate_bgm_volume_adjust ±3dB
|
||||
- 边界条件:空配置/未启用 BGM/未知 mood
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.bgm_pool import (
|
||||
BGM_POOL,
|
||||
STYLE_TO_MOOD,
|
||||
BGMPoolEntry,
|
||||
allocate_bgm_pool_for_variants,
|
||||
generate_bgm_segment_offset,
|
||||
generate_bgm_volume_adjust,
|
||||
get_bgm_pool_candidates,
|
||||
select_bgm_from_pool,
|
||||
)
|
||||
|
||||
|
||||
class TestBGMPoolDefinition:
|
||||
"""BGM 池定义测试。"""
|
||||
|
||||
def test_pool_has_at_least_10_entries(self):
|
||||
"""BGM 池至少 10 首(满足 5-10 首需求)。"""
|
||||
assert len(BGM_POOL) >= 10
|
||||
|
||||
def test_all_entries_have_required_fields(self):
|
||||
"""每条 BGM 池条目都有 id/preset_id/mood/duration。"""
|
||||
for entry in BGM_POOL:
|
||||
assert entry.id, "Missing id"
|
||||
assert entry.preset_id, f"Missing preset_id in {entry.id}"
|
||||
assert entry.mood, f"Missing mood in {entry.id}"
|
||||
assert entry.duration > 0, f"Duration must be > 0 in {entry.id}"
|
||||
|
||||
def test_pool_covers_multiple_moods(self):
|
||||
"""池覆盖至少 3 种不同 mood。"""
|
||||
moods = {e.mood for e in BGM_POOL}
|
||||
assert len(moods) >= 3, f"Expected >= 3 moods, got {moods}"
|
||||
|
||||
def test_each_mood_has_at_least_2_entries(self):
|
||||
"""每种 mood 至少有 2 首(保证差异化有意义)。"""
|
||||
mood_counts: dict[str, int] = {}
|
||||
for entry in BGM_POOL:
|
||||
mood_counts[entry.mood] = mood_counts.get(entry.mood, 0) + 1
|
||||
for mood, count in mood_counts.items():
|
||||
assert count >= 2, f"Mood '{mood}' only has {count} entries (need >= 2)"
|
||||
|
||||
|
||||
class TestGetBGMPoolCandidates:
|
||||
"""风格匹配测试。"""
|
||||
|
||||
def test_no_mood_returns_full_pool(self):
|
||||
"""不指定 mood 时返回全池。"""
|
||||
candidates = get_bgm_pool_candidates(None)
|
||||
assert len(candidates) == len(BGM_POOL)
|
||||
|
||||
def test_matching_mood_filters(self):
|
||||
"""指定已知 mood 时返回同 mood 条目。"""
|
||||
candidates = get_bgm_pool_candidates("upbeat")
|
||||
assert all(c.mood == "upbeat" for c in candidates)
|
||||
assert len(candidates) >= 2
|
||||
|
||||
def test_unknown_mood_returns_full_pool(self):
|
||||
"""未知 mood 降级返回全池。"""
|
||||
candidates = get_bgm_pool_candidates("nonexistent_style")
|
||||
# 如果没有匹配到同 mood 的(>=2),降级全池
|
||||
assert len(candidates) >= 2
|
||||
|
||||
def test_style_to_mood_mapping(self):
|
||||
"""STYLE_TO_MOOD 覆盖所有 preset_bgm 的 style。"""
|
||||
assert "upbeat" in STYLE_TO_MOOD
|
||||
assert "relax" in STYLE_TO_MOOD
|
||||
assert "tech" in STYLE_TO_MOOD
|
||||
assert "commerce" in STYLE_TO_MOOD
|
||||
|
||||
|
||||
class TestSelectBGMPool:
|
||||
"""变体 BGM 选择测试。"""
|
||||
|
||||
def test_same_seed_same_result(self):
|
||||
"""相同 seed 返回相同 BGM(可复现)。"""
|
||||
result1 = select_bgm_from_pool(42)
|
||||
result2 = select_bgm_from_pool(42)
|
||||
assert result1.id == result2.id
|
||||
|
||||
def test_different_seeds_may_differ(self):
|
||||
"""不同 seed 可能返回不同 BGM。"""
|
||||
results = set()
|
||||
for seed in range(50):
|
||||
entry = select_bgm_from_pool(seed)
|
||||
results.add(entry.id)
|
||||
assert len(results) >= 3, "Expected >= 3 different BGMs from 50 seeds"
|
||||
|
||||
def test_respects_candidates_filter(self):
|
||||
"""传入候选池时只从中选择。"""
|
||||
candidates = [e for e in BGM_POOL if e.mood == "tech"]
|
||||
for seed in range(20):
|
||||
entry = select_bgm_from_pool(seed, candidates)
|
||||
assert entry.mood == "tech"
|
||||
|
||||
|
||||
class TestBGMSegmentOffset:
|
||||
"""段落差异化测试(策略二)。"""
|
||||
|
||||
def test_offset_non_negative(self):
|
||||
"""偏移 >= 0。"""
|
||||
for seed in range(50):
|
||||
offset = generate_bgm_segment_offset(seed, 120.0)
|
||||
assert offset >= 0.0
|
||||
|
||||
def test_offset_bounded(self):
|
||||
"""偏移 <= min(30s, duration * 0.3)。"""
|
||||
for seed in range(50):
|
||||
duration = 100.0
|
||||
max_expected = min(30.0, duration * 0.3)
|
||||
offset = generate_bgm_segment_offset(seed, duration)
|
||||
assert offset <= max_expected + 0.1 # 容差
|
||||
|
||||
def test_different_seeds_different_offsets(self):
|
||||
"""不同 seed 产生不同偏移(统计验证)。"""
|
||||
offsets = set()
|
||||
for seed in range(30):
|
||||
offsets.add(generate_bgm_segment_offset(seed, 120.0))
|
||||
assert len(offsets) >= 3, "Expected >= 3 distinct offsets"
|
||||
|
||||
def test_offset_is_quantized(self):
|
||||
"""偏移是 5s 的整数倍。"""
|
||||
for seed in range(20):
|
||||
offset = generate_bgm_segment_offset(seed, 120.0)
|
||||
assert offset % 5.0 == 0.0
|
||||
|
||||
def test_short_bgm_zero_offset(self):
|
||||
"""极短 BGM 偏移为 0。"""
|
||||
offset = generate_bgm_segment_offset(42, 5.0)
|
||||
# max_offset = min(30, 5*0.3) = 1.5, steps = int(1.5//5) = 0 → return 0
|
||||
assert offset == 0.0
|
||||
|
||||
def test_offset_different_from_bgm_selection(self):
|
||||
"""偏移的 seed 序列与 BGM 选择的 seed 序列不同(加素数偏移)。"""
|
||||
# 同一 seed,偏移和 BGM 选择应该独立
|
||||
bgm = select_bgm_from_pool(42)
|
||||
offset = generate_bgm_segment_offset(42, 120.0)
|
||||
# 只是验证能正常运行,不直接断言独立性(统计测试需要大样本)
|
||||
assert isinstance(offset, float)
|
||||
|
||||
|
||||
class TestBGMVolumeAdjust:
|
||||
"""音量微调测试(策略三)。"""
|
||||
|
||||
def test_volume_adjust_in_range(self):
|
||||
"""音量调整在 -3 ~ +3 dB 范围内。"""
|
||||
for seed in range(50):
|
||||
adj = generate_bgm_volume_adjust(seed)
|
||||
assert -3.0 <= adj <= 3.0
|
||||
assert adj == int(adj) # 整数 dB 步进
|
||||
|
||||
def test_different_seeds_different_volumes(self):
|
||||
"""不同 seed 产生不同音量调整值。"""
|
||||
values = set()
|
||||
for seed in range(50):
|
||||
values.add(generate_bgm_volume_adjust(seed))
|
||||
assert len(values) >= 3, "Expected >= 3 distinct volume values"
|
||||
|
||||
def test_includes_zero(self):
|
||||
"""音量调整值集合包含 0(不变)。"""
|
||||
values = {generate_bgm_volume_adjust(seed) for seed in range(100)}
|
||||
assert 0.0 in values
|
||||
|
||||
|
||||
class TestAllocateBGMPoolForVariants:
|
||||
"""批量分配入口测试。"""
|
||||
|
||||
def test_disabled_bgm_returns_empty(self):
|
||||
"""源 BGM 未启用时返回空列表。"""
|
||||
result = allocate_bgm_pool_for_variants({"enabled": False}, [1, 2, 3])
|
||||
assert result == []
|
||||
|
||||
def test_empty_config_returns_empty(self):
|
||||
"""空配置返回空列表。"""
|
||||
result = allocate_bgm_pool_for_variants({}, [1, 2, 3])
|
||||
assert result == []
|
||||
|
||||
def test_empty_seeds_returns_empty(self):
|
||||
"""无变体时返回空列表。"""
|
||||
result = allocate_bgm_pool_for_variants({"enabled": True, "preset_id": "bgm_upbeat_001"}, [])
|
||||
assert result == []
|
||||
|
||||
def test_returns_correct_count(self):
|
||||
"""返回与 variant_seeds 等长的列表。"""
|
||||
config = {"enabled": True, "preset_id": "bgm_upbeat_001"}
|
||||
seeds = [100, 200, 300, 400]
|
||||
result = allocate_bgm_pool_for_variants(config, seeds)
|
||||
assert len(result) == 4
|
||||
|
||||
def test_each_entry_has_required_keys(self):
|
||||
"""每项都包含必要字段。"""
|
||||
config = {"enabled": True, "preset_id": "bgm_upbeat_001"}
|
||||
seeds = [100, 200, 300]
|
||||
result = allocate_bgm_pool_for_variants(config, seeds)
|
||||
for entry in result:
|
||||
assert "preset_id" in entry
|
||||
assert "audio_offset" in entry
|
||||
assert "volume_adjust_db" in entry
|
||||
assert "bgm_pool_entry_id" in entry
|
||||
assert "bgm_pool_mood" in entry
|
||||
|
||||
def test_batch_3_variants_at_least_2_different_bgm(self):
|
||||
"""批量 3 个变体,至少 2 个不同 BGM。"""
|
||||
config = {"enabled": True, "preset_id": "bgm_upbeat_001"}
|
||||
seeds = [100, 200, 300]
|
||||
result = allocate_bgm_pool_for_variants(config, seeds)
|
||||
bgm_ids = {r["bgm_pool_entry_id"] for r in result}
|
||||
assert len(bgm_ids) >= 2, f"Expected >= 2 different BGMs, got {bgm_ids}"
|
||||
|
||||
def test_style_matching_with_preset_id(self):
|
||||
"""源 BGM 有 preset_id 时按风格筛选。"""
|
||||
config = {"enabled": True, "preset_id": "bgm_tech_001"} # tech 风格
|
||||
seeds = [100, 200, 300]
|
||||
result = allocate_bgm_pool_for_variants(config, seeds)
|
||||
# tech mood 至少有 2 首,所以应该筛选到 tech
|
||||
for entry in result:
|
||||
assert entry["bgm_pool_mood"] == "tech"
|
||||
|
||||
def test_reproducible_with_same_seeds(self):
|
||||
"""相同 seeds 产生相同分配(可复现)。"""
|
||||
config = {"enabled": True, "preset_id": "bgm_upbeat_001"}
|
||||
seeds = [42, 100, 200]
|
||||
result1 = allocate_bgm_pool_for_variants(config, seeds)
|
||||
result2 = allocate_bgm_pool_for_variants(config, seeds)
|
||||
for r1, r2 in zip(result1, result2, strict=False):
|
||||
assert r1["bgm_pool_entry_id"] == r2["bgm_pool_entry_id"]
|
||||
assert r1["audio_offset"] == r2["audio_offset"]
|
||||
assert r1["volume_adjust_db"] == r2["volume_adjust_db"]
|
||||
|
||||
|
||||
class TestIssue1767Acceptance:
|
||||
"""Issue #1767 验收测试。"""
|
||||
|
||||
def test_batch_3_videos_bgm_different_or_segment_different(self):
|
||||
"""批量生成 3 个视频,BGM 不同或起始段落不同。"""
|
||||
config = {"enabled": True, "preset_id": "bgm_upbeat_001"}
|
||||
seeds = [100, 200, 300]
|
||||
result = allocate_bgm_pool_for_variants(config, seeds)
|
||||
|
||||
# 检查:BGM 不同 或 段落偏移不同
|
||||
unique_combos = set()
|
||||
for r in result:
|
||||
combo = (r["bgm_pool_entry_id"], r["audio_offset"])
|
||||
unique_combos.add(combo)
|
||||
|
||||
assert len(unique_combos) >= 2, f"Expected >= 2 unique (bgm, offset) combos, got {unique_combos}"
|
||||
|
||||
def test_bgm_volume_micro_adjust_doesnt_affect_voice_clarity(self):
|
||||
"""BGM 音量微调在 ±3dB 内,不影响配音清晰度。"""
|
||||
config = {"enabled": True, "preset_id": "bgm_upbeat_001"}
|
||||
seeds = [100, 200, 300]
|
||||
result = allocate_bgm_pool_for_variants(config, seeds)
|
||||
|
||||
for r in result:
|
||||
# ±3dB 是安全的微调范围,不会让 BGM 盖过配音
|
||||
assert abs(r["volume_adjust_db"]) <= 3.0
|
||||
|
||||
def test_single_video_mode_unaffected(self):
|
||||
"""单视频模式不受影响(不走池分配)。"""
|
||||
# 单视频不调用 allocate_bgm_pool_for_variants
|
||||
# 只要不主动调用,就不会改变行为
|
||||
# 这个测试验证函数签名和行为不会意外影响单视频
|
||||
config = {"enabled": True, "preset_id": "bgm_upbeat_001"}
|
||||
result = allocate_bgm_pool_for_variants(config, [42]) # 单变体
|
||||
assert len(result) == 1
|
||||
# 单项分配仍然有完整配置(不影响功能,只是差异化)
|
||||
assert "preset_id" in result[0]
|
||||
@@ -0,0 +1,177 @@
|
||||
"""默认项目/默认素材库幂等化测试(Issue #1775)。
|
||||
|
||||
覆盖:
|
||||
- get_or_create_default_project:同用户幂等、不同用户独立、重复调用返回同一个
|
||||
- get_or_create_default_library:同项目同 kind 幂等、IntegrityError 后重查
|
||||
- Project.is_default 字段传递
|
||||
- ensure-default-context 组合逻辑(用内存仓储)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.adapters.in_memory.asset_library_repository import InMemoryAssetLibraryRepository
|
||||
from packages.adapters.in_memory.project_repository import InMemoryProjectRepository
|
||||
from packages.adapters.sqlalchemy_impl.asset_library_repository import (
|
||||
SQLAlchemyAssetLibraryRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.project_repository import SQLAlchemyProjectRepository
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind, Project
|
||||
|
||||
|
||||
class TestDefaultProjectIdempotent:
|
||||
"""默认项目幂等。"""
|
||||
|
||||
def test_first_call_creates_default(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
project = repo.get_or_create_default_project("user-1")
|
||||
assert project.owner_user_id == "user-1"
|
||||
assert project.is_default is True
|
||||
assert project.name == "默认项目"
|
||||
|
||||
def test_second_call_returns_same_project(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
p1 = repo.get_or_create_default_project("user-1")
|
||||
p2 = repo.get_or_create_default_project("user-1")
|
||||
assert p1.id == p2.id
|
||||
|
||||
def test_different_users_independent(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
p1 = repo.get_or_create_default_project("user-1")
|
||||
p2 = repo.get_or_create_default_project("user-2")
|
||||
assert p1.id != p2.id
|
||||
assert p1.owner_user_id == "user-1"
|
||||
assert p2.owner_user_id == "user-2"
|
||||
|
||||
def test_find_default_by_owner(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
created = repo.get_or_create_default_project("user-1")
|
||||
found = repo.find_default_by_owner("user-1")
|
||||
assert found is not None
|
||||
assert found.id == created.id
|
||||
|
||||
def test_find_default_returns_none_when_no_default(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
# 手动建一个非默认项目
|
||||
normal = Project.create(owner_user_id="user-1", name="普通项目")
|
||||
normal.is_default = False
|
||||
repo.save(normal)
|
||||
assert repo.find_default_by_owner("user-1") is None
|
||||
|
||||
def test_repeated_calls_after_creation_return_same(self):
|
||||
"""先创建默认项目后,后续多次调用均返回已有项目(测试 find 路径)。
|
||||
注:真正的并发保护依赖 PostgreSQL partial unique index,
|
||||
InMemory 仓储不做并发测试(无 DB 约束),并发场景由
|
||||
TestSqlRepoIntegrityErrorRecovery 通过 SQLAlchemy + SQLite 验证。
|
||||
"""
|
||||
repo = InMemoryProjectRepository()
|
||||
first = repo.get_or_create_default_project("user-repeat")
|
||||
for _ in range(9):
|
||||
again = repo.get_or_create_default_project("user-repeat")
|
||||
assert again.id == first.id
|
||||
defaults = [p for p in repo.find_by_owner_user_id("user-repeat") if p.is_default]
|
||||
assert len(defaults) == 1
|
||||
|
||||
|
||||
class TestDefaultLibraryIdempotent:
|
||||
"""默认素材库幂等。"""
|
||||
|
||||
def test_first_call_creates(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
lib = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
assert lib.project_id == "proj-1"
|
||||
assert lib.kind == AssetLibraryKind.VIDEO
|
||||
|
||||
def test_second_call_returns_same(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
l1 = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
l2 = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
assert l1.id == l2.id
|
||||
|
||||
def test_different_kinds_independent(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
v = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
a = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VOICE)
|
||||
i = repo.get_or_create_default_library("proj-1", AssetLibraryKind.IMAGE)
|
||||
assert len({v.id, a.id, i.id}) == 3
|
||||
|
||||
def test_different_projects_independent(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
l1 = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
l2 = repo.get_or_create_default_library("proj-2", AssetLibraryKind.VIDEO)
|
||||
assert l1.id != l2.id
|
||||
|
||||
def test_repeated_calls_after_creation_return_same(self):
|
||||
"""先创建后多次调用均返回同一素材库(测试 find 路径)。"""
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
first = repo.get_or_create_default_library("proj-repeat", AssetLibraryKind.VOICE)
|
||||
for _ in range(9):
|
||||
again = repo.get_or_create_default_library("proj-repeat", AssetLibraryKind.VOICE)
|
||||
assert again.id == first.id
|
||||
|
||||
|
||||
class TestProjectIsDefaultField:
|
||||
"""Project.is_default 字段语义。"""
|
||||
|
||||
def test_create_non_default_by_default(self):
|
||||
p = Project.create(owner_user_id="u", name="普通项目")
|
||||
assert p.is_default is False
|
||||
|
||||
def test_create_default(self):
|
||||
p = Project.create(owner_user_id="u", name="默认项目", is_default=True)
|
||||
assert p.is_default is True
|
||||
|
||||
|
||||
class TestSqlRepoIntegrityErrorRecovery:
|
||||
"""SQLAlchemy 仓储:唯一约束冲突时回滚重查,返回已有记录(不报 500)。"""
|
||||
|
||||
def test_project_integrity_error_returns_existing(self):
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
session = MagicMock()
|
||||
# commit 第一次抛 IntegrityError(并发冲突),回滚后查询返回已有项目
|
||||
existing_model = MagicMock()
|
||||
existing_model.id = "existing-id"
|
||||
existing_model.owner_user_id = "user-1"
|
||||
existing_model.name = "默认项目"
|
||||
existing_model.description = ""
|
||||
existing_model.shared_users = []
|
||||
existing_model.is_default = True
|
||||
existing_model.created_at = None
|
||||
|
||||
session.commit.side_effect = [IntegrityError("stmt", {}, Exception("dup")), None]
|
||||
# 第一次 query(快速路径 find_default)返回 None;rollback 后第二次返回 existing
|
||||
session.query.return_value.filter.return_value.first.side_effect = [None, existing_model]
|
||||
|
||||
repo = SQLAlchemyProjectRepository(session)
|
||||
result = repo.get_or_create_default_project("user-1")
|
||||
|
||||
assert result.id == "existing-id"
|
||||
session.rollback.assert_called_once()
|
||||
|
||||
def test_library_integrity_error_returns_existing(self):
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
session = MagicMock()
|
||||
existing_model = MagicMock()
|
||||
existing_model.id = "lib-existing"
|
||||
existing_model.project_id = "proj-1"
|
||||
existing_model.name = "视频素材库"
|
||||
existing_model.kind = "video"
|
||||
existing_model.asset_count = 0
|
||||
existing_model.total_size = 0
|
||||
existing_model.created_at = None
|
||||
existing_model.updated_at = None
|
||||
|
||||
session.commit.side_effect = [IntegrityError("stmt", {}, Exception("dup")), None]
|
||||
session.query.return_value.filter.return_value.first.side_effect = [None, existing_model]
|
||||
|
||||
repo = SQLAlchemyAssetLibraryRepository(session)
|
||||
result = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
|
||||
assert result.id == "lib-existing"
|
||||
assert result.kind == AssetLibraryKind.VIDEO
|
||||
session.rollback.assert_called_once()
|
||||
@@ -239,8 +239,8 @@ class TestEditorClipsBySegments:
|
||||
assert not hasattr(mock_plan_svc, "create_clip") or not mock_plan_svc.create_clip.called
|
||||
|
||||
@patch("app.api.routes.templates_editor.clips.get_storage_service")
|
||||
def test_no_segments_raises_400(self, mock_storage):
|
||||
"""模板没有 segment 配置时返回 400。"""
|
||||
def test_no_segments_raises_422(self, mock_storage):
|
||||
"""模板存在但未配置片段时返回 422(配置错误,与 404 区分)。"""
|
||||
from app.api.routes.templates_editor.clips import (
|
||||
create_clips_from_assets_editor,
|
||||
)
|
||||
@@ -264,11 +264,44 @@ class TestEditorClipsBySegments:
|
||||
current_user=_make_auth_user(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "片段配置" in exc_info.value.detail
|
||||
assert exc_info.value.status_code == 422
|
||||
assert "片段" in exc_info.value.detail
|
||||
# 不应调用替换方法
|
||||
mock_plan_svc.replace_all_clips_transactional.assert_not_called()
|
||||
|
||||
@patch("app.api.routes.templates_editor.clips.get_storage_service")
|
||||
def test_template_not_found_raises_404(self, mock_storage):
|
||||
"""模板不存在/已删除/无权限(服务层抛 TemplateNotFoundError)时返回 404。"""
|
||||
from app.api.routes.templates_editor.clips import (
|
||||
create_clips_from_assets_editor,
|
||||
)
|
||||
from app.api.routes.templates_editor.schemas import ClipsFromAssetsRequest
|
||||
from app.services.edit_template_service import TemplateNotFoundError
|
||||
|
||||
mock_plan_svc = _make_plan_svc()
|
||||
mock_asset_repo = MagicMock()
|
||||
|
||||
body = ClipsFromAssetsRequest(asset_ids=["a1"])
|
||||
|
||||
with patch(
|
||||
"app.api.routes.templates_editor.clips._get_template_segments",
|
||||
side_effect=TemplateNotFoundError("tpl-missing"),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
create_clips_from_assets_editor(
|
||||
template_id="tpl-missing",
|
||||
body=body,
|
||||
background_tasks=MagicMock(),
|
||||
plan_id=TEST_PLAN_ID,
|
||||
services=(MagicMock(), mock_plan_svc),
|
||||
asset_repo=mock_asset_repo,
|
||||
db=MagicMock(),
|
||||
current_user=_make_auth_user(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
mock_plan_svc.replace_all_clips_transactional.assert_not_called()
|
||||
|
||||
|
||||
class TestEditorClipsDurationAndStartTime:
|
||||
"""测试素材时长获取、clip duration 缩短、start_time 传入。"""
|
||||
|
||||
@@ -1,13 +1,17 @@
|
||||
"""_get_template_segments 回退路径测试.
|
||||
"""模板片段配置读取路径测试(#1774).
|
||||
|
||||
验证三级回退链:
|
||||
1. 新模板系统(tpl_svc.list_clip_configs)正常 → 直接返回
|
||||
2. 新模板系统主表不存在(ValueError)→ 直接查 template_clip_configs 表兜底
|
||||
3. 直接查表也失败 → 回退旧模板系统(template_segments)
|
||||
4. 全部失败 → 返回空列表
|
||||
收敛后模板读取走单一数据源,不再有"新表抛异常→降级查旧表→再降级查 segments"
|
||||
的异常控制流:
|
||||
|
||||
覆盖 P0 修复:自建模板在 edit_templates 主表不存在但在 template_clip_configs 有记录时,
|
||||
from-assets 流程不再 400。
|
||||
- ``EditTemplateService.list_clip_configs_for_editor`` 显式判定模板归属/存在性:
|
||||
1. 用户自建模板在旧表 ``templates``(归属 user_id,is_active=True)→ 直接读
|
||||
``template_clip_configs``;
|
||||
2. 全局模板在新表 ``edit_templates``(无 user_id)→ 直接读 ``template_clip_configs``;
|
||||
3. 两表都没有 → 抛 ``TemplateNotFoundError``(路由层映射 404)。
|
||||
- ``_get_template_segments`` 仅做配置→(order, min, max) 的映射与排序,
|
||||
模板存在但无配置返回空列表(路由层映射 422)。
|
||||
|
||||
使用真实 SQLite 内存库 + 真实仓储,验证端到端读路径不抛 ``ValueError: 模板不存在``。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -15,21 +19,198 @@ from __future__ import annotations
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, PropertyMock
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
import pytest # noqa: E402
|
||||
from sqlalchemy import create_engine # noqa: E402
|
||||
from sqlalchemy.orm import sessionmaker # noqa: E402
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import ( # noqa: E402
|
||||
Base,
|
||||
EditTemplateModel,
|
||||
TemplateClipConfigModel,
|
||||
TemplateModel,
|
||||
)
|
||||
|
||||
TEST_TEMPLATE_ID = "tmpl-orphan-001"
|
||||
DEFAULT_DUR = 5.0 # _DEFAULT_EDITOR_CLIP_DURATION
|
||||
USER_ID = "user-001"
|
||||
OTHER_USER_ID = "user-002"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 真实内存 DB fixture
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_session():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
return sessionmaker(bind=engine)()
|
||||
|
||||
|
||||
def _seed_legacy_template(session, template_id: str, user_id: str, *, active: bool = True, clip_count: int = 3):
|
||||
"""创建旧表 templates 模板(+ template_clip_configs 片段配置)。"""
|
||||
session.add(
|
||||
TemplateModel(
|
||||
id=template_id,
|
||||
user_id=user_id,
|
||||
name=f"模板-{template_id}",
|
||||
mode="one_take",
|
||||
is_active=active,
|
||||
)
|
||||
)
|
||||
for order in range(clip_count):
|
||||
session.add(
|
||||
TemplateClipConfigModel(
|
||||
id=f"cc-{template_id}-{order}",
|
||||
template_id=template_id,
|
||||
clip_type="main",
|
||||
order=order,
|
||||
min_duration=5.0,
|
||||
max_duration=8.0,
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
|
||||
def _seed_global_template(session, template_id: str, *, status: str = "active", clip_count: int = 2):
|
||||
"""创建新表 edit_templates 全局模板(+ template_clip_configs 片段配置)。"""
|
||||
session.add(
|
||||
EditTemplateModel(
|
||||
id=template_id,
|
||||
name=f"全局模板-{template_id}",
|
||||
template_type="default",
|
||||
editing_mode="one_take",
|
||||
status=status,
|
||||
)
|
||||
)
|
||||
for order in range(clip_count):
|
||||
session.add(
|
||||
TemplateClipConfigModel(
|
||||
id=f"gcc-{template_id}-{order}",
|
||||
template_id=template_id,
|
||||
clip_type="main",
|
||||
order=order,
|
||||
min_duration=3.0,
|
||||
max_duration=6.0,
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Service 读路径测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestListClipConfigsForEditor:
|
||||
"""list_clip_configs_for_editor 单一数据源 + 归属/存在性判定。"""
|
||||
|
||||
def test_legacy_user_template_returns_configs(self):
|
||||
"""用户自建模板(templates 表 + 3 条 clip_configs)→ 正常返回,不抛异常。"""
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
|
||||
session = _make_session()
|
||||
_seed_legacy_template(session, "tmpl-legacy", USER_ID, clip_count=3)
|
||||
|
||||
svc = EditTemplateService(session)
|
||||
configs = svc.list_clip_configs_for_editor("tmpl-legacy", USER_ID)
|
||||
|
||||
assert len(configs) == 3
|
||||
assert [c.order for c in configs] == [0, 1, 2]
|
||||
assert all(c.min_duration == 5.0 for c in configs)
|
||||
|
||||
def test_missing_template_raises_not_found(self):
|
||||
"""模板不存在(两表都没有)→ TemplateNotFoundError。"""
|
||||
from app.services.edit_template_service import EditTemplateService, TemplateNotFoundError
|
||||
|
||||
session = _make_session()
|
||||
svc = EditTemplateService(session)
|
||||
|
||||
with pytest.raises(TemplateNotFoundError):
|
||||
svc.list_clip_configs_for_editor("tmpl-not-exist", USER_ID)
|
||||
|
||||
def test_other_users_template_raises_not_found(self):
|
||||
"""他人模板(user_id 不匹配)→ TemplateNotFoundError(归属校验)。"""
|
||||
from app.services.edit_template_service import EditTemplateService, TemplateNotFoundError
|
||||
|
||||
session = _make_session()
|
||||
_seed_legacy_template(session, "tmpl-owner", OTHER_USER_ID, clip_count=3)
|
||||
|
||||
svc = EditTemplateService(session)
|
||||
with pytest.raises(TemplateNotFoundError):
|
||||
svc.list_clip_configs_for_editor("tmpl-owner", USER_ID)
|
||||
|
||||
def test_deleted_legacy_template_raises_not_found(self):
|
||||
"""已软删除(is_active=False)的旧表模板 → TemplateNotFoundError。"""
|
||||
from app.services.edit_template_service import EditTemplateService, TemplateNotFoundError
|
||||
|
||||
session = _make_session()
|
||||
_seed_legacy_template(session, "tmpl-deleted", USER_ID, active=False, clip_count=3)
|
||||
|
||||
svc = EditTemplateService(session)
|
||||
with pytest.raises(TemplateNotFoundError):
|
||||
svc.list_clip_configs_for_editor("tmpl-deleted", USER_ID)
|
||||
|
||||
def test_legacy_template_without_configs_returns_empty(self):
|
||||
"""模板存在且归属正确但无片段配置 → 返回空列表(不抛异常,路由层映射 422)。"""
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
|
||||
session = _make_session()
|
||||
_seed_legacy_template(session, "tmpl-noconfig", USER_ID, clip_count=0)
|
||||
|
||||
svc = EditTemplateService(session)
|
||||
configs = svc.list_clip_configs_for_editor("tmpl-noconfig", USER_ID)
|
||||
assert configs == []
|
||||
|
||||
def test_global_template_returns_configs(self):
|
||||
"""新表 edit_templates 全局模板(无 user_id)→ 任意用户可读,正常返回。"""
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
|
||||
session = _make_session()
|
||||
_seed_global_template(session, "tmpl-global", clip_count=2)
|
||||
|
||||
svc = EditTemplateService(session)
|
||||
configs = svc.list_clip_configs_for_editor("tmpl-global", USER_ID)
|
||||
|
||||
assert len(configs) == 2
|
||||
assert [c.order for c in configs] == [0, 1]
|
||||
|
||||
def test_normal_legacy_request_does_not_raise_valueerror(self):
|
||||
"""正常旧表模板请求绝不在读路径抛 ValueError: 模板不存在(回归保护)。"""
|
||||
import logging
|
||||
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
|
||||
session = _make_session()
|
||||
_seed_legacy_template(session, "tmpl-ok", USER_ID, clip_count=3)
|
||||
svc = EditTemplateService(session)
|
||||
|
||||
with pytest.MonkeyPatch.context() as mp:
|
||||
# 若读路径意外抛 ValueError 并被记录为异常堆栈,测试能感知
|
||||
errors: list[str] = []
|
||||
mp.setattr(
|
||||
logging.getLogger("app.services.edit_template_service"),
|
||||
"exception",
|
||||
lambda *a, **k: errors.append(str(a)),
|
||||
)
|
||||
configs = svc.list_clip_configs_for_editor("tmpl-ok", USER_ID)
|
||||
|
||||
assert len(configs) == 3
|
||||
assert errors == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_template_segments 映射测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_clip_config(order: int, min_dur: float = 3.0, max_dur: float = 8.0):
|
||||
"""构造 mock TemplateClipConfig 领域实体."""
|
||||
cc = MagicMock()
|
||||
cc.order = order
|
||||
cc.min_duration = min_dur
|
||||
@@ -37,207 +218,53 @@ def _make_clip_config(order: int, min_dur: float = 3.0, max_dur: float = 8.0):
|
||||
return cc
|
||||
|
||||
|
||||
def _make_old_segment(segment_order: int, dur_min: float = 4.0, dur_max: float = 7.0):
|
||||
"""构造 mock 旧 TemplateSegment."""
|
||||
s = MagicMock()
|
||||
s.segment_order = segment_order
|
||||
s.duration_min = dur_min
|
||||
s.duration_max = dur_max
|
||||
return s
|
||||
class TestGetTemplateSegments:
|
||||
"""_get_template_segments 仅做映射/排序,异常与空配置语义明确。"""
|
||||
|
||||
def test_maps_and_sorts_configs(self):
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetTemplateSegmentsFallback:
|
||||
"""_get_template_segments 三级回退链."""
|
||||
|
||||
def test_new_system_works(self):
|
||||
"""路径1:新模板系统正常返回 → 直接使用."""
|
||||
configs = [_make_clip_config(0, 2.0, 6.0), _make_clip_config(1, 3.0, 9.0)]
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.return_value = configs
|
||||
db = MagicMock()
|
||||
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
|
||||
assert len(result) == 2
|
||||
assert result[0] == (0, 2.0, 6.0)
|
||||
assert result[1] == (1, 3.0, 9.0)
|
||||
tpl_svc.list_clip_configs.assert_called_once_with(TEST_TEMPLATE_ID)
|
||||
|
||||
def test_main_table_missing_direct_query_succeeds(self):
|
||||
"""路径2(P0修复):主表不存在 ValueError → 直接查表成功.
|
||||
|
||||
模拟自建模板在 edit_templates 主表已删除/不存在,
|
||||
但 template_clip_configs 表有记录。
|
||||
"""
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.side_effect = ValueError(f"模板不存在: {TEST_TEMPLATE_ID}")
|
||||
db = MagicMock()
|
||||
|
||||
# Mock SQLAlchemyTemplateClipConfigRepository
|
||||
direct_configs = [
|
||||
_make_clip_config(0, 2.0, 5.0),
|
||||
_make_clip_config(1, 3.0, 7.0),
|
||||
_make_clip_config(2, 4.0, 8.0),
|
||||
]
|
||||
with (
|
||||
__import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateClipConfigRepository"
|
||||
) as mock_repo_cls,
|
||||
):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_template.return_value = direct_configs
|
||||
mock_repo_cls.return_value = mock_repo
|
||||
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
|
||||
assert len(result) == 3
|
||||
assert result[0] == (0, 2.0, 5.0)
|
||||
assert result[1] == (1, 3.0, 7.0)
|
||||
assert result[2] == (2, 4.0, 8.0)
|
||||
mock_repo.list_by_template.assert_called_once_with(TEST_TEMPLATE_ID)
|
||||
|
||||
def test_main_table_missing_direct_query_empty_falls_to_old(self):
|
||||
"""路径2→3:主表不存在 + 直接查表为空 → 回退旧系统."""
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.side_effect = ValueError("模板不存在")
|
||||
db = MagicMock()
|
||||
|
||||
old_segments = [_make_old_segment(0, 3.0, 6.0)]
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateClipConfigRepository"
|
||||
) as mock_repo_cls:
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_template.return_value = [] # 新表也没记录
|
||||
mock_repo_cls.return_value = mock_repo
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateRepository"
|
||||
) as mock_old_cls:
|
||||
mock_old = MagicMock()
|
||||
mock_old.list_segments.return_value = old_segments
|
||||
mock_old_cls.return_value = mock_old
|
||||
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0] == (0, 3.0, 6.0)
|
||||
|
||||
def test_all_fail_returns_empty(self):
|
||||
"""路径4:三级全部失败 → 返回空列表."""
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.side_effect = ValueError("模板不存在")
|
||||
db = MagicMock()
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateClipConfigRepository"
|
||||
) as mock_repo_cls:
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_template.side_effect = Exception("DB error")
|
||||
mock_repo_cls.return_value = mock_repo
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateRepository"
|
||||
) as mock_old_cls:
|
||||
mock_old = MagicMock()
|
||||
mock_old.list_segments.return_value = [] # 旧表也空
|
||||
mock_old_cls.return_value = mock_old
|
||||
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
|
||||
assert result == []
|
||||
|
||||
def test_direct_query_sorts_by_order(self):
|
||||
"""直接查表返回的结果按 order 排序."""
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.side_effect = ValueError("模板不存在")
|
||||
db = MagicMock()
|
||||
|
||||
# 故意乱序
|
||||
configs = [
|
||||
tpl_svc.list_clip_configs_for_editor.return_value = [
|
||||
_make_clip_config(2, 5.0, 10.0),
|
||||
_make_clip_config(0, 2.0, 4.0),
|
||||
_make_clip_config(1, 3.0, 6.0),
|
||||
]
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateClipConfigRepository"
|
||||
) as mock_repo_cls:
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_template.return_value = configs
|
||||
mock_repo_cls.return_value = mock_repo
|
||||
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
result = _get_template_segments("tmpl-1", USER_ID, tpl_svc)
|
||||
|
||||
assert [r[0] for r in result] == [0, 1, 2]
|
||||
assert result[0] == (0, 2.0, 4.0)
|
||||
assert result[1] == (1, 3.0, 6.0)
|
||||
assert result[2] == (2, 5.0, 10.0)
|
||||
tpl_svc.list_clip_configs_for_editor.assert_called_once_with("tmpl-1", USER_ID)
|
||||
|
||||
def test_empty_configs_returns_empty(self):
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
|
||||
def test_direct_query_handles_none_durations(self):
|
||||
"""直接查表时 min/max_duration 为 None → 使用默认值."""
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.side_effect = ValueError("模板不存在")
|
||||
db = MagicMock()
|
||||
tpl_svc.list_clip_configs_for_editor.return_value = []
|
||||
|
||||
assert _get_template_segments("tmpl-1", USER_ID, tpl_svc) == []
|
||||
|
||||
def test_missing_template_propagates_not_found(self):
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
from app.services.edit_template_service import TemplateNotFoundError
|
||||
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs_for_editor.side_effect = TemplateNotFoundError("tmpl-x")
|
||||
|
||||
with pytest.raises(TemplateNotFoundError):
|
||||
_get_template_segments("tmpl-x", USER_ID, tpl_svc)
|
||||
|
||||
def test_none_durations_use_default(self):
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
|
||||
tpl_svc = MagicMock()
|
||||
cc = MagicMock()
|
||||
cc.order = 0
|
||||
cc.min_duration = None
|
||||
cc.max_duration = None
|
||||
tpl_svc.list_clip_configs_for_editor.return_value = [cc]
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateClipConfigRepository"
|
||||
) as mock_repo_cls:
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_template.return_value = [cc]
|
||||
mock_repo_cls.return_value = mock_repo
|
||||
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
|
||||
assert len(result) == 1
|
||||
# None → default (5.0), max(None or None) → default (5.0)
|
||||
assert result[0] == (0, DEFAULT_DUR, DEFAULT_DUR)
|
||||
|
||||
def test_new_system_returns_empty_tries_direct(self):
|
||||
"""新模板系统返回空列表(非异常)→ 继续尝试直接查表."""
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.return_value = [] # 空列表,非异常
|
||||
db = MagicMock()
|
||||
|
||||
direct_configs = [_make_clip_config(0, 3.0, 6.0)]
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateClipConfigRepository"
|
||||
) as mock_repo_cls:
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_template.return_value = direct_configs
|
||||
mock_repo_cls.return_value = mock_repo
|
||||
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
|
||||
# 新系统返回空 → 不走 except → 但也没 return → 继续往下走
|
||||
# 直接查表有数据 → 返回
|
||||
assert len(result) == 1
|
||||
assert result[0] == (0, 3.0, 6.0)
|
||||
|
||||
def test_existing_template_unaffected(self):
|
||||
"""正常模板(主表存在)行为不变."""
|
||||
configs = [_make_clip_config(0, 2.0, 5.0)]
|
||||
tpl_svc = MagicMock()
|
||||
tpl_svc.list_clip_configs.return_value = configs
|
||||
db = MagicMock()
|
||||
|
||||
with __import__("unittest.mock", fromlist=["patch"]).patch(
|
||||
"app.api.routes.templates_editor.clips.SQLAlchemyTemplateClipConfigRepository"
|
||||
) as mock_repo_cls:
|
||||
result = _get_template_segments(TEST_TEMPLATE_ID, tpl_svc, db)
|
||||
# 直接查表不应被调用(新系统已返回)
|
||||
mock_repo_cls.assert_not_called()
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0] == (0, 2.0, 5.0)
|
||||
result = _get_template_segments("tmpl-1", USER_ID, tpl_svc)
|
||||
assert result == [(0, DEFAULT_DUR, DEFAULT_DUR)]
|
||||
|
||||
@@ -128,12 +128,12 @@ class TestAssetLibraryRepoCounting:
|
||||
assert lib.asset_count == 0
|
||||
assert lib.total_size == 0
|
||||
|
||||
repo.increment_asset_count(lib.id, 1024)
|
||||
repo.increment_asset_count(lib.id, size_delta=1024)
|
||||
fetched = repo.get(lib.id)
|
||||
assert fetched.asset_count == 1
|
||||
assert fetched.total_size == 1024
|
||||
|
||||
repo.increment_asset_count(lib.id, 2048)
|
||||
repo.increment_asset_count(lib.id, size_delta=2048)
|
||||
fetched = repo.get(lib.id)
|
||||
assert fetched.asset_count == 2
|
||||
assert fetched.total_size == 3072
|
||||
@@ -144,7 +144,7 @@ class TestAssetLibraryRepoCounting:
|
||||
lib.total_size = 3000
|
||||
repo.create(lib)
|
||||
|
||||
repo.decrement_asset_count(lib.id, 1000)
|
||||
repo.decrement_asset_count(lib.id, size_delta=1000)
|
||||
fetched = repo.get(lib.id)
|
||||
assert fetched.asset_count == 2
|
||||
assert fetched.total_size == 2000
|
||||
@@ -157,14 +157,14 @@ class TestAssetLibraryRepoCounting:
|
||||
repo.create(lib)
|
||||
|
||||
# 减 2 次,应该被钳制到 0
|
||||
repo.decrement_asset_count(lib.id, 200)
|
||||
repo.decrement_asset_count(lib.id, size_delta=200)
|
||||
fetched = repo.get(lib.id)
|
||||
assert fetched.asset_count == 0
|
||||
assert fetched.total_size == 0
|
||||
|
||||
def test_increment_nonexistent_library_no_error(self, repo):
|
||||
"""对不存在的素材库操作,不抛异常也无效果."""
|
||||
repo.increment_asset_count("nonexistent", 100)
|
||||
repo.increment_asset_count("nonexistent", size_delta=100)
|
||||
# 不报错
|
||||
assert repo.get("nonexistent") is None
|
||||
|
||||
|
||||
@@ -97,40 +97,40 @@ class TestInMemoryAssetLibraryRepository:
|
||||
|
||||
def test_increment_asset_count(self, repo, lib_video):
|
||||
repo.create(lib_video)
|
||||
repo.increment_asset_count(lib_video.id, 1024)
|
||||
repo.increment_asset_count(lib_video.id, size_delta=1024)
|
||||
|
||||
lib = repo.get(lib_video.id)
|
||||
assert lib.asset_count == 1
|
||||
assert lib.total_size == 1024
|
||||
|
||||
repo.increment_asset_count(lib_video.id, 512)
|
||||
repo.increment_asset_count(lib_video.id, size_delta=512)
|
||||
lib = repo.get(lib_video.id)
|
||||
assert lib.asset_count == 2
|
||||
assert lib.total_size == 1536
|
||||
|
||||
def test_increment_asset_count_nonexistent(self, repo):
|
||||
# 不报错,静默忽略
|
||||
repo.increment_asset_count("nonexistent", 100)
|
||||
repo.increment_asset_count("nonexistent", size_delta=100)
|
||||
|
||||
def test_decrement_asset_count(self, repo, lib_video):
|
||||
repo.create(lib_video)
|
||||
repo.increment_asset_count(lib_video.id, 1024)
|
||||
repo.increment_asset_count(lib_video.id, 512)
|
||||
repo.increment_asset_count(lib_video.id, size_delta=1024)
|
||||
repo.increment_asset_count(lib_video.id, size_delta=512)
|
||||
|
||||
repo.decrement_asset_count(lib_video.id, 512)
|
||||
repo.decrement_asset_count(lib_video.id, size_delta=512)
|
||||
lib = repo.get(lib_video.id)
|
||||
assert lib.asset_count == 1
|
||||
assert lib.total_size == 1024
|
||||
|
||||
def test_decrement_asset_count_not_below_zero(self, repo, lib_video):
|
||||
repo.create(lib_video)
|
||||
repo.decrement_asset_count(lib_video.id, 9999)
|
||||
repo.decrement_asset_count(lib_video.id, size_delta=9999)
|
||||
lib = repo.get(lib_video.id)
|
||||
assert lib.asset_count == 0
|
||||
assert lib.total_size == 0
|
||||
|
||||
def test_decrement_asset_count_nonexistent(self, repo):
|
||||
repo.decrement_asset_count("nonexistent", 100)
|
||||
repo.decrement_asset_count("nonexistent", size_delta=100)
|
||||
|
||||
|
||||
# ==================== Tag ====================
|
||||
|
||||
@@ -226,10 +226,10 @@ class TestGetMediakitRecommendations:
|
||||
|
||||
|
||||
class TestGetTemplateSegments:
|
||||
"""测试模板片段配置查询。"""
|
||||
"""测试模板片段配置查询(单一数据源:template_clip_configs)。"""
|
||||
|
||||
def test_returns_segments_from_new_template_system(self):
|
||||
"""新模板系统(clip_configs)有数据时优先使用。"""
|
||||
def test_returns_segments_from_clip_configs(self):
|
||||
"""片段配置主表(clip_configs)有数据时按 order 排序返回。"""
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
|
||||
mock_tpl_svc = MagicMock()
|
||||
@@ -241,66 +241,34 @@ class TestGetTemplateSegments:
|
||||
cc2.order = 1
|
||||
cc2.min_duration = 4.0
|
||||
cc2.max_duration = 8.0
|
||||
mock_tpl_svc.list_clip_configs.return_value = [cc2, cc1] # 乱序返回
|
||||
mock_tpl_svc.list_clip_configs_for_editor.return_value = [cc2, cc1] # 乱序返回
|
||||
|
||||
result = _get_template_segments("tmpl-1", mock_tpl_svc, MagicMock())
|
||||
result = _get_template_segments("tmpl-1", "user-1", mock_tpl_svc)
|
||||
assert len(result) == 2
|
||||
assert result[0] == (0, 3.0, 5.0)
|
||||
assert result[1] == (1, 4.0, 8.0)
|
||||
mock_tpl_svc.list_clip_configs_for_editor.assert_called_once_with("tmpl-1", "user-1")
|
||||
|
||||
def test_falls_back_to_old_template_segments(self):
|
||||
"""新模板系统无数据时回退到旧系统。"""
|
||||
def test_returns_empty_when_no_configs(self):
|
||||
"""模板存在但没有片段配置时返回空列表(路由层据此返回 422)。"""
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
|
||||
mock_tpl_svc = MagicMock()
|
||||
mock_tpl_svc.list_clip_configs.return_value = []
|
||||
mock_tpl_svc.list_clip_configs_for_editor.return_value = []
|
||||
|
||||
with patch("app.api.routes.templates_editor.clips.SQLAlchemyTemplateRepository") as MockRepo:
|
||||
mock_repo = MagicMock()
|
||||
seg1 = MagicMock()
|
||||
seg1.segment_order = 0
|
||||
seg1.duration_min = 2.0
|
||||
seg1.duration_max = 4.0
|
||||
mock_repo.list_segments.return_value = [seg1]
|
||||
MockRepo.return_value = mock_repo
|
||||
result = _get_template_segments("tmpl-1", "user-1", mock_tpl_svc)
|
||||
assert result == []
|
||||
|
||||
result = _get_template_segments("tmpl-1", mock_tpl_svc, MagicMock())
|
||||
assert len(result) == 1
|
||||
assert result[0] == (0, 2.0, 4.0)
|
||||
|
||||
def test_returns_empty_when_no_segments(self):
|
||||
"""两套系统都没有片段配置时返回空列表。"""
|
||||
def test_missing_template_raises(self):
|
||||
"""模板不存在/无权限时服务层抛 TemplateNotFoundError(路由层据此返回 404)。"""
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
from app.services.edit_template_service import TemplateNotFoundError
|
||||
|
||||
mock_tpl_svc = MagicMock()
|
||||
mock_tpl_svc.list_clip_configs.return_value = []
|
||||
mock_tpl_svc.list_clip_configs_for_editor.side_effect = TemplateNotFoundError("tmpl-x")
|
||||
|
||||
with patch("app.api.routes.templates_editor.clips.SQLAlchemyTemplateRepository") as MockRepo:
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_segments.return_value = []
|
||||
MockRepo.return_value = mock_repo
|
||||
|
||||
result = _get_template_segments("tmpl-1", mock_tpl_svc, MagicMock())
|
||||
assert result == []
|
||||
|
||||
def test_new_system_exception_falls_back(self):
|
||||
"""新模板系统异常时回退到旧系统。"""
|
||||
from app.api.routes.templates_editor.clips import _get_template_segments
|
||||
|
||||
mock_tpl_svc = MagicMock()
|
||||
mock_tpl_svc.list_clip_configs.side_effect = RuntimeError("db error")
|
||||
|
||||
with patch("app.api.routes.templates_editor.clips.SQLAlchemyTemplateRepository") as MockRepo:
|
||||
mock_repo = MagicMock()
|
||||
seg = MagicMock()
|
||||
seg.segment_order = 0
|
||||
seg.duration_min = 1.0
|
||||
seg.duration_max = 3.0
|
||||
mock_repo.list_segments.return_value = [seg]
|
||||
MockRepo.return_value = mock_repo
|
||||
|
||||
result = _get_template_segments("tmpl-1", mock_tpl_svc, MagicMock())
|
||||
assert len(result) == 1
|
||||
with pytest.raises(TemplateNotFoundError):
|
||||
_get_template_segments("tmpl-x", "user-1", mock_tpl_svc)
|
||||
|
||||
|
||||
# ── from-assets 端点集成测试 ────────────────────────────────────────────────
|
||||
@@ -358,7 +326,7 @@ def _make_tpl_svc_with_segments(segments):
|
||||
"""segments: list of (order, min_dur, max_dur)"""
|
||||
svc = MagicMock()
|
||||
clip_configs = [_make_clip_config(o, mn, mx) for o, mn, mx in segments]
|
||||
svc.list_clip_configs.return_value = clip_configs
|
||||
svc.list_clip_configs_for_editor.return_value = clip_configs
|
||||
return svc
|
||||
|
||||
|
||||
@@ -540,34 +508,56 @@ class TestFromAssetsByTemplateSegments:
|
||||
orders = [c["order"] for c in clips_data]
|
||||
assert orders == [0, 1, 2]
|
||||
|
||||
def test_no_segments_raises_400(self):
|
||||
"""模板没有 segment 配置时返回 400。"""
|
||||
def test_no_segments_raises_422(self):
|
||||
"""模板存在但未配置片段时返回 422(与模板不存在的 404 区分)。"""
|
||||
from app.api.routes.templates_editor.clips import create_clips_from_assets_editor
|
||||
from app.api.routes.templates_editor.schemas import ClipsFromAssetsRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_tpl_svc = MagicMock()
|
||||
mock_tpl_svc.list_clip_configs.return_value = []
|
||||
mock_tpl_svc.list_clip_configs_for_editor.return_value = []
|
||||
mock_plan_svc = _make_plan_svc()
|
||||
|
||||
with patch("app.api.routes.templates_editor.clips.SQLAlchemyTemplateRepository") as MockRepo:
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_segments.return_value = []
|
||||
MockRepo.return_value = mock_repo
|
||||
body = ClipsFromAssetsRequest(asset_ids=["a1"])
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
create_clips_from_assets_editor(
|
||||
template_id="tmpl-1",
|
||||
body=body,
|
||||
background_tasks=MagicMock(),
|
||||
plan_id="plan-1",
|
||||
services=(mock_tpl_svc, mock_plan_svc),
|
||||
asset_repo=MagicMock(),
|
||||
db=MagicMock(),
|
||||
current_user=_make_auth_user(),
|
||||
)
|
||||
assert exc_info.value.status_code == 422
|
||||
|
||||
body = ClipsFromAssetsRequest(asset_ids=["a1"])
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
create_clips_from_assets_editor(
|
||||
template_id="tmpl-1",
|
||||
body=body,
|
||||
background_tasks=MagicMock(),
|
||||
plan_id="plan-1",
|
||||
services=(mock_tpl_svc, mock_plan_svc),
|
||||
asset_repo=MagicMock(),
|
||||
db=MagicMock(),
|
||||
current_user=_make_auth_user(),
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_plan_svc.replace_all_clips_transactional.assert_not_called()
|
||||
|
||||
def test_template_not_found_raises_404(self):
|
||||
"""模板不存在/已删除/无权限时返回 404。"""
|
||||
from app.api.routes.templates_editor.clips import create_clips_from_assets_editor
|
||||
from app.api.routes.templates_editor.schemas import ClipsFromAssetsRequest
|
||||
from app.services.edit_template_service import TemplateNotFoundError
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_tpl_svc = MagicMock()
|
||||
mock_tpl_svc.list_clip_configs_for_editor.side_effect = TemplateNotFoundError("tmpl-x")
|
||||
mock_plan_svc = _make_plan_svc()
|
||||
|
||||
body = ClipsFromAssetsRequest(asset_ids=["a1"])
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
create_clips_from_assets_editor(
|
||||
template_id="tmpl-x",
|
||||
body=body,
|
||||
background_tasks=MagicMock(),
|
||||
plan_id="plan-1",
|
||||
services=(mock_tpl_svc, mock_plan_svc),
|
||||
asset_repo=MagicMock(),
|
||||
db=MagicMock(),
|
||||
current_user=_make_auth_user(),
|
||||
)
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
mock_plan_svc.replace_all_clips_transactional.assert_not_called()
|
||||
|
||||
|
||||
@@ -110,3 +110,32 @@ class TestPixelPerturbationAcceptance:
|
||||
# 至少 2 种不同组合
|
||||
unique = len(set(results))
|
||||
assert unique >= 2, f"Expected >= 2 unique filter combos, got {unique}: {results}"
|
||||
|
||||
|
||||
class TestIntSeedSupport:
|
||||
"""int seed 入参支持(与 get_rhythm_template(seed) 接口一致)。"""
|
||||
|
||||
def test_int_seed_returns_dict(self):
|
||||
"""int seed 正常返回 dict。"""
|
||||
result = generate_pixel_perturbation(42)
|
||||
assert isinstance(result, dict)
|
||||
assert "filters" in result
|
||||
|
||||
def test_int_seed_reproducible(self):
|
||||
"""相同 int seed 结果一致。"""
|
||||
assert generate_pixel_perturbation(42) == generate_pixel_perturbation(42)
|
||||
|
||||
def test_int_seed_differs_across_seeds(self):
|
||||
"""不同 int seed 大概率不同(遍历确认至少 2 种组合)。"""
|
||||
results = {tuple(generate_pixel_perturbation(s)["filters"]) for s in range(30)}
|
||||
assert len(results) >= 2
|
||||
|
||||
def test_int_seed_matches_random_obj(self):
|
||||
"""int seed 与等价 random.Random(seed) 结果一致。"""
|
||||
assert generate_pixel_perturbation(7) == generate_pixel_perturbation(random.Random(7))
|
||||
|
||||
def test_none_seed_works(self):
|
||||
"""None 入参(默认随机)正常返回。"""
|
||||
result = generate_pixel_perturbation(None)
|
||||
assert isinstance(result, dict)
|
||||
assert len(result["filters"]) in [2, 3]
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""节奏模板单元测试(Issue #1764)。
|
||||
|
||||
覆盖:
|
||||
- RHYTHM_TEMPLATES 池定义(6 种模板)
|
||||
- RHYTHM_TEMPLATES 池定义(8 种模板,#1764 原始 6 种 + #1768 新增 2 种)
|
||||
- get_rhythm_template:根据 seed 选择模板
|
||||
- adapt_template_length:适配不同片段数
|
||||
- plan_clip_durations:按权重分配时长
|
||||
@@ -26,8 +26,8 @@ class TestRhythmTemplates:
|
||||
"""节奏模板池测试。"""
|
||||
|
||||
def test_six_templates_defined(self):
|
||||
"""预设 6 种节奏模板。"""
|
||||
assert len(RHYTHM_TEMPLATES) == 6
|
||||
"""预设 8 种节奏模板(#1764 原始 6 种 + #1768 新增 2 种)。"""
|
||||
assert len(RHYTHM_TEMPLATES) == 8
|
||||
|
||||
def test_average_template_is_all_ones(self):
|
||||
"""第一种模板是平均(全 1)。"""
|
||||
|
||||
@@ -188,6 +188,7 @@ class TestListTemplatesUseCase:
|
||||
tag="tag1",
|
||||
keyword="test",
|
||||
mode="one_take",
|
||||
valid_only=False,
|
||||
)
|
||||
|
||||
def test_list_pagination(self):
|
||||
@@ -226,6 +227,7 @@ class TestCountTemplatesUseCase:
|
||||
tag="tag1",
|
||||
keyword="kw",
|
||||
mode="pip",
|
||||
valid_only=False,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,319 @@
|
||||
"""#1766 转场位置与类型随机化测试 (packages/domain/transition_randomizer.py).
|
||||
|
||||
覆盖:
|
||||
- TRANSITION_POOL 定义(5 种效果)
|
||||
- generate_transition_plan:硬切比例 30%-50%
|
||||
- generate_transition_plan:转场时长 0.3s ~ 0.8s
|
||||
- generate_transition_plan:jitter 在 ±0.5s 范围内
|
||||
- 不同 seed 产生不同转场序列
|
||||
- 与片段时长协同:长片段间转场更长
|
||||
- 边界情况:num_transitions=0、1 个片段
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
for sub in ("packages", ""):
|
||||
p = str(REPO_ROOT / sub) if sub else str(REPO_ROOT)
|
||||
if p not in sys.path:
|
||||
sys.path.insert(0, p)
|
||||
|
||||
from packages.domain import transition_randomizer as tr # noqa: E402
|
||||
from packages.domain.transition_randomizer import ( # noqa: E402
|
||||
CUT_RATIO_MAX,
|
||||
CUT_RATIO_MIN,
|
||||
HARD_CUT,
|
||||
TIMING_JITTER_MAX,
|
||||
TRANSITION_DURATION_MAX,
|
||||
TRANSITION_DURATION_MIN,
|
||||
TRANSITION_POOL,
|
||||
_cut_probability_for_pair,
|
||||
generate_transition_plan,
|
||||
)
|
||||
|
||||
|
||||
class TestTransitionPool:
|
||||
"""转场类型池定义测试。"""
|
||||
|
||||
def test_pool_has_5_effects(self):
|
||||
assert len(TRANSITION_POOL) == 5
|
||||
|
||||
def test_pool_contains_expected_effects(self):
|
||||
assert "dissolve" in TRANSITION_POOL
|
||||
assert "zoomin" in TRANSITION_POOL
|
||||
assert "slideleft" in TRANSITION_POOL
|
||||
assert "wipeleft" in TRANSITION_POOL
|
||||
assert "fade" in TRANSITION_POOL
|
||||
|
||||
def test_pool_does_not_contain_cut(self):
|
||||
assert HARD_CUT not in TRANSITION_POOL
|
||||
|
||||
|
||||
class TestConstants:
|
||||
"""常量约束测试。"""
|
||||
|
||||
def test_duration_range(self):
|
||||
assert TRANSITION_DURATION_MIN == 0.3
|
||||
assert TRANSITION_DURATION_MAX == 0.8
|
||||
|
||||
def test_cut_ratio_range(self):
|
||||
assert CUT_RATIO_MIN == 0.3
|
||||
assert CUT_RATIO_MAX == 0.5
|
||||
|
||||
def test_jitter_max(self):
|
||||
assert TIMING_JITTER_MAX == 0.5
|
||||
|
||||
|
||||
class TestCutProbability:
|
||||
"""硬切概率计算测试。"""
|
||||
|
||||
def test_short_clips_higher_cut_prob(self):
|
||||
"""短片段(<3s)→ 硬切概率偏高。"""
|
||||
prob = _cut_probability_for_pair(2.0, 2.5)
|
||||
assert prob > 0.4 # 高于基准
|
||||
|
||||
def test_long_clips_lower_cut_prob(self):
|
||||
"""长片段(>6s)→ 硬切概率偏低。"""
|
||||
prob = _cut_probability_for_pair(8.0, 7.0)
|
||||
assert prob < 0.4 # 低于基准
|
||||
|
||||
def test_medium_clips_base_prob(self):
|
||||
"""中等片段(3-6s)→ 基准概率。"""
|
||||
prob = _cut_probability_for_pair(4.0, 5.0)
|
||||
assert abs(prob - 0.4) < 1e-6
|
||||
|
||||
def test_probability_in_range(self):
|
||||
"""概率始终在 [MIN, MAX] 范围内。"""
|
||||
for prev in [1.0, 3.0, 5.0, 8.0, 15.0]:
|
||||
for next_ in [1.0, 3.0, 5.0, 8.0, 15.0]:
|
||||
prob = _cut_probability_for_pair(prev, next_)
|
||||
assert CUT_RATIO_MIN <= prob <= CUT_RATIO_MAX
|
||||
|
||||
|
||||
class TestGenerateTransitionPlan:
|
||||
"""generate_transition_plan 核心测试。"""
|
||||
|
||||
def test_zero_transitions_returns_empty(self):
|
||||
assert generate_transition_plan(0) == []
|
||||
|
||||
def test_negative_transitions_returns_empty(self):
|
||||
assert generate_transition_plan(-1) == []
|
||||
|
||||
def test_returns_correct_count(self):
|
||||
plan = generate_transition_plan(5, rng=random.Random(42))
|
||||
assert len(plan) == 5
|
||||
|
||||
def test_each_item_has_required_keys(self):
|
||||
plan = generate_transition_plan(3, rng=random.Random(42))
|
||||
for item in plan:
|
||||
assert "effect" in item
|
||||
assert "duration" in item
|
||||
assert "jitter" in item
|
||||
|
||||
def test_effects_are_valid(self):
|
||||
"""所有 effect 要么是 cut 要么是 TRANSITION_POOL 中的。"""
|
||||
plan = generate_transition_plan(20, rng=random.Random(42))
|
||||
valid_effects = set(TRANSITION_POOL) | {HARD_CUT}
|
||||
for item in plan:
|
||||
assert item["effect"] in valid_effects
|
||||
|
||||
def test_hard_cut_ratio_in_range_many_samples(self):
|
||||
"""100 个转场点,硬切比例在 30%-50%(统计保证)。"""
|
||||
plan = generate_transition_plan(
|
||||
100,
|
||||
clip_durations=[5.0] * 101,
|
||||
rng=random.Random(42),
|
||||
)
|
||||
num_cuts = sum(1 for item in plan if item["effect"] == HARD_CUT)
|
||||
ratio = num_cuts / len(plan)
|
||||
# 统计波动允许 ±10% 的宽松范围
|
||||
assert 0.20 <= ratio <= 0.60, f"硬切比例 {ratio:.2%} 超出宽松范围"
|
||||
# 更严格的范围检查(±5%)
|
||||
assert CUT_RATIO_MIN - 0.05 <= ratio <= CUT_RATIO_MAX + 0.05, f"硬切比例 {ratio:.2%} 超出 [25%, 55%] 范围"
|
||||
|
||||
def test_transition_duration_in_range(self):
|
||||
"""非硬切转场的时长在 [0.3, 0.8] 范围内。"""
|
||||
plan = generate_transition_plan(30, rng=random.Random(42))
|
||||
for item in plan:
|
||||
if item["effect"] != HARD_CUT:
|
||||
assert (
|
||||
TRANSITION_DURATION_MIN <= item["duration"] <= TRANSITION_DURATION_MAX
|
||||
), f"转场时长 {item['duration']} 超出 [{TRANSITION_DURATION_MIN}, {TRANSITION_DURATION_MAX}]"
|
||||
|
||||
def test_cut_duration_is_zero(self):
|
||||
"""硬切转场的时长必须为 0。"""
|
||||
plan = generate_transition_plan(20, rng=random.Random(42))
|
||||
for item in plan:
|
||||
if item["effect"] == HARD_CUT:
|
||||
assert item["duration"] == 0.0
|
||||
|
||||
def test_jitter_in_range(self):
|
||||
"""jitter 在 [-0.5, +0.5] 范围内。"""
|
||||
plan = generate_transition_plan(30, rng=random.Random(42))
|
||||
for item in plan:
|
||||
assert -TIMING_JITTER_MAX <= item["jitter"] <= TIMING_JITTER_MAX, f"jitter {item['jitter']} 超出范围"
|
||||
|
||||
def test_cut_jitter_is_zero(self):
|
||||
"""硬切转场的 jitter 必须为 0。"""
|
||||
plan = generate_transition_plan(20, rng=random.Random(42))
|
||||
for item in plan:
|
||||
if item["effect"] == HARD_CUT:
|
||||
assert item["jitter"] == 0.0
|
||||
|
||||
def test_different_seeds_produce_different_plans(self):
|
||||
"""不同 seed 产生不同的转场序列(至少 2 组不同)。"""
|
||||
plans_seen = set()
|
||||
for seed in range(20):
|
||||
plan = generate_transition_plan(5, rng=random.Random(seed))
|
||||
plan_sig = tuple((item["effect"], item["duration"]) for item in plan)
|
||||
plans_seen.add(plan_sig)
|
||||
assert len(plans_seen) >= 2, "20 个 seed 只产生 1 种转场序列"
|
||||
|
||||
def test_same_seed_same_plan(self):
|
||||
"""相同 seed 产生相同的转场序列(确定性)。"""
|
||||
plan1 = generate_transition_plan(5, rng=random.Random(42))
|
||||
plan2 = generate_transition_plan(5, rng=random.Random(42))
|
||||
assert plan1 == plan2
|
||||
|
||||
def test_long_clips_longer_transitions(self):
|
||||
"""长片段(>6s)之间的转场倾向于比短片段更长。"""
|
||||
# 长片段
|
||||
long_plan = generate_transition_plan(
|
||||
20,
|
||||
clip_durations=[10.0] * 21,
|
||||
rng=random.Random(42),
|
||||
)
|
||||
# 短片段
|
||||
short_plan = generate_transition_plan(
|
||||
20,
|
||||
clip_durations=[2.0] * 21,
|
||||
rng=random.Random(42),
|
||||
)
|
||||
# 长片段的非硬切转场平均时长
|
||||
long_durs = [item["duration"] for item in long_plan if item["effect"] != HARD_CUT]
|
||||
short_durs = [item["duration"] for item in short_plan if item["effect"] != HARD_CUT]
|
||||
|
||||
if long_durs and short_durs:
|
||||
avg_long = sum(long_durs) / len(long_durs)
|
||||
avg_short = sum(short_durs) / len(short_durs)
|
||||
# 长片段平均转场时长 >= 短片段(协同节奏)
|
||||
assert avg_long >= avg_short * 0.95, f"长片段转场 {avg_long:.3f}s 不应显著短于短片段 {avg_short:.3f}s"
|
||||
|
||||
def test_clip_durations_none_uses_default(self):
|
||||
"""clip_durations=None 时使用默认值 5.0。"""
|
||||
plan = generate_transition_plan(3, rng=random.Random(42))
|
||||
assert len(plan) == 3
|
||||
|
||||
def test_fewer_clip_durations_than_needed(self):
|
||||
"""clip_durations 长度不足时用默认值补齐。"""
|
||||
plan = generate_transition_plan(
|
||||
5,
|
||||
clip_durations=[4.0, 5.0], # 只需前 2 个
|
||||
rng=random.Random(42),
|
||||
)
|
||||
assert len(plan) == 5
|
||||
|
||||
|
||||
class TestTransitionRandomizationIntegration:
|
||||
"""转场随机化与变体生成集成测试。"""
|
||||
|
||||
def test_reselect_produces_different_transitions(self):
|
||||
"""多次 reselect_clips_for_variant 产生不同的转场序列。"""
|
||||
from packages.domain.variant_plan_selector import reselect_clips_for_variant
|
||||
|
||||
source_clips = [
|
||||
{
|
||||
"order": i,
|
||||
"asset_id": f"asset_{i}",
|
||||
"start_time": 0.0,
|
||||
"duration": 5.0,
|
||||
"clip_type": "main",
|
||||
"playback_speed": 1.0,
|
||||
"transition_effect": "cut",
|
||||
"transition_duration": 0.0,
|
||||
"text_content": "",
|
||||
"config": {},
|
||||
}
|
||||
for i in range(4)
|
||||
]
|
||||
asset_durations = {f"asset_{i}": 30.0 for i in range(4)}
|
||||
|
||||
transition_seqs = set()
|
||||
for seed in range(5):
|
||||
rng = random.Random(seed)
|
||||
result = reselect_clips_for_variant(
|
||||
source_clips,
|
||||
list(asset_durations.keys()),
|
||||
asset_durations=asset_durations,
|
||||
rng=rng,
|
||||
)
|
||||
seq = tuple(
|
||||
(c.get("transition_effect"), round(c.get("transition_duration", 0), 2))
|
||||
for c in result
|
||||
if c.get("clip_type") == "main"
|
||||
)
|
||||
transition_seqs.add(seq)
|
||||
|
||||
assert len(transition_seqs) >= 2, f"5 个 seed 只产生 {len(transition_seqs)} 种转场序列"
|
||||
|
||||
def test_reselect_preserves_non_main_transitions(self):
|
||||
"""非 main 片段(intro/outro)的转场不被随机化。"""
|
||||
from packages.domain.variant_plan_selector import reselect_clips_for_variant
|
||||
|
||||
source_clips = [
|
||||
{
|
||||
"order": 0,
|
||||
"asset_id": "intro_asset",
|
||||
"start_time": 0.0,
|
||||
"duration": 3.0,
|
||||
"clip_type": "intro",
|
||||
"playback_speed": 1.0,
|
||||
"transition_effect": "fade",
|
||||
"transition_duration": 0.5,
|
||||
"text_content": "",
|
||||
"config": {},
|
||||
},
|
||||
{
|
||||
"order": 1,
|
||||
"asset_id": "a1",
|
||||
"start_time": 0.0,
|
||||
"duration": 5.0,
|
||||
"clip_type": "main",
|
||||
"playback_speed": 1.0,
|
||||
"transition_effect": "cut",
|
||||
"transition_duration": 0.0,
|
||||
"text_content": "",
|
||||
"config": {},
|
||||
},
|
||||
{
|
||||
"order": 2,
|
||||
"asset_id": "a2",
|
||||
"start_time": 0.0,
|
||||
"duration": 5.0,
|
||||
"clip_type": "main",
|
||||
"playback_speed": 1.0,
|
||||
"transition_effect": "cut",
|
||||
"transition_duration": 0.0,
|
||||
"text_content": "",
|
||||
"config": {},
|
||||
},
|
||||
]
|
||||
asset_durations = {"intro_asset": 10.0, "a1": 30.0, "a2": 30.0}
|
||||
|
||||
result = reselect_clips_for_variant(
|
||||
source_clips,
|
||||
list(asset_durations.keys()),
|
||||
asset_durations=asset_durations,
|
||||
rng=random.Random(42),
|
||||
)
|
||||
|
||||
# intro 片段的转场保持不变
|
||||
intro_clip = next(c for c in result if c["clip_type"] == "intro")
|
||||
assert intro_clip["transition_effect"] == "fade"
|
||||
assert intro_clip["transition_duration"] == 0.5
|
||||
@@ -158,6 +158,52 @@ class TestListByUser:
|
||||
assert len(result[0].segments) == 1
|
||||
assert result[0].segments[0].duration_min == 2.0
|
||||
|
||||
def test_valid_only_filters_templates_without_segments(self, repo, session):
|
||||
"""#1769: valid_only=True 时排除两张片段表都没有记录的无效模板."""
|
||||
# 有效模板:有 clip_configs
|
||||
valid_clip = _make_template(name="有效模板-clip_configs")
|
||||
repo.create(valid_clip)
|
||||
repo.create_segments([_make_segment(valid_clip.id, order=1)])
|
||||
# 有效模板:仅有旧表 template_segments 记录
|
||||
valid_old = _make_template(name="有效模板-old_segments")
|
||||
repo.create(valid_old)
|
||||
old = TemplateSegmentModel(
|
||||
id=str(uuid.uuid4()),
|
||||
template_id=valid_old.id,
|
||||
segment_order=1,
|
||||
duration_min=2.0,
|
||||
duration_max=6.0,
|
||||
)
|
||||
session.add(old)
|
||||
session.commit()
|
||||
# 无效模板:两张表都没有记录
|
||||
invalid = _make_template(name="无效模板-无片段")
|
||||
repo.create(invalid)
|
||||
|
||||
# 默认不过滤:编辑器视角能看到全部 3 个模板
|
||||
all_templates = repo.list_by_user("u1")
|
||||
assert len(all_templates) == 3
|
||||
assert repo.count_by_user("u1") == 3
|
||||
|
||||
# valid_only=True:剪辑页视角只返回 2 个有效模板
|
||||
valid_templates = repo.list_by_user("u1", valid_only=True)
|
||||
assert {t.name for t in valid_templates} == {"有效模板-clip_configs", "有效模板-old_segments"}
|
||||
assert all(len(t.segments) > 0 for t in valid_templates)
|
||||
assert repo.count_by_user("u1", valid_only=True) == 2
|
||||
|
||||
def test_valid_only_with_filters_and_pagination(self, repo, session):
|
||||
"""valid_only 与其他过滤/分页条件组合使用."""
|
||||
tpl = _make_template(name="口播模板", mode="voice_over")
|
||||
repo.create(tpl)
|
||||
repo.create_segments([_make_segment(tpl.id, order=1, material_type="人物")])
|
||||
_invalid = _make_template(name="口播无效模板", mode="voice_over")
|
||||
repo.create(_invalid)
|
||||
|
||||
result = repo.list_by_user("u1", mode="voice_over", valid_only=True)
|
||||
assert len(result) == 1
|
||||
assert result[0].name == "口播模板"
|
||||
assert repo.count_by_user("u1", mode="voice_over", valid_only=True) == 1
|
||||
|
||||
|
||||
class TestCopyTemplate:
|
||||
def test_copy_writes_to_clip_configs(self, repo, session):
|
||||
|
||||
@@ -8,6 +8,7 @@ from __future__ import annotations
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest import TestCase
|
||||
from unittest.mock import patch
|
||||
|
||||
# 修正 import 路径
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
@@ -62,13 +63,14 @@ class _StubPlan:
|
||||
self,
|
||||
plan_id: str = "plan-1",
|
||||
status: EditPlanStatus = EditPlanStatus.EDITING,
|
||||
config: dict | None = None,
|
||||
):
|
||||
self.id = plan_id
|
||||
self.template_id = "tpl-1"
|
||||
self.name = "测试计划"
|
||||
self.status = status
|
||||
self.total_duration = 0.0
|
||||
self.config = {}
|
||||
self.config = config or {}
|
||||
|
||||
|
||||
# ── Stub 仓储 ─────────────────────────────────────────────────────────────────
|
||||
@@ -712,6 +714,65 @@ class TestHasAudioTitleSubtitleFix(TestCase):
|
||||
self.assertEqual(chain.audio_label, "a0")
|
||||
|
||||
|
||||
# ── #1789 标题 drawtext 集成测试 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestComposeCommandTitleDrawtext(TestCase):
|
||||
"""build_compose_command 中标题 drawtext 集成测试。"""
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_title_config_injected(self, mock_font):
|
||||
"""plan.config 有 title 时,filter_complex 包含 drawtext。"""
|
||||
mock_font.return_value = ""
|
||||
plan = _StubPlan(config={"title": {"text": "测试标题", "font_size": 48, "position": "top"}})
|
||||
clips = [_make_ready_clip(plan_id=plan.id)]
|
||||
svc = _make_service(plan, clips)
|
||||
cmd = svc.build_compose_command(plan.id, "/tmp/out.mp4")
|
||||
|
||||
self.assertIn("drawtext=", cmd.filter_complex)
|
||||
self.assertIn("[composed]", cmd.filter_complex)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_title_config_alt_key(self, mock_font):
|
||||
"""plan.config['title'] 无文本时回退到 title_config。"""
|
||||
mock_font.return_value = ""
|
||||
plan = _StubPlan(config={"title": {}, "title_config": {"text": "备用标题", "font_size": 36}})
|
||||
clips = [_make_ready_clip(plan_id=plan.id)]
|
||||
svc = _make_service(plan, clips)
|
||||
cmd = svc.build_compose_command(plan.id, "/tmp/out.mp4")
|
||||
|
||||
self.assertIn("drawtext=", cmd.filter_complex)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_no_title_no_drawtext(self, mock_font):
|
||||
"""无标题配置时,filter_complex 不包含 drawtext。"""
|
||||
mock_font.return_value = ""
|
||||
plan = _StubPlan(config={})
|
||||
clips = [_make_ready_clip(plan_id=plan.id)]
|
||||
svc = _make_service(plan, clips)
|
||||
cmd = svc.build_compose_command(plan.id, "/tmp/out.mp4")
|
||||
|
||||
self.assertNotIn("drawtext=", cmd.filter_complex)
|
||||
|
||||
def test_title_config_not_dict(self):
|
||||
"""title config 为非 dict 值时不崩溃。"""
|
||||
plan = _StubPlan(config={"title": "not a dict"})
|
||||
clips = [_make_ready_clip(plan_id=plan.id)]
|
||||
svc = _make_service(plan, clips)
|
||||
cmd = svc.build_compose_command(plan.id, "/tmp/out.mp4")
|
||||
self.assertIsNotNone(cmd)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_title_content_field(self, mock_font):
|
||||
"""title config 使用 content 字段(前端 TitleConfig 命名)。"""
|
||||
mock_font.return_value = ""
|
||||
plan = _StubPlan(config={"title": {"content": "内容标题", "font_size": 36}})
|
||||
clips = [_make_ready_clip(plan_id=plan.id)]
|
||||
svc = _make_service(plan, clips)
|
||||
cmd = svc.build_compose_command(plan.id, "/tmp/out.mp4")
|
||||
self.assertIn("drawtext=", cmd.filter_complex)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import unittest
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from dataclasses import FrozenInstanceError
|
||||
from unittest.mock import patch
|
||||
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
@@ -26,9 +27,12 @@ from packages.domain.video_filter_builder import (
|
||||
DEFAULT_TRANSITION_DURATION,
|
||||
XFADE_TRANSITION_MAP,
|
||||
ClipFilterChain,
|
||||
_escape_drawtext_text,
|
||||
_resolve_font_path,
|
||||
build_clip_filter,
|
||||
build_concat_filter,
|
||||
build_filter_complex,
|
||||
build_title_drawtext_filter,
|
||||
build_xfade_filter,
|
||||
chain_filters,
|
||||
has_audio,
|
||||
@@ -852,5 +856,265 @@ class TestEndToEndFilterBuilding(unittest.TestCase):
|
||||
self.assertNotIn("[a0]", filter_str)
|
||||
|
||||
|
||||
# ── #1789 标题 drawtext 滤镜补充覆盖率 ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestEscapeDrawtextText(unittest.TestCase):
|
||||
"""直接测试转义函数,覆盖每一行。"""
|
||||
|
||||
def test_backslash_escape(self):
|
||||
result = _escape_drawtext_text("a\\b")
|
||||
self.assertIn("\\\\", result)
|
||||
|
||||
def test_single_quote_escape(self):
|
||||
result = _escape_drawtext_text("it's")
|
||||
self.assertIn("\\'", result)
|
||||
|
||||
def test_colon_escape(self):
|
||||
result = _escape_drawtext_text("a:b")
|
||||
self.assertIn("\\:", result)
|
||||
|
||||
def test_percent_escape(self):
|
||||
result = _escape_drawtext_text("100%")
|
||||
self.assertIn("%%", result)
|
||||
|
||||
def test_all_special_chars_combined(self):
|
||||
result = _escape_drawtext_text("\\':%")
|
||||
self.assertIn("\\\\", result)
|
||||
self.assertIn("\\'", result)
|
||||
self.assertIn("\\:", result)
|
||||
self.assertIn("%%", result)
|
||||
|
||||
def test_no_special_chars(self):
|
||||
result = _escape_drawtext_text("hello world")
|
||||
self.assertEqual(result, "hello world")
|
||||
|
||||
|
||||
class TestResolveFontPath(unittest.TestCase):
|
||||
"""测试字体路径解析逻辑。"""
|
||||
|
||||
@patch("os.path.isfile")
|
||||
def test_known_font_found(self, mock_isfile):
|
||||
mock_isfile.side_effect = lambda p: "NotoSansCJK" in p
|
||||
result = _resolve_font_path("思源黑体")
|
||||
self.assertNotEqual(result, "")
|
||||
self.assertIn("NotoSansCJK", result)
|
||||
|
||||
@patch("os.path.isfile")
|
||||
def test_unknown_font_fallback(self, mock_isfile):
|
||||
mock_isfile.side_effect = lambda p: "DejaVu" in p
|
||||
result = _resolve_font_path("UnknownFont")
|
||||
self.assertIn("DejaVu", result)
|
||||
|
||||
@patch("os.path.isfile")
|
||||
def test_no_fonts_available(self, mock_isfile):
|
||||
mock_isfile.return_value = False
|
||||
result = _resolve_font_path("思源黑体")
|
||||
self.assertEqual(result, "")
|
||||
|
||||
@patch("os.path.isfile")
|
||||
def test_passthrough_font_name(self, mock_isfile):
|
||||
mock_isfile.side_effect = lambda p: "NotoSansCJK" in p
|
||||
result = _resolve_font_path("NotoSansCJK")
|
||||
self.assertNotEqual(result, "")
|
||||
|
||||
@patch("os.path.isfile")
|
||||
def test_font_search_first_match(self, mock_isfile):
|
||||
mock_isfile.side_effect = lambda p: "opentype" in p
|
||||
result = _resolve_font_path("思源黑体")
|
||||
self.assertNotEqual(result, "")
|
||||
self.assertIn("opentype", result)
|
||||
|
||||
@patch("os.path.isfile")
|
||||
def test_font_fallback_skips_nonexistent(self, mock_isfile):
|
||||
mock_isfile.side_effect = lambda p: "DejaVu" in p
|
||||
result = _resolve_font_path("不存在字体")
|
||||
self.assertIn("DejaVu", result)
|
||||
|
||||
|
||||
class TestDrawtextFontFileIncluded(unittest.TestCase):
|
||||
"""当字体文件存在时,fontfile 参数出现在输出中。"""
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_fontfile_in_output(self, mock_font):
|
||||
mock_font.return_value = "/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc"
|
||||
result = build_title_drawtext_filter({"text": "标题"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("fontfile=", result)
|
||||
self.assertIn("NotoSansCJK", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_fontfile_escaped(self, mock_font):
|
||||
mock_font.return_value = "/path/with:special'chars.ttf"
|
||||
result = build_title_drawtext_filter({"text": "标题"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("fontfile=", result)
|
||||
|
||||
|
||||
class TestDrawtextFontFileNotIncluded(unittest.TestCase):
|
||||
"""当字体文件不存在时,无 fontfile 参数。"""
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_no_fontfile(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertNotIn("fontfile=", result)
|
||||
|
||||
|
||||
class TestDrawtextStrokeBranches(unittest.TestCase):
|
||||
"""stroke 各分支覆盖。"""
|
||||
|
||||
def test_stroke_non_bool_non_dict(self):
|
||||
result = build_title_drawtext_filter({"text": "标题", "stroke": "yes"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertNotIn("borderw", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_stroke_dict_default_color(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "stroke": {"width": 4}})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("borderw=4", result)
|
||||
self.assertIn("bordercolor=000000", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_stroke_dict_enabled_false(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "stroke": {"enabled": False, "width": 5}})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertNotIn("borderw", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_stroke_dict_custom_color(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "stroke": {"width": 2, "color": "#ff0000"}})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("bordercolor=ff0000", result)
|
||||
|
||||
|
||||
class TestDrawtextShadowBranches(unittest.TestCase):
|
||||
"""shadow 各分支覆盖。"""
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_shadow_dict_default_color(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "shadow": {"offset_x": 5, "offset_y": 5}})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("shadowcolor=000000", result)
|
||||
self.assertIn("shadowx=5", result)
|
||||
self.assertIn("shadowy=5", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_shadow_dict_disabled(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "shadow": {"enabled": False}})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertNotIn("shadowcolor", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_shadow_dict_custom_color(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter(
|
||||
{"text": "标题", "shadow": {"color": "#555555", "offset_x": 1, "offset_y": 1}}
|
||||
)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("shadowcolor=555555", result)
|
||||
|
||||
|
||||
class TestDrawtextBoldFalse(unittest.TestCase):
|
||||
def test_bold_false(self):
|
||||
result = build_title_drawtext_filter({"text": "标题", "bold": False})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertNotIn("font=bold", result)
|
||||
|
||||
|
||||
class TestDrawtextPositionBranches(unittest.TestCase):
|
||||
"""位置相关分支覆盖。"""
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_position_top_explicit(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "position": "top"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("y=50", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_position_center_explicit(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "position": "center"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("y=(h-text_h)/2", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_position_bottom_explicit(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "position": "bottom"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("y=h-text_h-50", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_position_custom_with_float_coords(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "position": "custom", "pos_x": 100.7, "pos_y": 200.3})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("x=100", result)
|
||||
self.assertIn("y=200", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_position_custom_bool_coords_fallback(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "position": "custom", "pos_x": True, "pos_y": True})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("x=(w-text_w)/2", result)
|
||||
self.assertIn("y=50", result)
|
||||
|
||||
|
||||
class TestDrawtextColorNoHash(unittest.TestCase):
|
||||
def test_color_without_hash(self):
|
||||
result = build_title_drawtext_filter({"text": "标题", "font_color": "red"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("fontcolor=red", result)
|
||||
|
||||
|
||||
class TestDrawtextFieldNormalization(unittest.TestCase):
|
||||
"""字段归一化覆盖更多分支。"""
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_content_fallback(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"content": "备用标题"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("备用标题", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_font_preset_fallback(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "font_preset": "楷体"})
|
||||
self.assertIsNotNone(result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_size_fallback(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "size": 72})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("fontsize=72", result)
|
||||
|
||||
@patch("packages.domain.video_filter_builder._resolve_font_path")
|
||||
def test_color_fallback(self, mock_font):
|
||||
mock_font.return_value = ""
|
||||
result = build_title_drawtext_filter({"text": "标题", "color": "#abcdef"})
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIn("fontcolor=abcdef", result)
|
||||
|
||||
|
||||
class TestDrawtextNotDictConfig(unittest.TestCase):
|
||||
def test_string_config_returns_none(self):
|
||||
self.assertIsNone(build_title_drawtext_filter("not a dict"))
|
||||
|
||||
def test_list_config_returns_none(self):
|
||||
self.assertIsNone(build_title_drawtext_filter([1, 2, 3]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,261 @@
|
||||
"""#1766 xfade_builder 逐转场时长与位置微调测试.
|
||||
|
||||
覆盖:
|
||||
- build_xfade_filter_chain:transition_durations 参数(逐转场独立时长)
|
||||
- build_xfade_filter_chain:jitters 参数(位置微调偏移)
|
||||
- 向后兼容:不传新参数时行为不变
|
||||
- TransitionEngine.build_xfade_chain:透传新参数
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
for sub in ("packages", ""):
|
||||
p = str(REPO_ROOT / sub) if sub else str(REPO_ROOT)
|
||||
if p not in sys.path:
|
||||
sys.path.insert(0, p)
|
||||
|
||||
from packages.domain.xfade_builder import ( # noqa: E402
|
||||
DEFAULT_TRANSITION_DURATION,
|
||||
build_xfade_filter_chain,
|
||||
)
|
||||
|
||||
|
||||
class TestBackwardCompatibility:
|
||||
"""向后兼容测试:不传新参数时行为不变。"""
|
||||
|
||||
def test_single_transition_duration(self):
|
||||
"""全局 transition_duration 仍有效。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
)
|
||||
assert "xfade" in result
|
||||
assert "duration=0.500" in result
|
||||
assert dur > 0
|
||||
|
||||
def test_default_transition_duration(self):
|
||||
"""不传 transition_duration 时使用默认值 0.5。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
)
|
||||
assert "duration=0.500" in result
|
||||
|
||||
|
||||
class TestPerTransitionDurations:
|
||||
"""逐转场独立时长测试。"""
|
||||
|
||||
def test_different_durations_per_transition(self):
|
||||
"""每个转场使用不同的时长。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1", "v2"],
|
||||
transitions=["cut", "fade", "dissolve"],
|
||||
transition_durations=[0.3, 0.8], # 第 1 个转场 0.3s,第 2 个 0.8s
|
||||
)
|
||||
assert "duration=0.300" in result
|
||||
assert "duration=0.800" in result
|
||||
|
||||
def test_partial_durations_fallback_to_global(self):
|
||||
"""transition_durations 长度不足时 fallback 到 transition_duration。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1", "v2"],
|
||||
transitions=["cut", "fade", "dissolve"],
|
||||
transition_duration=0.5,
|
||||
transition_durations=[0.3], # 只有第一个,第二个 fallback 到 0.5
|
||||
)
|
||||
assert "duration=0.300" in result
|
||||
assert "duration=0.500" in result
|
||||
|
||||
def test_empty_durations_uses_global(self):
|
||||
"""transition_durations=[] 时使用全局 transition_duration。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.6,
|
||||
transition_durations=[],
|
||||
)
|
||||
assert "duration=0.600" in result
|
||||
|
||||
def test_none_durations_uses_global(self):
|
||||
"""transition_durations=None 时使用全局 transition_duration。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.6,
|
||||
transition_durations=None,
|
||||
)
|
||||
assert "duration=0.600" in result
|
||||
|
||||
def test_duration_clamped_by_clip_length(self):
|
||||
"""转场时长不能超过相邻片段时长。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[2.0, 2.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_durations=[1.5], # 1.5s > 2.0s * 0.4,会被钳制
|
||||
)
|
||||
# 钳制到可用范围内
|
||||
assert "duration=" in result
|
||||
|
||||
def test_total_duration_reflects_per_transition(self):
|
||||
"""总时长反映逐转场的重叠量。"""
|
||||
# 使用 0.3s 转场
|
||||
_, dur_short = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_durations=[0.3],
|
||||
)
|
||||
# 使用 0.8s 转场
|
||||
_, dur_long = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_durations=[0.8],
|
||||
)
|
||||
# 更长转场 → 更多重叠 → 总时长更短
|
||||
assert dur_long < dur_short
|
||||
|
||||
|
||||
class TestJitters:
|
||||
"""位置微调 jitter 测试。"""
|
||||
|
||||
def test_positive_jitter_delays_transition(self):
|
||||
"""正 jitter 推迟转场(offset 增大)。"""
|
||||
result_no_jitter, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
)
|
||||
result_with_jitter, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
jitters=[0.3], # 正 jitter:推迟转场
|
||||
)
|
||||
# 提取 offset 值
|
||||
import re
|
||||
|
||||
offset_no = float(re.search(r"offset=([\d.]+)", result_no_jitter).group(1))
|
||||
offset_yes = float(re.search(r"offset=([\d.]+)", result_with_jitter).group(1))
|
||||
assert offset_yes > offset_no
|
||||
|
||||
def test_negative_jitter_advances_transition(self):
|
||||
"""负 jitter 提前转场(offset 减小)。"""
|
||||
result_no_jitter, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
)
|
||||
result_with_jitter, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
jitters=[-0.3], # 负 jitter:提前转场
|
||||
)
|
||||
import re
|
||||
|
||||
offset_no = float(re.search(r"offset=([\d.]+)", result_no_jitter).group(1))
|
||||
offset_yes = float(re.search(r"offset=([\d.]+)", result_with_jitter).group(1))
|
||||
assert offset_yes < offset_no
|
||||
|
||||
def test_jitter_clamped_to_valid_range(self):
|
||||
"""jitter 不会使 offset 超出有效范围(>=0)。"""
|
||||
result, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
jitters=[-100.0], # 极大负 jitter
|
||||
)
|
||||
import re
|
||||
|
||||
offset = float(re.search(r"offset=([\d.]+)", result).group(1))
|
||||
assert offset >= 0.0
|
||||
|
||||
def test_per_transition_jitters(self):
|
||||
"""每个转场可以有独立的 jitter。"""
|
||||
result, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1", "v2"],
|
||||
transitions=["cut", "fade", "dissolve"],
|
||||
transition_duration=0.5,
|
||||
jitters=[0.2, -0.1],
|
||||
)
|
||||
import re
|
||||
|
||||
offsets = [float(m.group(1)) for m in re.finditer(r"offset=([\d.]+)", result)]
|
||||
assert len(offsets) == 2
|
||||
|
||||
def test_empty_jitters_no_effect(self):
|
||||
"""jitters=[] 等同于无 jitter。"""
|
||||
result_no, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
)
|
||||
result_empty, _ = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "fade"],
|
||||
transition_duration=0.5,
|
||||
jitters=[],
|
||||
)
|
||||
assert result_no == result_empty
|
||||
|
||||
|
||||
class TestEdgeCases:
|
||||
"""边界情况测试。"""
|
||||
|
||||
def test_single_clip_with_durations(self):
|
||||
"""单片段传入 transition_durations 不报错。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0],
|
||||
clip_video_labels=["v0"],
|
||||
transitions=["cut"],
|
||||
transition_durations=[0.5],
|
||||
)
|
||||
assert "copy" in result
|
||||
assert dur == 5.0
|
||||
|
||||
def test_empty_clips(self):
|
||||
"""空片段列表返回空字符串。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[],
|
||||
clip_video_labels=[],
|
||||
transitions=[],
|
||||
transition_durations=[],
|
||||
jitters=[],
|
||||
)
|
||||
assert result == ""
|
||||
assert dur == 0.0
|
||||
|
||||
def test_all_cut_transitions(self):
|
||||
"""全硬切场景(不进入 xfade,由调用方处理 concat)。"""
|
||||
result, dur = build_xfade_filter_chain(
|
||||
clip_durations=[5.0, 5.0],
|
||||
clip_video_labels=["v0", "v1"],
|
||||
transitions=["cut", "cut"],
|
||||
transition_durations=[0.0, 0.0],
|
||||
)
|
||||
# cut 转场仍然会生成 xfade 滤镜(因为底层不区分 cut)
|
||||
# 调用方(unified_render_service)负责检测全硬切并走 concat
|
||||
assert dur > 0
|
||||
Reference in New Issue
Block a user