Compare commits
19 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f02ea7207f | |||
| 9f153eec54 | |||
| b392bb1b78 | |||
| b852603664 | |||
| bda2170b9d | |||
| bb0d01f080 | |||
| f6a458564f | |||
| 48bf66c298 | |||
| 8e2c563b42 | |||
| b99437fd81 | |||
| 81eff29e7b | |||
| 02d60a02f2 | |||
| 5b0ba40ecd | |||
| 9b9dc243ed | |||
| 6514b8c34d | |||
| 4e8265ec5e | |||
| 23de3906c2 | |||
| 50b05db03e | |||
| 7f3c462617 |
@@ -16,6 +16,11 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
SQLAlchemyEditPlanRepository,
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.domain.clip_operations import calculate_merge as _calc_merge
|
||||
from packages.domain.clip_operations import calculate_shift_orders as _calc_shift_orders
|
||||
from packages.domain.clip_operations import calculate_split as _calc_split
|
||||
from packages.domain.clip_operations import validate_merge_clips as _validate_merge
|
||||
from packages.domain.clip_operations import validate_split_time as _validate_split
|
||||
from packages.domain.edit_plan import EditPlan, EditPlanStatus
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
|
||||
@@ -384,36 +389,45 @@ class EditPlanService:
|
||||
clip = self.get_clip_or_raise(clip_id)
|
||||
plan_id = clip.plan_id
|
||||
|
||||
if split_time <= 0 or split_time >= clip.duration:
|
||||
raise ValueError(f"分割时间必须在 (0, {clip.duration:.3f}) 范围内,当前: {split_time}")
|
||||
# 纯逻辑:校验 + 计算
|
||||
_validate_split(split_time, clip.duration)
|
||||
split = _calc_split(
|
||||
duration=clip.duration,
|
||||
split_time=split_time,
|
||||
start_time=clip.start_time,
|
||||
)
|
||||
|
||||
self._auto_resume_editing(plan_id)
|
||||
|
||||
original_duration = clip.duration
|
||||
left_duration = round(split_time, 3)
|
||||
right_duration = round(original_duration - split_time, 3)
|
||||
original_order = clip.order
|
||||
|
||||
# 更新左半部分(原片段)
|
||||
clip.duration = left_duration
|
||||
clip.duration = split.left_duration
|
||||
left_clip = self._clip_repo.update(clip)
|
||||
|
||||
# 后面片段的 order 全部 +1(给右半部分腾位置)
|
||||
all_clips = self._clip_repo.list_by_plan(plan_id)
|
||||
for c in all_clips:
|
||||
if c.order > original_order and c.id != clip_id:
|
||||
c.order += 1
|
||||
self._clip_repo.update(c)
|
||||
shifts = _calc_shift_orders(
|
||||
all_clips,
|
||||
threshold_order=original_order,
|
||||
shift=1,
|
||||
excluded_ids={clip_id},
|
||||
id_attr="id",
|
||||
order_attr="order",
|
||||
)
|
||||
for c, new_order in shifts:
|
||||
c.order = new_order
|
||||
self._clip_repo.update(c)
|
||||
|
||||
# 创建右半部分新片段(继承原片段的大部分属性)
|
||||
right_config = dict(clip.config) if clip.config else {}
|
||||
# 素材裁剪信息
|
||||
if clip.asset_id:
|
||||
# 右半部分从 split_time 开始播放
|
||||
right_config["trim_start"] = left_duration
|
||||
right_config["trim_start"] = split.right_trim_start
|
||||
# 左半部分在 split_time 处结束
|
||||
left_config = dict(left_clip.config) if left_clip.config else {}
|
||||
left_config["trim_end"] = right_duration
|
||||
left_config["trim_end"] = split.left_trim_end
|
||||
left_clip.config = left_config
|
||||
left_clip = self._clip_repo.update(left_clip)
|
||||
|
||||
@@ -424,8 +438,8 @@ class EditPlanService:
|
||||
template_clip_config_id=clip.template_clip_config_id,
|
||||
asset_id=clip.asset_id,
|
||||
text_content=clip.text_content,
|
||||
start_time=clip.start_time + left_duration,
|
||||
duration=right_duration,
|
||||
start_time=split.right_start_time,
|
||||
duration=split.right_duration,
|
||||
transition_effect=clip.transition_effect,
|
||||
transition_duration=clip.transition_duration,
|
||||
playback_speed=clip.playback_speed,
|
||||
@@ -438,8 +452,8 @@ class EditPlanService:
|
||||
clip_id,
|
||||
plan_id,
|
||||
split_time,
|
||||
left_duration,
|
||||
right_duration,
|
||||
split.left_duration,
|
||||
split.right_duration,
|
||||
)
|
||||
|
||||
return {
|
||||
@@ -468,70 +482,45 @@ class EditPlanService:
|
||||
clip = self.get_clip_or_raise(cid)
|
||||
clips.append(clip)
|
||||
|
||||
# 校验:同一计划
|
||||
plan_id = clips[0].plan_id
|
||||
for c in clips[1:]:
|
||||
if c.plan_id != plan_id:
|
||||
raise ValueError("只能合并同一计划下的片段")
|
||||
|
||||
# 按 order 排序
|
||||
clips.sort(key=lambda c: c.order)
|
||||
|
||||
# 校验:order 连续
|
||||
for i in range(1, len(clips)):
|
||||
if clips[i].order != clips[i - 1].order + 1:
|
||||
raise ValueError(f"片段不连续:order {clips[i-1].order} → {clips[i].order}")
|
||||
|
||||
# 校验:类型一致
|
||||
clip_type = clips[0].clip_type
|
||||
for c in clips[1:]:
|
||||
if c.clip_type != clip_type:
|
||||
raise ValueError("只能合并相同类型的片段")
|
||||
# 纯逻辑:校验 + 计算
|
||||
plan_id, first_order = _validate_merge(clips)
|
||||
merge = _calc_merge(clips)
|
||||
|
||||
self._auto_resume_editing(plan_id)
|
||||
|
||||
# 计算合并后的属性
|
||||
first_clip = clips[0]
|
||||
total_duration = round(sum(c.duration for c in clips), 3)
|
||||
first_order = first_clip.order
|
||||
|
||||
# 合并文案(用换行连接)
|
||||
merged_text = "\n".join(c.text_content for c in clips if c.text_content.strip())
|
||||
|
||||
# 合并 config(后面的覆盖前面的)
|
||||
merged_config: Dict[str, Any] = {}
|
||||
for c in clips:
|
||||
if c.config:
|
||||
merged_config.update(c.config)
|
||||
# 清理 trim 相关字段(合并后就是完整片段了)
|
||||
merged_config.pop("trim_start", None)
|
||||
merged_config.pop("trim_end", None)
|
||||
|
||||
# 更新第一个片段(保留它作为合并结果)
|
||||
first_clip.duration = total_duration
|
||||
first_clip.text_content = merged_text
|
||||
first_clip.config = merged_config
|
||||
first_clip = sorted(clips, key=lambda c: c.order)[0]
|
||||
first_clip.duration = merge.total_duration
|
||||
first_clip.text_content = merge.merged_text
|
||||
first_clip.config = merge.merged_config
|
||||
# 转场保留第一个的(合并后的入点转场)
|
||||
# playback_speed 取第一个的
|
||||
merged_clip = self._clip_repo.update(first_clip)
|
||||
|
||||
# 删除其余片段
|
||||
for c in clips[1:]:
|
||||
self._clip_repo.delete(c.id)
|
||||
rest_ids = [c.id for c in clips if c.id != merged_clip.id]
|
||||
for cid in rest_ids:
|
||||
self._clip_repo.delete(cid)
|
||||
|
||||
# 后面的片段 order 前移 (len - 1) 位
|
||||
shift = len(clips) - 1
|
||||
all_clips = self._clip_repo.list_by_plan(plan_id)
|
||||
for c in all_clips:
|
||||
if c.order > first_order and c.id != merged_clip.id:
|
||||
c.order -= shift
|
||||
self._clip_repo.update(c)
|
||||
shifts = _calc_shift_orders(
|
||||
all_clips,
|
||||
threshold_order=first_order,
|
||||
shift=-merge.shift_amount,
|
||||
excluded_ids={merged_clip.id},
|
||||
id_attr="id",
|
||||
order_attr="order",
|
||||
)
|
||||
for c, new_order in shifts:
|
||||
c.order = new_order
|
||||
self._clip_repo.update(c)
|
||||
|
||||
logger.info(
|
||||
"合并片段: plan_id=%s count=%d total_duration=%.3fs",
|
||||
plan_id,
|
||||
len(clips),
|
||||
total_duration,
|
||||
merge.total_duration,
|
||||
)
|
||||
|
||||
return merged_clip
|
||||
|
||||
@@ -23,6 +23,13 @@ from packages.domain.template_clip_config import (
|
||||
TemplateClipConfig,
|
||||
TransitionEffect,
|
||||
)
|
||||
from packages.domain.template_clip_converter import (
|
||||
clip_configs_to_snapshots,
|
||||
clips_to_template_clip_configs,
|
||||
filter_plan_config_to_template,
|
||||
snapshots_to_template_clip_configs,
|
||||
validate_template_name,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -121,9 +128,7 @@ class EditTemplateService:
|
||||
ValueError: 名称为空或重复
|
||||
"""
|
||||
# 名称校验
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
raise ValueError("模板名称不能为空")
|
||||
clean_name = validate_template_name(name)
|
||||
|
||||
# 名称重复检查
|
||||
existing = self._template_repo.list_all(skip=0, limit=1000)
|
||||
@@ -471,12 +476,7 @@ class EditTemplateService:
|
||||
raise ValueError(f"模板名称已存在: {clean_name}")
|
||||
|
||||
# 从计划 config 中提取模板级配置,去掉运行时/素材相关字段
|
||||
plan_config = plan.config or {}
|
||||
template_config: dict[str, Any] = {}
|
||||
for key, value in plan_config.items():
|
||||
# 跳过明显的运行时/实例字段,保留风格/模式类配置
|
||||
if key not in {"asset_ids", "source_edit_plan_id", "generation_task_id"}:
|
||||
template_config[key] = value
|
||||
template_config = filter_plan_config_to_template(plan.config)
|
||||
|
||||
template = EditTemplate.create(
|
||||
name=clean_name,
|
||||
@@ -497,40 +497,7 @@ class EditTemplateService:
|
||||
|
||||
# 5. 转换每个片段为模板片段配置
|
||||
created_configs: List[TemplateClipConfig] = []
|
||||
for clip in clips:
|
||||
clip_config: dict[str, Any] = {}
|
||||
# 播放速度存入 config
|
||||
if clip.playback_speed and clip.playback_speed != 1.0:
|
||||
clip_config["playback_speed"] = clip.playback_speed
|
||||
# 片段自有 config 合并(优先级:clip.config 覆盖上面的)
|
||||
if clip.config:
|
||||
clip_config.update(clip.config)
|
||||
# 去掉素材相关字段
|
||||
clip_config.pop("asset_info", None)
|
||||
clip_config.pop("source_asset_id", None)
|
||||
|
||||
# 转场效果兼容校验
|
||||
try:
|
||||
transition = TransitionEffect(clip.transition_effect)
|
||||
except ValueError:
|
||||
transition = TransitionEffect.CUT
|
||||
|
||||
# 片段类型兼容校验
|
||||
try:
|
||||
clip_type = ClipType(clip.clip_type)
|
||||
except ValueError:
|
||||
clip_type = ClipType.MAIN
|
||||
|
||||
clip_config_obj = TemplateClipConfig.create(
|
||||
template_id=created_template.id,
|
||||
clip_type=clip_type,
|
||||
order=clip.order,
|
||||
min_duration=clip.duration,
|
||||
max_duration=clip.duration,
|
||||
text_template=clip.text_content or "",
|
||||
transition_effect=transition,
|
||||
config=clip_config,
|
||||
)
|
||||
for clip_config_obj in clips_to_template_clip_configs(created_template.id, clips):
|
||||
created = self._clip_config_repo.create(clip_config_obj)
|
||||
created_configs.append(created)
|
||||
|
||||
@@ -679,8 +646,6 @@ class EditTemplateService:
|
||||
Raises:
|
||||
ValueError: 模板/草稿不存在,或草稿不属于该模板
|
||||
"""
|
||||
from packages.domain.template_clip_config import TemplateClipConfig
|
||||
|
||||
# 1. 校验模板和草稿
|
||||
template = self.get_template_or_raise(template_id)
|
||||
draft = self._plan_repo.get(draft_plan_id)
|
||||
@@ -700,39 +665,14 @@ class EditTemplateService:
|
||||
editing_mode = config.get("editing_mode", "one_take")
|
||||
|
||||
# 4. 提取模板配置(去掉草稿/运行时字段)
|
||||
draft_config = draft.config or {}
|
||||
template_config: dict[str, Any] = {}
|
||||
skip_keys = {
|
||||
"is_template_draft",
|
||||
"asset_ids",
|
||||
"source_edit_plan_id",
|
||||
"generation_task_id",
|
||||
}
|
||||
for key, value in draft_config.items():
|
||||
if key not in skip_keys:
|
||||
template_config[key] = value
|
||||
template_config = filter_plan_config_to_template(draft.config)
|
||||
|
||||
# 5. 事务更新
|
||||
try:
|
||||
# 5.0 先保存旧版快照(发布前的状态),用于回滚
|
||||
old_version = template.version or 1
|
||||
old_clip_configs = self._clip_config_repo.list_by_template(template_id)
|
||||
old_clip_snapshots = [
|
||||
{
|
||||
"clip_type": cfg.clip_type.value if hasattr(cfg.clip_type, "value") else cfg.clip_type,
|
||||
"order": cfg.order,
|
||||
"min_duration": cfg.min_duration,
|
||||
"max_duration": cfg.max_duration,
|
||||
"text_template": cfg.text_template or "",
|
||||
"transition_effect": (
|
||||
cfg.transition_effect.value
|
||||
if hasattr(cfg.transition_effect, "value")
|
||||
else cfg.transition_effect
|
||||
),
|
||||
"config": cfg.config or {},
|
||||
}
|
||||
for cfg in old_clip_configs
|
||||
]
|
||||
old_clip_snapshots = clip_configs_to_snapshots(old_clip_configs)
|
||||
|
||||
from packages.domain.template_version import EditTemplateVersion
|
||||
|
||||
@@ -759,46 +699,7 @@ class EditTemplateService:
|
||||
|
||||
# 创建新的片段配置
|
||||
created_configs: list[TemplateClipConfig] = []
|
||||
for clip in draft_clips:
|
||||
clip_config: dict[str, Any] = {}
|
||||
# 播放速度存入 config
|
||||
if clip.playback_speed and clip.playback_speed != 1.0:
|
||||
clip_config["playback_speed"] = clip.playback_speed
|
||||
# 片段自有 config 合并
|
||||
if clip.config:
|
||||
clip_config.update(clip.config)
|
||||
# 去掉素材相关字段
|
||||
clip_config.pop("asset_info", None)
|
||||
clip_config.pop("source_asset_id", None)
|
||||
|
||||
# 转场效果兼容校验
|
||||
try:
|
||||
from packages.domain.template_clip_config import (
|
||||
TransitionEffect,
|
||||
)
|
||||
|
||||
transition = TransitionEffect(clip.transition_effect)
|
||||
except (ValueError, ImportError):
|
||||
transition = TransitionEffect.CUT # type: ignore
|
||||
|
||||
# 片段类型兼容校验
|
||||
try:
|
||||
from packages.domain.template_clip_config import ClipType
|
||||
|
||||
clip_type = ClipType(clip.clip_type)
|
||||
except (ValueError, ImportError):
|
||||
clip_type = ClipType.MAIN # type: ignore
|
||||
|
||||
config_obj = TemplateClipConfig.create(
|
||||
template_id=template_id,
|
||||
clip_type=clip_type,
|
||||
order=clip.order,
|
||||
min_duration=clip.duration,
|
||||
max_duration=clip.duration,
|
||||
text_template=clip.text_content or "",
|
||||
transition_effect=transition,
|
||||
config=clip_config,
|
||||
)
|
||||
for config_obj in clips_to_template_clip_configs(template_id, draft_clips):
|
||||
created = self._clip_config_repo.create(config_obj)
|
||||
created_configs.append(created)
|
||||
|
||||
@@ -843,8 +744,6 @@ class EditTemplateService:
|
||||
Raises:
|
||||
ValueError: 模板/版本不存在
|
||||
"""
|
||||
from packages.domain.template_clip_config import TemplateClipConfig
|
||||
|
||||
template = self.get_template_or_raise(template_id)
|
||||
|
||||
# 1. 读取目标版本快照
|
||||
@@ -857,22 +756,7 @@ class EditTemplateService:
|
||||
try:
|
||||
# 2. 先保存当前状态快照(当前版本号),确保回滚可撤销
|
||||
old_clip_configs = self._clip_config_repo.list_by_template(template_id)
|
||||
old_clip_snapshots = [
|
||||
{
|
||||
"clip_type": cfg.clip_type.value if hasattr(cfg.clip_type, "value") else cfg.clip_type,
|
||||
"order": cfg.order,
|
||||
"min_duration": cfg.min_duration,
|
||||
"max_duration": cfg.max_duration,
|
||||
"text_template": cfg.text_template or "",
|
||||
"transition_effect": (
|
||||
cfg.transition_effect.value
|
||||
if hasattr(cfg.transition_effect, "value")
|
||||
else cfg.transition_effect
|
||||
),
|
||||
"config": cfg.config or {},
|
||||
}
|
||||
for cfg in old_clip_configs
|
||||
]
|
||||
old_clip_snapshots = clip_configs_to_snapshots(old_clip_configs)
|
||||
|
||||
from packages.domain.template_version import EditTemplateVersion
|
||||
|
||||
@@ -905,37 +789,7 @@ class EditTemplateService:
|
||||
synchronize_session=False
|
||||
)
|
||||
|
||||
for clip_snap in target_version.clip_configs:
|
||||
# 转场效果兼容校验
|
||||
try:
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
|
||||
transition = TransitionEffect(clip_snap.get("transition_effect", "cut"))
|
||||
except (ValueError, ImportError):
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
|
||||
transition = TransitionEffect.CUT
|
||||
|
||||
# 片段类型兼容校验
|
||||
try:
|
||||
from packages.domain.template_clip_config import ClipType
|
||||
|
||||
clip_type = ClipType(clip_snap.get("clip_type", "main"))
|
||||
except (ValueError, ImportError):
|
||||
from packages.domain.template_clip_config import ClipType
|
||||
|
||||
clip_type = ClipType.MAIN
|
||||
|
||||
config_obj = TemplateClipConfig.create(
|
||||
template_id=template_id,
|
||||
clip_type=clip_type,
|
||||
order=clip_snap.get("order", 0),
|
||||
min_duration=clip_snap.get("min_duration", 0.0),
|
||||
max_duration=clip_snap.get("max_duration", 0.0),
|
||||
text_template=clip_snap.get("text_template", ""),
|
||||
transition_effect=transition,
|
||||
config=clip_snap.get("config", {}) or {},
|
||||
)
|
||||
for config_obj in snapshots_to_template_clip_configs(template_id, target_version.clip_configs):
|
||||
self._clip_config_repo.create(config_obj)
|
||||
|
||||
self._db.commit()
|
||||
|
||||
@@ -206,4 +206,3 @@ class PlanGeneratorService:
|
||||
委托给 plan_generator_utils.distribute_assets 纯函数。
|
||||
"""
|
||||
distribute_assets(clips, asset_ids, editing_mode)
|
||||
|
||||
|
||||
@@ -21,14 +21,14 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from packages.domain.asset_scoring import MEDIUM_BUCKET_MAX as _MEDIUM_BUCKET_MAX
|
||||
from packages.domain.asset_scoring import MIN_QUALITY_SCORE as _MIN_QUALITY_SCORE
|
||||
from packages.domain.asset_scoring import OPTIMAL_DURATION_MAX as _OPTIMAL_DURATION_MAX
|
||||
from packages.domain.asset_scoring import OPTIMAL_DURATION_MIN as _OPTIMAL_DURATION_MIN
|
||||
from packages.domain.asset_scoring import SHORT_BUCKET_MAX as _SHORT_BUCKET_MAX
|
||||
from packages.domain.asset_scoring import TARGET_HEIGHT as _TARGET_HEIGHT
|
||||
from packages.domain.asset_scoring import TARGET_WIDTH as _TARGET_WIDTH
|
||||
from packages.domain.asset_scoring import (
|
||||
MEDIUM_BUCKET_MAX as _MEDIUM_BUCKET_MAX,
|
||||
MIN_QUALITY_SCORE as _MIN_QUALITY_SCORE,
|
||||
OPTIMAL_DURATION_MAX as _OPTIMAL_DURATION_MAX,
|
||||
OPTIMAL_DURATION_MIN as _OPTIMAL_DURATION_MIN,
|
||||
SHORT_BUCKET_MAX as _SHORT_BUCKET_MAX,
|
||||
TARGET_HEIGHT as _TARGET_HEIGHT,
|
||||
TARGET_WIDTH as _TARGET_WIDTH,
|
||||
AssetScoreDetail,
|
||||
SmartSelectResult,
|
||||
diverse_selection,
|
||||
|
||||
@@ -29,47 +29,33 @@ from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
|
||||
)
|
||||
from packages.domain.edit_plan import EditPlanStatus
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
from packages.domain.video_filter_builder import (
|
||||
DEFAULT_FPS,
|
||||
DEFAULT_OUTPUT_HEIGHT,
|
||||
DEFAULT_OUTPUT_WIDTH,
|
||||
DEFAULT_TRANSITION_DURATION,
|
||||
ClipFilterChain,
|
||||
build_clip_filter,
|
||||
)
|
||||
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_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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
# ── 常量(向后兼容别名) ──────────────────────────────────────────────────────
|
||||
# 实际定义已迁移至 packages/domain/video_filter_builder.py
|
||||
|
||||
DEFAULT_OUTPUT_WIDTH = 1280
|
||||
DEFAULT_OUTPUT_HEIGHT = 720
|
||||
DEFAULT_FPS = 25
|
||||
DEFAULT_CODEC = "libx264"
|
||||
DEFAULT_CRF = 23
|
||||
DEFAULT_PRESET = "medium"
|
||||
|
||||
# xfade 转场映射:TransitionEffect → FFmpeg xfade transition 名称
|
||||
_XFADE_TRANSITION_MAP: dict[str, str] = {
|
||||
TransitionEffect.FADE: "fade",
|
||||
TransitionEffect.SLIDE_LEFT: "slideleft",
|
||||
TransitionEffect.SLIDE_RIGHT: "slideright",
|
||||
TransitionEffect.DISSOLVE: "dissolve",
|
||||
TransitionEffect.WIPE: "wipeleft",
|
||||
}
|
||||
|
||||
# 转场默认时长(秒)
|
||||
DEFAULT_TRANSITION_DURATION = 0.5
|
||||
|
||||
|
||||
# ── 数据结构 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ClipFilterChain:
|
||||
"""单个片段的滤镜链描述。"""
|
||||
|
||||
clip_id: str
|
||||
input_index: int
|
||||
video_label: str
|
||||
audio_label: str | None
|
||||
filters: list[str]
|
||||
duration: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ComposeCommand:
|
||||
"""完整的 FFmpeg 合成命令描述。"""
|
||||
@@ -401,62 +387,8 @@ class VideoComposeService:
|
||||
output_height: int,
|
||||
fps: int,
|
||||
) -> ClipFilterChain:
|
||||
"""为单个片段构建滤镜链。
|
||||
|
||||
滤镜顺序:
|
||||
1. scale — 等比缩放到目标分辨率(保证覆盖)
|
||||
2. crop — 居中裁剪到目标分辨率
|
||||
3. fps — 统一输出帧率(concat 要求所有输入帧率一致)
|
||||
4. setpts — 重置时间戳 + 偏移
|
||||
5. trim — 视频时长裁剪
|
||||
6. atrim — 音频时长裁剪(如有音频流)
|
||||
"""
|
||||
duration = clip.duration if clip.duration > 0 else 5.0 # 默认 5 秒
|
||||
start = clip.start_time
|
||||
|
||||
filters: list[str] = []
|
||||
|
||||
# 1. scale: 等比缩放(保持比例,不裁剪)
|
||||
filters.append(f"scale={output_width}:{output_height}" f":force_original_aspect_ratio=decrease")
|
||||
|
||||
# 2. pad: 居中+留黑边到目标分辨率(保持原始比例,不裁剪内容)
|
||||
filters.append(f"pad={output_width}:{output_height}:(ow-iw)/2:(oh-ih)/2:black")
|
||||
|
||||
# 3. format: 统一像素格式为 yuv420p(H.264 标准格式,concat 要求所有输入像素格式一致)
|
||||
# 不同素材可能是 yuv420p / yuv422p / yuv444p / nv12 等,必须统一
|
||||
filters.append("format=yuv420p")
|
||||
|
||||
# 4. fps: 统一帧率(concat 要求所有输入帧率一致)
|
||||
# 放在 pad 之后、setpts 之前,确保分辨率和帧率都已统一
|
||||
if fps and fps > 0:
|
||||
filters.append(f"fps={fps}")
|
||||
|
||||
# 3. setpts: 重置时间戳
|
||||
if start > 0:
|
||||
filters.append(f"setpts=PTS-STARTPTS+{start}/TB")
|
||||
else:
|
||||
filters.append("setpts=PTS-STARTPTS")
|
||||
|
||||
# 4. trim: 视频时长
|
||||
filters.append(f"trim=0:{duration}")
|
||||
filters.append("setpts=PTS-STARTPTS") # trim 后需要重置 PTS
|
||||
|
||||
video_label = f"v{input_index}"
|
||||
|
||||
# 5. 音频标签:仅当片段类型可能有音频时才设置
|
||||
# title/subtitle 是纯文字/图片卡片,没有音频流
|
||||
clip_type = clip.clip_type.lower() if clip.clip_type else ""
|
||||
has_audio_stream = clip_type not in ("title", "subtitle")
|
||||
audio_label = f"a{input_index}" if has_audio_stream else None
|
||||
|
||||
return ClipFilterChain(
|
||||
clip_id=clip.id,
|
||||
input_index=input_index,
|
||||
video_label=video_label,
|
||||
audio_label=audio_label,
|
||||
filters=filters,
|
||||
duration=duration,
|
||||
)
|
||||
"""向后兼容:委托给 video_filter_builder.build_clip_filter。"""
|
||||
return build_clip_filter(clip, input_index, output_width, output_height, fps)
|
||||
|
||||
@staticmethod
|
||||
def _build_filter_complex(
|
||||
@@ -466,102 +398,30 @@ class VideoComposeService:
|
||||
transition_duration: float,
|
||||
transitions: list[str],
|
||||
) -> tuple[str, float]:
|
||||
"""构建完整的 filter_complex 字符串。
|
||||
|
||||
策略:
|
||||
- 单片段:直接输出
|
||||
- 多片段 + 全 cut:使用 concat 滤镜(高效)
|
||||
- 多片段 + 有转场:使用 xfade 滤镜链
|
||||
|
||||
返回 (filter_complex_string, estimated_total_duration)。
|
||||
"""
|
||||
n = len(clip_chains)
|
||||
|
||||
if n == 0:
|
||||
return "", 0.0
|
||||
|
||||
# ── 单片段 ─────────────────────────────────────────────────────
|
||||
if n == 1:
|
||||
chain = clip_chains[0]
|
||||
filter_str = _chain_filters(chain.filters, chain.video_label)
|
||||
# 音频
|
||||
if chain.audio_label:
|
||||
filter_str += f";[0:a]{chain.audio_label}"
|
||||
total_duration = chain.duration
|
||||
return filter_str, total_duration
|
||||
|
||||
# ── 检查是否有转场 ─────────────────────────────────────────────
|
||||
has_transitions = any(t != TransitionEffect.CUT and t != "cut" for t in transitions)
|
||||
|
||||
if not has_transitions:
|
||||
return _build_concat_filter(clip_chains)
|
||||
|
||||
# ── 有转场:使用 xfade ─────────────────────────────────────────
|
||||
return _build_xfade_filter(
|
||||
clip_chains=clip_chains,
|
||||
transition_duration=transition_duration,
|
||||
transitions=transitions,
|
||||
)
|
||||
"""向后兼容:委托给 video_filter_builder.build_filter_complex。"""
|
||||
return _build_filter_complex(clip_chains, output_width, output_height, transition_duration, transitions)
|
||||
|
||||
@staticmethod
|
||||
def _has_audio(clip_chains: list[ClipFilterChain]) -> bool:
|
||||
"""是否有任何片段包含音频流。"""
|
||||
return any(c.audio_label is not None for c in clip_chains)
|
||||
"""向后兼容:委托给 video_filter_builder.has_audio。"""
|
||||
return _has_audio_func(clip_chains)
|
||||
|
||||
|
||||
# ── 模块级辅助函数 ────────────────────────────────────────────────────────────
|
||||
# ── 模块级辅助函数(向后兼容别名) ──────────────────────────────────────────
|
||||
# 实际实现已迁移至 packages/domain/video_filter_builder.py
|
||||
# 保留此处别名以兼容现有测试与调用方
|
||||
|
||||
|
||||
def _chain_filters(filters: list[str], output_label: str) -> str:
|
||||
"""将滤镜列表串联为 FFmpeg 滤镜字符串。"""
|
||||
filter_body = ",".join(filters)
|
||||
return f"[0:v]{filter_body}[{output_label}]"
|
||||
"""向后兼容:委托给 video_filter_builder.chain_filters。"""
|
||||
return _chain_filters_func(filters, output_label)
|
||||
|
||||
|
||||
def _build_concat_filter(
|
||||
clip_chains: list[ClipFilterChain],
|
||||
) -> tuple[str, float]:
|
||||
"""构建 concat 滤镜(无转场,高效拼接)。
|
||||
|
||||
格式:
|
||||
[0:v]filters[v0]; [1:v]filters[v1]; ...
|
||||
[v0][v1]...[vN]concat=n=N:v=1:a=0[outv]
|
||||
"""
|
||||
n = len(clip_chains)
|
||||
parts: list[str] = []
|
||||
total_duration = 0.0
|
||||
|
||||
# 每个片段的滤镜链
|
||||
for idx, chain in enumerate(clip_chains):
|
||||
filter_body = ",".join(chain.filters)
|
||||
parts.append(f"[{idx}:v]{filter_body}[{chain.video_label}]")
|
||||
total_duration += chain.duration
|
||||
|
||||
# concat 滤镜
|
||||
concat_inputs = "".join(f"[{c.video_label}]" for c in clip_chains)
|
||||
concat_filter = f"{concat_inputs}concat=n={n}:v=1:a=0[outv]"
|
||||
parts.append(concat_filter)
|
||||
|
||||
# 音频 concat(如果有)— 先统一音频格式再拼接,否则不同采样率/声道会导致concat失败
|
||||
audio_parts: list[str] = []
|
||||
for idx, chain in enumerate(clip_chains):
|
||||
if chain.audio_label:
|
||||
# aformat: 统一采样率48000Hz + 双声道stereo + fltp采样格式(AAC标准格式)
|
||||
audio_filters = [
|
||||
"aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp",
|
||||
f"atrim=0:{chain.duration}",
|
||||
"asetpts=PTS-STARTPTS",
|
||||
]
|
||||
audio_parts.append(f"[{idx}:a]{','.join(audio_filters)}[{chain.audio_label}]")
|
||||
|
||||
if audio_parts:
|
||||
parts.extend(audio_parts)
|
||||
audio_inputs = "".join(f"[{c.audio_label}]" for c in clip_chains if c.audio_label)
|
||||
audio_count = sum(1 for c in clip_chains if c.audio_label)
|
||||
if audio_count > 0:
|
||||
parts.append(f"{audio_inputs}concat=n={audio_count}:v=0:a=1[outa]")
|
||||
|
||||
return ";".join(parts), total_duration
|
||||
"""向后兼容:委托给 video_filter_builder.build_concat_filter。"""
|
||||
return _build_concat_filter_func(clip_chains)
|
||||
|
||||
|
||||
def _build_xfade_filter(
|
||||
@@ -569,80 +429,5 @@ def _build_xfade_filter(
|
||||
transition_duration: float,
|
||||
transitions: list[str],
|
||||
) -> tuple[str, float]:
|
||||
"""构建 xfade 转场滤镜链。
|
||||
|
||||
每两个相邻片段之间插入 xfade 转场。
|
||||
offset = 前一个片段的累积时长 - 转场时长。
|
||||
|
||||
格式(2 片段):
|
||||
[0:v]filters[v0]; [1:v]filters[v1];
|
||||
[v0][v1]xfade=transition=fade:duration=0.5:offset=4.5[outv]
|
||||
|
||||
格式(3+ 片段):
|
||||
[v0][v1]xfade=...[tmp1]; [tmp1][v2]xfade=...[outv]
|
||||
"""
|
||||
n = len(clip_chains)
|
||||
parts: list[str] = []
|
||||
total_duration = 0.0
|
||||
|
||||
# 每个片段的滤镜链
|
||||
for idx, chain in enumerate(clip_chains):
|
||||
filter_body = ",".join(chain.filters)
|
||||
parts.append(f"[{idx}:v]{filter_body}[{chain.video_label}]")
|
||||
total_duration += chain.duration
|
||||
|
||||
# xfade 链
|
||||
if n == 1:
|
||||
# 单片段不需要 xfade
|
||||
parts.append(f"[{clip_chains[0].video_label}]copy[outv]")
|
||||
return ";".join(parts), total_duration
|
||||
|
||||
# 计算每个转场的 offset
|
||||
cumulative = 0.0
|
||||
prev_label = clip_chains[0].video_label
|
||||
|
||||
for i in range(1, n):
|
||||
cumulative += clip_chains[i - 1].duration
|
||||
offset = max(0.0, cumulative - transition_duration * i)
|
||||
|
||||
# 获取转场类型
|
||||
transition = transitions[i] if i < len(transitions) else "cut"
|
||||
xfade_transition = _XFADE_TRANSITION_MAP.get(transition, "fade")
|
||||
|
||||
if i == n - 1:
|
||||
# 最后一个转场,输出到 [outv]
|
||||
out_label = "outv"
|
||||
else:
|
||||
out_label = f"xf{i}"
|
||||
|
||||
parts.append(
|
||||
f"[{prev_label}][{clip_chains[i].video_label}]"
|
||||
f"xfade=transition={xfade_transition}"
|
||||
f":duration={transition_duration}"
|
||||
f":offset={offset:.3f}"
|
||||
f"[{out_label}]"
|
||||
)
|
||||
prev_label = out_label
|
||||
|
||||
# 总时长需要减去转场重叠部分
|
||||
total_duration -= transition_duration * (n - 1)
|
||||
|
||||
# 音频:先 aformat 归一化再 concat(不同采样率/声道/采样格式会导致concat失败)
|
||||
audio_chains_with_label = [(c, c.audio_label) for c in clip_chains if c.audio_label]
|
||||
if len(audio_chains_with_label) >= 2:
|
||||
normalized_audio_labels: list[str] = []
|
||||
for chain, _ in audio_chains_with_label:
|
||||
norm_label = f"anorm_{chain.video_label}"
|
||||
audio_filters = [
|
||||
"aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp",
|
||||
f"atrim=0:{chain.duration}",
|
||||
"asetpts=PTS-STARTPTS",
|
||||
]
|
||||
parts.append(f"[{chain.audio_label}]{','.join(audio_filters)}[{norm_label}]")
|
||||
normalized_audio_labels.append(norm_label)
|
||||
audio_inputs = "".join(f"[{label}]" for label in normalized_audio_labels)
|
||||
parts.append(f"{audio_inputs}concat=n={len(normalized_audio_labels)}:v=0:a=1[outa]")
|
||||
elif len(audio_chains_with_label) == 1:
|
||||
parts.append(f"[{audio_chains_with_label[0][0].audio_label}]acopy[outa]")
|
||||
|
||||
return ";".join(parts), max(0.0, total_duration)
|
||||
"""向后兼容:委托给 video_filter_builder.build_xfade_filter。"""
|
||||
return _build_xfade_filter_func(clip_chains, transition_duration, transitions)
|
||||
|
||||
@@ -2,7 +2,12 @@
|
||||
* 混剪单图层配置区
|
||||
*/
|
||||
import React from "react"
|
||||
import type { PipLayer, PipAnimType, PipSlideDirection, PipGridPosition } from "@/pages/editing-planner/types"
|
||||
import type {
|
||||
PipLayer,
|
||||
PipAnimType,
|
||||
PipSlideDirection,
|
||||
PipGridPosition,
|
||||
} from "@/pages/editing-planner/types"
|
||||
import {
|
||||
GRID_POSITIONS,
|
||||
ANIM_OPTIONS,
|
||||
|
||||
@@ -3,126 +3,43 @@
|
||||
* 卡片视图展示用户已保存的剪辑模板
|
||||
* 支持搜索、分类筛选、编辑/复制/删除/使用模板生成
|
||||
*/
|
||||
import React, { useState } from "react"
|
||||
import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import {
|
||||
Typography,
|
||||
Card,
|
||||
Input,
|
||||
Select,
|
||||
Tag,
|
||||
Button,
|
||||
Space,
|
||||
Empty,
|
||||
Spin,
|
||||
Tooltip,
|
||||
message,
|
||||
Popconfirm,
|
||||
Row,
|
||||
Col,
|
||||
} from "antd"
|
||||
import React from "react"
|
||||
import { Typography, Input, Select, Button, Empty, Spin, Row, Col } from "antd"
|
||||
import {
|
||||
SearchOutlined,
|
||||
EditOutlined,
|
||||
CopyOutlined,
|
||||
DeleteOutlined,
|
||||
VideoCameraOutlined,
|
||||
AppstoreOutlined,
|
||||
PlusOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { useNavigate } from "react-router-dom"
|
||||
import {
|
||||
getEditingTemplates,
|
||||
getTemplateCategories,
|
||||
deleteEditingTemplate,
|
||||
createEditingTemplate,
|
||||
MODE_LABELS,
|
||||
MODE_COLORS,
|
||||
type EditingTemplate,
|
||||
type TemplateMode,
|
||||
} from "@/api/editing-planner"
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
import { useMyTemplates } from "./hooks/useMyTemplates"
|
||||
import { TemplateCard } from "./components/TemplateCard"
|
||||
import "./MyTemplates.css"
|
||||
|
||||
const { Title, Text } = Typography
|
||||
|
||||
const MyTemplates: React.FC = () => {
|
||||
const navigate = useNavigate()
|
||||
const queryClient = useQueryClient()
|
||||
const {
|
||||
searchText,
|
||||
setSearchText,
|
||||
filterCategory,
|
||||
setFilterCategory,
|
||||
templates,
|
||||
categories,
|
||||
isLoading,
|
||||
handleCopy,
|
||||
handleDelete,
|
||||
} = useMyTemplates()
|
||||
|
||||
const [searchText, setSearchText] = useState("")
|
||||
const [filterCategory, setFilterCategory] = useState("")
|
||||
|
||||
/* ── 数据查询 ── */
|
||||
const { data: templates = [], isLoading } = useQuery({
|
||||
queryKey: ["editing-templates", filterCategory, searchText],
|
||||
queryFn: () =>
|
||||
getEditingTemplates({
|
||||
category: filterCategory || undefined,
|
||||
tag: searchText || undefined,
|
||||
}),
|
||||
})
|
||||
|
||||
const { data: categories = [] } = useQuery({
|
||||
queryKey: ["template-categories"],
|
||||
queryFn: getTemplateCategories,
|
||||
})
|
||||
|
||||
/* ── Mutations ── */
|
||||
const deleteMutation = useMutation({
|
||||
mutationFn: deleteEditingTemplate,
|
||||
onSuccess: () => {
|
||||
message.success("模板已删除")
|
||||
queryClient.invalidateQueries({ queryKey: ["editing-templates"] })
|
||||
},
|
||||
onError: (err: unknown) => {
|
||||
if (!(err as { __msgShown?: boolean })?.__msgShown) message.error("删除失败")
|
||||
},
|
||||
})
|
||||
|
||||
const copyMutation = useMutation({
|
||||
mutationFn: (tpl: EditingTemplate) =>
|
||||
createEditingTemplate({
|
||||
name: `${tpl.name}(副本)`,
|
||||
mode: tpl.mode,
|
||||
category: tpl.category,
|
||||
tags: tpl.tags,
|
||||
title_config: tpl.title_config,
|
||||
subtitle_config: tpl.subtitle_config,
|
||||
bgm_config: tpl.bgm_config,
|
||||
estimated_duration:
|
||||
tpl.estimated_duration ??
|
||||
Math.round(
|
||||
tpl.segments.reduce((s, seg) => s + (seg.duration_min + seg.duration_max) / 2, 0),
|
||||
),
|
||||
segments: tpl.segments.map(({ id: _id, ...rest }) => rest),
|
||||
}),
|
||||
onSuccess: () => {
|
||||
message.success("模板已复制")
|
||||
queryClient.invalidateQueries({ queryKey: ["editing-templates"] })
|
||||
},
|
||||
onError: (err: unknown) => {
|
||||
if (!(err as { __msgShown?: boolean })?.__msgShown) message.error("复制失败")
|
||||
},
|
||||
})
|
||||
|
||||
/* ── 操作 ── */
|
||||
const handleEdit = (tpl: EditingTemplate) => {
|
||||
navigate(`/editing-planner?template=${tpl.id}`)
|
||||
}
|
||||
|
||||
const handleGenerate = (tpl: EditingTemplate) => {
|
||||
// 跳转到智能剪辑页面,统一从智能剪辑出片
|
||||
navigate(`/generate?templateId=${tpl.id}`)
|
||||
}
|
||||
|
||||
const handleCopy = (tpl: EditingTemplate) => {
|
||||
copyMutation.mutate(tpl)
|
||||
}
|
||||
|
||||
const handleDelete = (id: string) => {
|
||||
deleteMutation.mutate(id)
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="mt-page">
|
||||
{/* 页面头部 */}
|
||||
@@ -179,69 +96,13 @@ const MyTemplates: React.FC = () => {
|
||||
<Row gutter={[16, 16]}>
|
||||
{templates.map((tpl) => (
|
||||
<Col key={tpl.id} xs={24} sm={12} md={8} lg={6}>
|
||||
<Card
|
||||
className="mt-card"
|
||||
hoverable
|
||||
actions={[
|
||||
<Tooltip title="编辑" key="edit">
|
||||
<EditOutlined onClick={() => handleEdit(tpl)} />
|
||||
</Tooltip>,
|
||||
<Tooltip title="复制" key="copy">
|
||||
<CopyOutlined onClick={() => handleCopy(tpl)} />
|
||||
</Tooltip>,
|
||||
<Tooltip title="使用模板生成" key="generate">
|
||||
<VideoCameraOutlined onClick={() => handleGenerate(tpl)} />
|
||||
</Tooltip>,
|
||||
<Popconfirm
|
||||
key="delete"
|
||||
title="确定删除此模板?"
|
||||
onConfirm={() => handleDelete(tpl.id)}
|
||||
okText="删除"
|
||||
cancelText="取消"
|
||||
>
|
||||
<Tooltip title="删除">
|
||||
<DeleteOutlined style={{ color: "#ff4d4f" }} />
|
||||
</Tooltip>
|
||||
</Popconfirm>,
|
||||
]}
|
||||
>
|
||||
<div className="mt-card-head">
|
||||
<Text strong ellipsis style={{ fontSize: 15 }}>
|
||||
{tpl.name}
|
||||
</Text>
|
||||
<Tag color={MODE_COLORS[tpl.mode as TemplateMode] || "default"}>
|
||||
{MODE_LABELS[tpl.mode as TemplateMode] || tpl.mode}
|
||||
</Tag>
|
||||
<Tag color="green">用户自制</Tag>
|
||||
</div>
|
||||
|
||||
<div className="mt-card-meta">
|
||||
<Text type="secondary" style={{ fontSize: 12 }}>
|
||||
{tpl.segments.length} 片段 · 预估 ~{tpl.estimated_duration}s
|
||||
</Text>
|
||||
{tpl.category && (
|
||||
<Tag style={{ fontSize: 11, marginTop: 4 }}>{tpl.category}</Tag>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{tpl.tags.length > 0 && (
|
||||
<div className="mt-card-tags">
|
||||
{tpl.tags.map((tag) => (
|
||||
<Tag key={tag} style={{ fontSize: 11 }}>
|
||||
{tag}
|
||||
</Tag>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div className="mt-card-config">
|
||||
<Space size={4} wrap>
|
||||
{tpl.title_config.ai_auto_select && <Tag color="cyan">AI标题</Tag>}
|
||||
{tpl.subtitle_config.enabled && <Tag color="geekblue">字幕</Tag>}
|
||||
{tpl.bgm_config.enabled && <Tag color="pink">BGM</Tag>}
|
||||
</Space>
|
||||
</div>
|
||||
</Card>
|
||||
<TemplateCard
|
||||
tpl={tpl}
|
||||
onEdit={handleEdit}
|
||||
onCopy={handleCopy}
|
||||
onGenerate={handleGenerate}
|
||||
onDelete={handleDelete}
|
||||
/>
|
||||
</Col>
|
||||
))}
|
||||
</Row>
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
import React from "react"
|
||||
import { Card, Tag, Tooltip, Popconfirm, Space, Typography } from "antd"
|
||||
import { EditOutlined, CopyOutlined, DeleteOutlined, VideoCameraOutlined } from "@ant-design/icons"
|
||||
import {
|
||||
MODE_LABELS,
|
||||
MODE_COLORS,
|
||||
type EditingTemplate,
|
||||
type TemplateMode,
|
||||
} from "@/api/editing-planner"
|
||||
|
||||
const { Text } = Typography
|
||||
|
||||
interface TemplateCardProps {
|
||||
tpl: EditingTemplate
|
||||
onEdit: (tpl: EditingTemplate) => void
|
||||
onCopy: (tpl: EditingTemplate) => void
|
||||
onGenerate: (tpl: EditingTemplate) => void
|
||||
onDelete: (id: string) => void
|
||||
}
|
||||
|
||||
/**
|
||||
* 单个模板卡片组件
|
||||
*/
|
||||
export const TemplateCard: React.FC<TemplateCardProps> = ({
|
||||
tpl,
|
||||
onEdit,
|
||||
onCopy,
|
||||
onGenerate,
|
||||
onDelete,
|
||||
}) => (
|
||||
<Card
|
||||
className="mt-card"
|
||||
hoverable
|
||||
actions={[
|
||||
<Tooltip title="编辑" key="edit">
|
||||
<EditOutlined onClick={() => onEdit(tpl)} />
|
||||
</Tooltip>,
|
||||
<Tooltip title="复制" key="copy">
|
||||
<CopyOutlined onClick={() => onCopy(tpl)} />
|
||||
</Tooltip>,
|
||||
<Tooltip title="使用模板生成" key="generate">
|
||||
<VideoCameraOutlined onClick={() => onGenerate(tpl)} />
|
||||
</Tooltip>,
|
||||
<Popconfirm
|
||||
key="delete"
|
||||
title="确定删除此模板?"
|
||||
onConfirm={() => onDelete(tpl.id)}
|
||||
okText="删除"
|
||||
cancelText="取消"
|
||||
>
|
||||
<Tooltip title="删除">
|
||||
<DeleteOutlined style={{ color: "#ff4d4f" }} />
|
||||
</Tooltip>
|
||||
</Popconfirm>,
|
||||
]}
|
||||
>
|
||||
<div className="mt-card-head">
|
||||
<Text strong ellipsis style={{ fontSize: 15 }}>
|
||||
{tpl.name}
|
||||
</Text>
|
||||
<Tag color={MODE_COLORS[tpl.mode as TemplateMode] || "default"}>
|
||||
{MODE_LABELS[tpl.mode as TemplateMode] || tpl.mode}
|
||||
</Tag>
|
||||
<Tag color="green">用户自制</Tag>
|
||||
</div>
|
||||
|
||||
<div className="mt-card-meta">
|
||||
<Text type="secondary" style={{ fontSize: 12 }}>
|
||||
{tpl.segments.length} 片段 · 预估 ~{tpl.estimated_duration}s
|
||||
</Text>
|
||||
{tpl.category && <Tag style={{ fontSize: 11, marginTop: 4 }}>{tpl.category}</Tag>}
|
||||
</div>
|
||||
|
||||
{tpl.tags.length > 0 && (
|
||||
<div className="mt-card-tags">
|
||||
{tpl.tags.map((tag) => (
|
||||
<Tag key={tag} style={{ fontSize: 11 }}>
|
||||
{tag}
|
||||
</Tag>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div className="mt-card-config">
|
||||
<Space size={4} wrap>
|
||||
{tpl.title_config.ai_auto_select && <Tag color="cyan">AI标题</Tag>}
|
||||
{tpl.subtitle_config.enabled && <Tag color="geekblue">字幕</Tag>}
|
||||
{tpl.bgm_config.enabled && <Tag color="pink">BGM</Tag>}
|
||||
</Space>
|
||||
</div>
|
||||
</Card>
|
||||
)
|
||||
@@ -0,0 +1,100 @@
|
||||
import { useState } from "react"
|
||||
import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import {
|
||||
getEditingTemplates,
|
||||
getTemplateCategories,
|
||||
deleteEditingTemplate,
|
||||
createEditingTemplate,
|
||||
type EditingTemplate,
|
||||
} from "@/api/editing-planner"
|
||||
|
||||
/**
|
||||
* 我的模板数据 Hook
|
||||
* 封装模板列表查询、筛选、删除、复制等数据操作
|
||||
*/
|
||||
export function useMyTemplates() {
|
||||
const queryClient = useQueryClient()
|
||||
const [searchText, setSearchText] = useState("")
|
||||
const [filterCategory, setFilterCategory] = useState("")
|
||||
|
||||
/* 模板列表 */
|
||||
const { data: templates = [], isLoading } = useQuery({
|
||||
queryKey: ["editing-templates", filterCategory, searchText],
|
||||
queryFn: () =>
|
||||
getEditingTemplates({
|
||||
category: filterCategory || undefined,
|
||||
tag: searchText || undefined,
|
||||
}),
|
||||
})
|
||||
|
||||
/* 分类列表 */
|
||||
const { data: categories = [] } = useQuery({
|
||||
queryKey: ["template-categories"],
|
||||
queryFn: getTemplateCategories,
|
||||
})
|
||||
|
||||
/* 删除 mutation */
|
||||
const deleteMutation = useMutation({
|
||||
mutationFn: deleteEditingTemplate,
|
||||
onSuccess: () => {
|
||||
message.success("模板已删除")
|
||||
queryClient.invalidateQueries({ queryKey: ["editing-templates"] })
|
||||
},
|
||||
onError: (err: unknown) => {
|
||||
if (!(err as { __msgShown?: boolean })?.__msgShown) message.error("删除失败")
|
||||
},
|
||||
})
|
||||
|
||||
/* 复制 mutation */
|
||||
const copyMutation = useMutation({
|
||||
mutationFn: (tpl: EditingTemplate) =>
|
||||
createEditingTemplate({
|
||||
name: `${tpl.name}(副本)`,
|
||||
mode: tpl.mode,
|
||||
category: tpl.category,
|
||||
tags: tpl.tags,
|
||||
title_config: tpl.title_config,
|
||||
subtitle_config: tpl.subtitle_config,
|
||||
bgm_config: tpl.bgm_config,
|
||||
estimated_duration:
|
||||
tpl.estimated_duration ??
|
||||
Math.round(
|
||||
tpl.segments.reduce((s, seg) => s + (seg.duration_min + seg.duration_max) / 2, 0),
|
||||
),
|
||||
segments: tpl.segments.map(({ id: _id, ...rest }) => rest),
|
||||
}),
|
||||
onSuccess: () => {
|
||||
message.success("模板已复制")
|
||||
queryClient.invalidateQueries({ queryKey: ["editing-templates"] })
|
||||
},
|
||||
onError: (err: unknown) => {
|
||||
if (!(err as { __msgShown?: boolean })?.__msgShown) message.error("复制失败")
|
||||
},
|
||||
})
|
||||
|
||||
const handleCopy = (tpl: EditingTemplate) => {
|
||||
copyMutation.mutate(tpl)
|
||||
}
|
||||
|
||||
const handleDelete = (id: string) => {
|
||||
deleteMutation.mutate(id)
|
||||
}
|
||||
|
||||
return {
|
||||
// 状态
|
||||
searchText,
|
||||
setSearchText,
|
||||
filterCategory,
|
||||
setFilterCategory,
|
||||
// 数据
|
||||
templates,
|
||||
categories,
|
||||
isLoading,
|
||||
// 操作
|
||||
handleCopy,
|
||||
handleDelete,
|
||||
isDeleting: deleteMutation.isPending,
|
||||
isCopying: copyMutation.isPending,
|
||||
}
|
||||
}
|
||||
@@ -6,708 +6,99 @@
|
||||
* - 复制模板 / 从模板生成
|
||||
* - 卡片网格布局 + 类型筛选 + 搜索 + 收藏
|
||||
*/
|
||||
import React, { useState, useMemo, useCallback } from "react"
|
||||
import { useNavigate } from "react-router-dom"
|
||||
import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { Button, message, Pagination, Tooltip, Tag, Descriptions } from "antd"
|
||||
import {
|
||||
LoadingOutlined,
|
||||
ExclamationCircleOutlined,
|
||||
InboxOutlined,
|
||||
SearchOutlined,
|
||||
CopyOutlined,
|
||||
ThunderboltOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import {
|
||||
getTemplates,
|
||||
getTemplate,
|
||||
toggleFavoriteTemplate,
|
||||
copyTemplate,
|
||||
type TemplateItem,
|
||||
type TemplateListParams,
|
||||
type TemplateSegment,
|
||||
} from "@/api/templates"
|
||||
import React from "react"
|
||||
import { useTemplateLibrary } from "./hooks/useTemplateLibrary"
|
||||
import { useTemplateDetail } from "./hooks/useTemplateDetail"
|
||||
import { TemplateHeader } from "./components/template-library/TemplateHeader"
|
||||
import { TemplateToolbar } from "./components/template-library/TemplateToolbar"
|
||||
import { TemplateGrid } from "./components/template-library/TemplateGrid"
|
||||
import { TemplateDetailModal } from "./components/template-library/TemplateDetailModal"
|
||||
import "./templates.css"
|
||||
|
||||
/* ============================================================
|
||||
* 类型定义
|
||||
* ============================================================ */
|
||||
|
||||
/** 模板类型 */
|
||||
type EditTemplateType = "口播" | "种草" | "产品" | "品牌" | "混剪" | "Vlog"
|
||||
|
||||
/* ============================================================
|
||||
* 模板类型配置
|
||||
* ============================================================ */
|
||||
|
||||
const TEMPLATE_TYPES: Array<{
|
||||
type: EditTemplateType | "全部"
|
||||
label: string
|
||||
icon: string
|
||||
color: string
|
||||
}> = [
|
||||
{ type: "全部", label: "全部", icon: "📋", color: "#6366f1" },
|
||||
{ type: "口播", label: "口播", icon: "🎙️", color: "#6366f1" },
|
||||
{ type: "种草", label: "种草", icon: "🌱", color: "#10b981" },
|
||||
{ type: "产品", label: "产品", icon: "📦", color: "#0ea5e9" },
|
||||
{ type: "品牌", label: "品牌", icon: "🏷️", color: "#f59e0b" },
|
||||
{ type: "混剪", label: "混剪", icon: "🎬", color: "#8b5cf6" },
|
||||
{ type: "Vlog", label: "Vlog", icon: "📹", color: "#ec4899" },
|
||||
]
|
||||
|
||||
/** 时长筛选选项 */
|
||||
const DURATION_OPTIONS: Array<{
|
||||
value: "" | "short" | "medium" | "long"
|
||||
label: string
|
||||
}> = [
|
||||
{ value: "", label: "全部时长" },
|
||||
{ value: "short", label: "30秒以内" },
|
||||
{ value: "medium", label: "30秒-2分钟" },
|
||||
{ value: "long", label: "2分钟以上" },
|
||||
]
|
||||
|
||||
/* ============================================================
|
||||
* 辅助函数
|
||||
* ============================================================ */
|
||||
|
||||
/** 获取类型对应颜色 */
|
||||
const getTypeColor = (type: string): string => {
|
||||
const found = TEMPLATE_TYPES.find((t) => t.type === type)
|
||||
return found?.color ?? "#6366f1"
|
||||
}
|
||||
|
||||
/** 根据 category 生成占位渐变色 */
|
||||
const gradientForCategory = (category: string): string => {
|
||||
const gradients: Record<string, string> = {
|
||||
口播: "linear-gradient(135deg, #6366f1, #8b5cf6)",
|
||||
种草: "linear-gradient(135deg, #10b981, #059669)",
|
||||
产品: "linear-gradient(135deg, #0ea5e9, #0284c7)",
|
||||
品牌: "linear-gradient(135deg, #f59e0b, #d97706)",
|
||||
混剪: "linear-gradient(135deg, #8b5cf6, #6d28d9)",
|
||||
Vlog: "linear-gradient(135deg, #ec4899, #db2777)",
|
||||
}
|
||||
return gradients[category] ?? "linear-gradient(135deg, #6366f1, #8b5cf6)"
|
||||
}
|
||||
|
||||
/** 格式化时长 */
|
||||
const formatDuration = (seconds: number | undefined | null): string => {
|
||||
if (!seconds || seconds <= 0) return "0秒"
|
||||
const totalSec = Math.round(seconds)
|
||||
const m = Math.floor(totalSec / 60)
|
||||
const s = totalSec % 60
|
||||
if (m === 0) return `${s}秒`
|
||||
return `${m}分${s > 0 ? `${s}秒` : ""}`
|
||||
}
|
||||
|
||||
/** 配置展示字段(formatConfig 提取通用配置的可读属性) */
|
||||
interface ConfigDisplayFields {
|
||||
font_size?: string | number
|
||||
font_family?: string
|
||||
color?: string
|
||||
position?: string
|
||||
volume?: string | number
|
||||
name?: string
|
||||
}
|
||||
|
||||
/** 格式化配置对象为可读文本 */
|
||||
const formatConfig = (config?: object): string => {
|
||||
if (!config || Object.keys(config).length === 0) return "默认"
|
||||
const c = config as ConfigDisplayFields
|
||||
const parts: string[] = []
|
||||
if (c.font_size) parts.push(`字号: ${c.font_size}`)
|
||||
if (c.font_family) parts.push(`字体: ${c.font_family}`)
|
||||
if (c.color) parts.push(`颜色: ${c.color}`)
|
||||
if (c.position) parts.push(`位置: ${c.position}`)
|
||||
if (c.volume !== undefined) parts.push(`音量: ${c.volume}%`)
|
||||
if (c.name) parts.push(String(c.name))
|
||||
return parts.length > 0 ? parts.join(" / ") : JSON.stringify(config)
|
||||
}
|
||||
|
||||
/** 素材类型标签 */
|
||||
const MATERIAL_TYPE_LABELS: Record<string, string> = {
|
||||
video: "视频",
|
||||
image: "图片",
|
||||
audio: "音频",
|
||||
voiceover: "配音",
|
||||
subtitle: "字幕",
|
||||
null: "不限",
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
* 模板详情弹窗组件
|
||||
* ============================================================ */
|
||||
|
||||
interface TemplateDetailModalProps {
|
||||
template: TemplateItem
|
||||
isFavorite: boolean
|
||||
onClose: () => void
|
||||
onToggleFavorite: (id: string) => void
|
||||
onUse: (template: TemplateItem) => void
|
||||
onCopy: (template: TemplateItem) => void
|
||||
}
|
||||
|
||||
const TemplateDetailModal: React.FC<TemplateDetailModalProps> = ({
|
||||
template,
|
||||
isFavorite,
|
||||
onClose,
|
||||
onToggleFavorite,
|
||||
onUse,
|
||||
onCopy,
|
||||
}) => {
|
||||
const segments = template.segments ?? []
|
||||
const totalSegmentDuration = segments.reduce(
|
||||
(sum, s) => sum + (s.duration_min + s.duration_max) / 2,
|
||||
0,
|
||||
)
|
||||
|
||||
return (
|
||||
<div className="xx-template-modal-overlay" onClick={onClose}>
|
||||
<div
|
||||
className="xx-template-modal xx-template-modal-wide"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
>
|
||||
{/* 关闭按钮 */}
|
||||
<button className="xx-template-modal-close" onClick={onClose} title="关闭">
|
||||
✕
|
||||
</button>
|
||||
|
||||
{/* 预览区域 */}
|
||||
<div
|
||||
className="xx-template-modal-preview"
|
||||
style={{ background: gradientForCategory(template.category) }}
|
||||
>
|
||||
{template.thumbnail_url ? (
|
||||
<img
|
||||
src={template.thumbnail_url}
|
||||
alt={template.name}
|
||||
className="xx-template-modal-thumb-img"
|
||||
/>
|
||||
) : (
|
||||
<div className="xx-template-modal-preview-content">
|
||||
<span className="xx-template-preview-icon">
|
||||
{TEMPLATE_TYPES.find((t) => t.type === template.category)?.icon ?? "📋"}
|
||||
</span>
|
||||
<span className="xx-template-preview-title">{template.name}</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 内容区域 */}
|
||||
<div className="xx-template-modal-content">
|
||||
{/* 标题行 */}
|
||||
<div className="xx-template-modal-title-row">
|
||||
<h3>{template.name}</h3>
|
||||
<span
|
||||
className="xx-template-modal-type-badge"
|
||||
style={{
|
||||
color: getTypeColor(template.category),
|
||||
background: `${getTypeColor(template.category)}18`,
|
||||
}}
|
||||
>
|
||||
{TEMPLATE_TYPES.find((t) => t.type === template.category)?.icon} {template.category}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{/* 描述 */}
|
||||
<p className="xx-template-modal-desc">{template.description}</p>
|
||||
|
||||
{/* 标签 */}
|
||||
{(template.tags?.length ?? 0) > 0 && (
|
||||
<div className="xx-template-modal-tags">
|
||||
{template.tags!.map((tag) => (
|
||||
<span key={tag} className="xx-template-modal-tag">
|
||||
#{tag}
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 基本信息 */}
|
||||
<Descriptions
|
||||
column={2}
|
||||
size="small"
|
||||
className="xx-template-modal-desc-table"
|
||||
items={[
|
||||
{
|
||||
key: "duration",
|
||||
label: "目标时长",
|
||||
children: formatDuration(template.estimated_duration ?? template.target_duration),
|
||||
},
|
||||
{
|
||||
key: "clips",
|
||||
label: "片段数量",
|
||||
children: `${template.clip_count} 个`,
|
||||
},
|
||||
{
|
||||
key: "ratio",
|
||||
label: "视频比例",
|
||||
children: template.aspect_ratio ?? "16:9",
|
||||
},
|
||||
{
|
||||
key: "usage",
|
||||
label: "使用次数",
|
||||
children: `${template.usage_count ?? 0} 次`,
|
||||
},
|
||||
]}
|
||||
/>
|
||||
|
||||
{/* 素材规则(片段配置) */}
|
||||
{segments.length > 0 && (
|
||||
<div className="xx-template-modal-section">
|
||||
<h4>🎬 素材规则</h4>
|
||||
<div className="xx-template-modal-clip-list">
|
||||
{segments
|
||||
.sort((a, b) => a.segment_order - b.segment_order)
|
||||
.map((seg: TemplateSegment, idx: number) => (
|
||||
<div key={seg.id ?? idx} className="xx-template-modal-clip-item">
|
||||
<span className="xx-template-modal-clip-order">#{seg.segment_order}</span>
|
||||
<span
|
||||
className="xx-template-modal-clip-badge"
|
||||
style={{
|
||||
color: seg.material_type ? getTypeColor(seg.material_type) : "#64748b",
|
||||
background: seg.material_type
|
||||
? `${getTypeColor(seg.material_type)}18`
|
||||
: "#f1f5f9",
|
||||
}}
|
||||
>
|
||||
{MATERIAL_TYPE_LABELS[seg.material_type ?? "null"] ??
|
||||
seg.material_type ??
|
||||
"不限"}
|
||||
</span>
|
||||
<span className="xx-template-modal-clip-desc">
|
||||
{seg.description || `片段 ${seg.segment_order}`}
|
||||
</span>
|
||||
<Tooltip title={`时长范围: ${seg.duration_min}秒 - ${seg.duration_max}秒`}>
|
||||
<span className="xx-template-modal-clip-duration">
|
||||
{seg.duration_min}-{seg.duration_max}秒
|
||||
</span>
|
||||
</Tooltip>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
<div className="xx-template-modal-total-duration">
|
||||
预估总时长:{formatDuration(Math.round(totalSegmentDuration))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 样式配置 */}
|
||||
<div className="xx-template-modal-section">
|
||||
<h4>🎨 样式配置</h4>
|
||||
<div className="xx-template-modal-style-grid">
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">字幕样式</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{formatConfig(template.subtitle_config)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">标题样式</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{formatConfig(template.title_config)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">BGM 配置</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{formatConfig(template.bgm_config)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">视频比例</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{template.aspect_ratio ?? "16:9"}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 统计信息 */}
|
||||
<div className="xx-template-modal-stats">
|
||||
<span>已使用 {template.usage_count ?? 0} 次</span>
|
||||
<button
|
||||
className={`xx-template-modal-fav-btn${isFavorite ? " is-favorite" : ""}`}
|
||||
onClick={() => onToggleFavorite(template.id)}
|
||||
>
|
||||
{isFavorite ? "★ 已收藏" : "☆ 收藏"}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* 操作按钮 */}
|
||||
<div className="xx-template-modal-actions">
|
||||
<Button icon={<CopyOutlined />} onClick={() => onCopy(template)}>
|
||||
复制模板
|
||||
</Button>
|
||||
<Button type="primary" icon={<ThunderboltOutlined />} onClick={() => onUse(template)}>
|
||||
使用此模板生成
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
* 模板卡片组件
|
||||
* ============================================================ */
|
||||
|
||||
interface TemplateCardProps {
|
||||
template: TemplateItem
|
||||
isFavorite: boolean
|
||||
onPreview: (template: TemplateItem) => void
|
||||
onToggleFavorite: (id: string, e: React.MouseEvent) => void
|
||||
onUse: (template: TemplateItem) => void
|
||||
}
|
||||
|
||||
const TemplateCard: React.FC<TemplateCardProps> = ({
|
||||
template,
|
||||
isFavorite,
|
||||
onPreview,
|
||||
onToggleFavorite,
|
||||
onUse,
|
||||
}) => {
|
||||
return (
|
||||
<div className="xx-template-card" onClick={() => onPreview(template)}>
|
||||
{/* 缩略图 */}
|
||||
<div className="xx-template-thumb">
|
||||
{template.thumbnail_url ? (
|
||||
<img src={template.thumbnail_url} alt={template.name} className="xx-template-thumb-img" />
|
||||
) : (
|
||||
<div
|
||||
className="xx-template-thumb-bg"
|
||||
style={{ background: gradientForCategory(template.category) }}
|
||||
>
|
||||
{(template.description ?? "").slice(0, 80)}
|
||||
{(template.description ?? "").length > 80 ? "..." : ""}
|
||||
</div>
|
||||
)}
|
||||
<div className="xx-template-thumb-overlay" />
|
||||
<div className="xx-template-thumb-name">{template.name}</div>
|
||||
<div className="xx-template-thumb-meta">
|
||||
<span className="xx-template-thumb-duration">
|
||||
{formatDuration(template.estimated_duration ?? template.target_duration)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-preview-hint">点击查看详情</div>
|
||||
<button
|
||||
className={`xx-template-fav-btn${isFavorite ? " is-favorite" : ""}`}
|
||||
onClick={(e) => onToggleFavorite(template.id, e)}
|
||||
title={isFavorite ? "取消收藏" : "收藏"}
|
||||
>
|
||||
{isFavorite ? "★" : "☆"}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* 信息区 */}
|
||||
<div className="xx-template-info">
|
||||
<div className="xx-template-info-top">
|
||||
<span
|
||||
className="xx-template-category-pill"
|
||||
style={{
|
||||
color: getTypeColor(template.category),
|
||||
background: `${getTypeColor(template.category)}18`,
|
||||
}}
|
||||
>
|
||||
{template.category}
|
||||
</span>
|
||||
{(template.tags ?? []).slice(0, 2).map((tag) => (
|
||||
<Tag key={tag} className="xx-template-tag-pill" bordered={false}>
|
||||
{tag}
|
||||
</Tag>
|
||||
))}
|
||||
</div>
|
||||
<p className="xx-template-desc">{template.description ?? ""}</p>
|
||||
<div className="xx-template-meta">
|
||||
<span className="xx-template-usage">已使用 {template.usage_count ?? 0} 次</span>
|
||||
<button
|
||||
className="xx-template-use-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onUse(template)
|
||||
}}
|
||||
>
|
||||
使用此模板
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
* 主组件
|
||||
* ============================================================ */
|
||||
|
||||
const TemplateLibrary: React.FC = () => {
|
||||
const navigate = useNavigate()
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
// 筛选状态
|
||||
const [searchText, setSearchText] = useState("")
|
||||
const [activeType, setActiveType] = useState<EditTemplateType | "全部">("全部")
|
||||
const [durationRange, setDurationRange] = useState<"" | "short" | "medium" | "long">("")
|
||||
const [page, setPage] = useState(1)
|
||||
const [pageSize] = useState(12)
|
||||
|
||||
// 弹窗状态
|
||||
const [previewTemplate, setPreviewTemplate] = useState<TemplateItem | null>(null)
|
||||
const [detailLoading, setDetailLoading] = useState(false)
|
||||
|
||||
// ── 构建查询参数 ──
|
||||
const queryParams: TemplateListParams = useMemo(() => {
|
||||
const params: TemplateListParams = {
|
||||
page,
|
||||
page_size: pageSize,
|
||||
}
|
||||
if (activeType !== "全部") params.category = activeType
|
||||
if (searchText.trim()) params.keyword = searchText.trim()
|
||||
if (durationRange) params.duration_range = durationRange
|
||||
return params
|
||||
}, [page, pageSize, activeType, searchText, durationRange])
|
||||
|
||||
// ── 获取模板列表(后端分页 + 筛选) ──
|
||||
const {
|
||||
data: templateData,
|
||||
templates,
|
||||
totalTemplates,
|
||||
isLoading,
|
||||
isError,
|
||||
error,
|
||||
} = useQuery({
|
||||
queryKey: ["templates", queryParams],
|
||||
queryFn: () => getTemplates(queryParams),
|
||||
staleTime: 30_000,
|
||||
searchText,
|
||||
activeType,
|
||||
durationRange,
|
||||
page,
|
||||
pageSize,
|
||||
setPage,
|
||||
toggleFavorite,
|
||||
handleCopy,
|
||||
handleUse,
|
||||
handleCreate,
|
||||
handleSearchChange,
|
||||
handleCategoryChange,
|
||||
handleDurationChange,
|
||||
} = useTemplateLibrary()
|
||||
|
||||
const {
|
||||
previewTemplate,
|
||||
detailLoading,
|
||||
handlePreview,
|
||||
handleClose,
|
||||
handleToggleFavorite,
|
||||
handleUse: handleUseFromDetail,
|
||||
handleCopy: handleCopyFromDetail,
|
||||
} = useTemplateDetail({
|
||||
onToggleFavorite: (id) => toggleFavorite(id),
|
||||
onUse: handleUse,
|
||||
onCopy: handleCopy,
|
||||
})
|
||||
|
||||
const templates = templateData?.items ?? []
|
||||
const totalTemplates = templateData?.total ?? 0
|
||||
|
||||
// ── 收藏 mutation ──
|
||||
const favMutation = useMutation({
|
||||
mutationFn: toggleFavoriteTemplate,
|
||||
onSuccess: (_data, templateId) => {
|
||||
queryClient.invalidateQueries({ queryKey: ["templates"] })
|
||||
if (previewTemplate && previewTemplate.id === templateId) {
|
||||
setPreviewTemplate((prev) => (prev ? { ...prev, is_favorite: !prev.is_favorite } : prev))
|
||||
}
|
||||
},
|
||||
})
|
||||
|
||||
// ── 复制模板 mutation ──
|
||||
const copyMutation = useMutation({
|
||||
mutationFn: copyTemplate,
|
||||
onSuccess: (data) => {
|
||||
message.success(`模板「${data.name}」已复制到「我的模板」`)
|
||||
queryClient.invalidateQueries({ queryKey: ["templates"] })
|
||||
},
|
||||
onError: () => {
|
||||
message.error("复制模板失败,请稍后重试")
|
||||
},
|
||||
})
|
||||
|
||||
/** 切换收藏 */
|
||||
const toggleFavorite = useCallback(
|
||||
(id: string, e?: React.MouseEvent) => {
|
||||
e?.stopPropagation()
|
||||
favMutation.mutate(id)
|
||||
},
|
||||
[favMutation],
|
||||
)
|
||||
|
||||
/** 点击卡片 → 获取详情并展示弹窗 */
|
||||
const handlePreview = useCallback(async (template: TemplateItem) => {
|
||||
setDetailLoading(true)
|
||||
setPreviewTemplate(template)
|
||||
try {
|
||||
const detail = await getTemplate(template.id)
|
||||
setPreviewTemplate(detail)
|
||||
} catch {
|
||||
// 详情加载失败时使用列表数据
|
||||
message.warning("模板详情加载失败,显示摘要信息")
|
||||
} finally {
|
||||
setDetailLoading(false)
|
||||
}
|
||||
}, [])
|
||||
|
||||
/** 复制模板 */
|
||||
const handleCopy = useCallback(
|
||||
(template: TemplateItem) => {
|
||||
copyMutation.mutate(template.id)
|
||||
},
|
||||
[copyMutation],
|
||||
)
|
||||
|
||||
/** 使用模板 → 进入剪辑编辑器配置 */
|
||||
const handleUse = useCallback(
|
||||
(template: TemplateItem) => {
|
||||
navigate(`/app/editing-planner?templateId=${template.id}`)
|
||||
},
|
||||
[navigate],
|
||||
)
|
||||
|
||||
/** 搜索防抖处理 */
|
||||
const handleSearchChange = useCallback((e: React.ChangeEvent<HTMLInputElement>) => {
|
||||
setSearchText(e.target.value)
|
||||
setPage(1)
|
||||
}, [])
|
||||
|
||||
/** 切换分类 */
|
||||
const handleCategoryChange = useCallback((type: EditTemplateType | "全部") => {
|
||||
setActiveType(type)
|
||||
setPage(1)
|
||||
}, [])
|
||||
|
||||
/** 切换时长筛选 */
|
||||
const handleDurationChange = useCallback((value: "" | "short" | "medium" | "long") => {
|
||||
setDurationRange(value)
|
||||
setPage(1)
|
||||
}, [])
|
||||
|
||||
// ── Loading 状态 ──
|
||||
if (isLoading) {
|
||||
return (
|
||||
<div className="xx-templates-page">
|
||||
<div className="xx-templates-empty">
|
||||
<div className="xx-templates-empty-icon">
|
||||
<LoadingOutlined />
|
||||
</div>
|
||||
<h3>加载模板中...</h3>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
// ── Error 状态 ──
|
||||
if (isError) {
|
||||
return (
|
||||
<div className="xx-templates-page">
|
||||
<div className="xx-templates-empty">
|
||||
<div className="xx-templates-empty-icon">
|
||||
<ExclamationCircleOutlined />
|
||||
</div>
|
||||
<h3>加载失败</h3>
|
||||
<p>{error?.message || "网络异常,请稍后重试"}</p>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="xx-templates-page">
|
||||
{/* ── 页面头部 ──────────────────────────────────────────── */}
|
||||
<div className="xx-templates-header">
|
||||
<div className="xx-templates-header-text">
|
||||
<h2>模板库</h2>
|
||||
<p>选择模板快速创建,支持自定义修改</p>
|
||||
</div>
|
||||
<Button type="primary" onClick={() => navigate("/app/editing-planner")}>
|
||||
+ 创建模板
|
||||
</Button>
|
||||
</div>
|
||||
{/* 页面头部 */}
|
||||
<TemplateHeader onCreateClick={handleCreate} />
|
||||
|
||||
{/* ── 工具栏:搜索 + 类型按钮组 + 时长筛选 ─────────────── */}
|
||||
<div className="xx-templates-toolbar">
|
||||
<div className="xx-templates-search">
|
||||
<span className="xx-templates-search-icon">
|
||||
<SearchOutlined />
|
||||
</span>
|
||||
<input
|
||||
className="xx-templates-search-input"
|
||||
type="text"
|
||||
placeholder="搜索模板名称、描述或标签..."
|
||||
value={searchText}
|
||||
onChange={handleSearchChange}
|
||||
/>
|
||||
</div>
|
||||
<div className="xx-templates-categories">
|
||||
{TEMPLATE_TYPES.map((cat) => (
|
||||
<button
|
||||
key={cat.type}
|
||||
className={`xx-templates-cat-btn${activeType === cat.type ? " active" : ""}`}
|
||||
onClick={() => handleCategoryChange(cat.type)}
|
||||
>
|
||||
<span className="xx-templates-cat-icon">{cat.icon}</span>
|
||||
{cat.label}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
{/* 时长筛选 */}
|
||||
<div className="xx-templates-duration-filter">
|
||||
{DURATION_OPTIONS.map((opt) => (
|
||||
<button
|
||||
key={opt.value}
|
||||
className={`xx-templates-duration-btn${durationRange === opt.value ? " active" : ""}`}
|
||||
onClick={() => handleDurationChange(opt.value)}
|
||||
>
|
||||
{opt.label}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
{/* 工具栏:搜索 + 类型按钮组 + 时长筛选 */}
|
||||
<TemplateToolbar
|
||||
searchText={searchText}
|
||||
onSearchChange={handleSearchChange}
|
||||
activeType={activeType}
|
||||
onTypeChange={handleCategoryChange}
|
||||
durationRange={durationRange}
|
||||
onDurationChange={handleDurationChange}
|
||||
/>
|
||||
|
||||
{/* ── 模板展示区 ────────────────────────────────────────── */}
|
||||
{templates.length === 0 ? (
|
||||
<div className="xx-templates-empty">
|
||||
<div className="xx-templates-empty-icon">
|
||||
<InboxOutlined />
|
||||
</div>
|
||||
<h3>
|
||||
{searchText || activeType !== "全部" || durationRange ? "未找到匹配的模板" : "暂无模板"}
|
||||
</h3>
|
||||
<p>
|
||||
{searchText || activeType !== "全部" || durationRange
|
||||
? "试试调整搜索条件或切换类型"
|
||||
: "点击上方「创建模板」开始创作"}
|
||||
</p>
|
||||
</div>
|
||||
) : (
|
||||
<>
|
||||
<div className="xx-templates-grid">
|
||||
{templates.map((tpl) => (
|
||||
<TemplateCard
|
||||
key={tpl.id}
|
||||
template={tpl}
|
||||
isFavorite={tpl.is_favorite ?? false}
|
||||
onPreview={handlePreview}
|
||||
onToggleFavorite={toggleFavorite}
|
||||
onUse={handleUse}
|
||||
/>
|
||||
))}
|
||||
</div>
|
||||
{/* 模板展示区 */}
|
||||
<TemplateGrid
|
||||
templates={templates}
|
||||
total={totalTemplates}
|
||||
page={page}
|
||||
pageSize={pageSize}
|
||||
isLoading={isLoading}
|
||||
isError={isError}
|
||||
errorMessage={error?.message}
|
||||
searchText={searchText}
|
||||
activeType={activeType}
|
||||
durationRange={durationRange}
|
||||
onPageChange={setPage}
|
||||
onPreview={handlePreview}
|
||||
onToggleFavorite={toggleFavorite}
|
||||
onUse={handleUse}
|
||||
/>
|
||||
|
||||
{/* 分页 */}
|
||||
{totalTemplates > pageSize && (
|
||||
<div className="xx-templates-pagination">
|
||||
<Pagination
|
||||
current={page}
|
||||
pageSize={pageSize}
|
||||
total={totalTemplates}
|
||||
showSizeChanger={false}
|
||||
showQuickJumper
|
||||
showTotal={(total) => `共 ${total} 个模板`}
|
||||
onChange={(p) => setPage(p)}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
|
||||
{/* ── 详情弹窗 ──────────────────────────────────────────── */}
|
||||
{/* 详情弹窗 */}
|
||||
{previewTemplate && (
|
||||
<TemplateDetailModal
|
||||
template={previewTemplate}
|
||||
isFavorite={previewTemplate.is_favorite ?? false}
|
||||
onClose={() => setPreviewTemplate(null)}
|
||||
onToggleFavorite={toggleFavorite}
|
||||
onUse={handleUse}
|
||||
onCopy={handleCopy}
|
||||
onClose={handleClose}
|
||||
onToggleFavorite={handleToggleFavorite}
|
||||
onUse={handleUseFromDetail}
|
||||
onCopy={handleCopyFromDetail}
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* 详情加载中的提示(可选覆盖层) */}
|
||||
{/* 详情加载中的提示 */}
|
||||
{detailLoading && previewTemplate && (
|
||||
<div className="xx-template-detail-loading">
|
||||
<LoadingOutlined /> 加载中...
|
||||
</div>
|
||||
<div className="xx-template-detail-loading">加载中...</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
import React from "react"
|
||||
import { Tag } from "antd"
|
||||
import type { TemplateItem } from "@/api/templates"
|
||||
import { gradientForCategory, getTypeColor, formatDuration } from "../../utils/templateLibrary"
|
||||
|
||||
interface TemplateCardProps {
|
||||
template: TemplateItem
|
||||
isFavorite: boolean
|
||||
onPreview: (template: TemplateItem) => void
|
||||
onToggleFavorite: (id: string, e: React.MouseEvent) => void
|
||||
onUse: (template: TemplateItem) => void
|
||||
}
|
||||
|
||||
export const TemplateCard: React.FC<TemplateCardProps> = ({
|
||||
template,
|
||||
isFavorite,
|
||||
onPreview,
|
||||
onToggleFavorite,
|
||||
onUse,
|
||||
}) => {
|
||||
return (
|
||||
<div className="xx-template-card" onClick={() => onPreview(template)}>
|
||||
{/* 缩略图 */}
|
||||
<div className="xx-template-thumb">
|
||||
{template.thumbnail_url ? (
|
||||
<img src={template.thumbnail_url} alt={template.name} className="xx-template-thumb-img" />
|
||||
) : (
|
||||
<div
|
||||
className="xx-template-thumb-bg"
|
||||
style={{ background: gradientForCategory(template.category) }}
|
||||
>
|
||||
{(template.description ?? "").slice(0, 80)}
|
||||
{(template.description ?? "").length > 80 ? "..." : ""}
|
||||
</div>
|
||||
)}
|
||||
<div className="xx-template-thumb-overlay" />
|
||||
<div className="xx-template-thumb-name">{template.name}</div>
|
||||
<div className="xx-template-thumb-meta">
|
||||
<span className="xx-template-thumb-duration">
|
||||
{formatDuration(template.estimated_duration ?? template.target_duration)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-preview-hint">点击查看详情</div>
|
||||
<button
|
||||
className={`xx-template-fav-btn${isFavorite ? " is-favorite" : ""}`}
|
||||
onClick={(e) => onToggleFavorite(template.id, e)}
|
||||
title={isFavorite ? "取消收藏" : "收藏"}
|
||||
>
|
||||
{isFavorite ? "★" : "☆"}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* 信息区 */}
|
||||
<div className="xx-template-info">
|
||||
<div className="xx-template-info-top">
|
||||
<span
|
||||
className="xx-template-category-pill"
|
||||
style={{
|
||||
color: getTypeColor(template.category),
|
||||
background: `${getTypeColor(template.category)}18`,
|
||||
}}
|
||||
>
|
||||
{template.category}
|
||||
</span>
|
||||
{(template.tags ?? []).slice(0, 2).map((tag) => (
|
||||
<Tag key={tag} className="xx-template-tag-pill" bordered={false}>
|
||||
{tag}
|
||||
</Tag>
|
||||
))}
|
||||
</div>
|
||||
<p className="xx-template-desc">{template.description ?? ""}</p>
|
||||
<div className="xx-template-meta">
|
||||
<span className="xx-template-usage">已使用 {template.usage_count ?? 0} 次</span>
|
||||
<button
|
||||
className="xx-template-use-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onUse(template)
|
||||
}}
|
||||
>
|
||||
使用此模板
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,219 @@
|
||||
import React from "react"
|
||||
import { Button, Descriptions, Tooltip } from "antd"
|
||||
import { CopyOutlined, ThunderboltOutlined } from "@ant-design/icons"
|
||||
import type { TemplateItem, TemplateSegment } from "@/api/templates"
|
||||
import {
|
||||
gradientForCategory,
|
||||
getTypeColor,
|
||||
formatDuration,
|
||||
formatConfig,
|
||||
getMaterialTypeLabel,
|
||||
calcTotalSegmentDuration,
|
||||
} from "../../utils/templateLibrary"
|
||||
import { TEMPLATE_TYPES } from "../../constants/templateLibrary"
|
||||
|
||||
interface TemplateDetailModalProps {
|
||||
template: TemplateItem
|
||||
isFavorite: boolean
|
||||
onClose: () => void
|
||||
onToggleFavorite: (id: string) => void
|
||||
onUse: (template: TemplateItem) => void
|
||||
onCopy: (template: TemplateItem) => void
|
||||
}
|
||||
|
||||
export const TemplateDetailModal: React.FC<TemplateDetailModalProps> = ({
|
||||
template,
|
||||
isFavorite,
|
||||
onClose,
|
||||
onToggleFavorite,
|
||||
onUse,
|
||||
onCopy,
|
||||
}) => {
|
||||
const segments = template.segments ?? []
|
||||
const totalSegmentDuration = calcTotalSegmentDuration(segments)
|
||||
|
||||
return (
|
||||
<div className="xx-template-modal-overlay" onClick={onClose}>
|
||||
<div
|
||||
className="xx-template-modal xx-template-modal-wide"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
>
|
||||
{/* 关闭按钮 */}
|
||||
<button className="xx-template-modal-close" onClick={onClose} title="关闭">
|
||||
✕
|
||||
</button>
|
||||
|
||||
{/* 预览区域 */}
|
||||
<div
|
||||
className="xx-template-modal-preview"
|
||||
style={{ background: gradientForCategory(template.category) }}
|
||||
>
|
||||
{template.thumbnail_url ? (
|
||||
<img
|
||||
src={template.thumbnail_url}
|
||||
alt={template.name}
|
||||
className="xx-template-modal-thumb-img"
|
||||
/>
|
||||
) : (
|
||||
<div className="xx-template-modal-preview-content">
|
||||
<span className="xx-template-preview-icon">
|
||||
{TEMPLATE_TYPES.find((t) => t.type === template.category)?.icon ?? "📋"}
|
||||
</span>
|
||||
<span className="xx-template-preview-title">{template.name}</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 内容区域 */}
|
||||
<div className="xx-template-modal-content">
|
||||
{/* 标题行 */}
|
||||
<div className="xx-template-modal-title-row">
|
||||
<h3>{template.name}</h3>
|
||||
<span
|
||||
className="xx-template-modal-type-badge"
|
||||
style={{
|
||||
color: getTypeColor(template.category),
|
||||
background: `${getTypeColor(template.category)}18`,
|
||||
}}
|
||||
>
|
||||
{TEMPLATE_TYPES.find((t) => t.type === template.category)?.icon} {template.category}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{/* 描述 */}
|
||||
<p className="xx-template-modal-desc">{template.description}</p>
|
||||
|
||||
{/* 标签 */}
|
||||
{(template.tags?.length ?? 0) > 0 && (
|
||||
<div className="xx-template-modal-tags">
|
||||
{template.tags!.map((tag) => (
|
||||
<span key={tag} className="xx-template-modal-tag">
|
||||
#{tag}
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 基本信息 */}
|
||||
<Descriptions
|
||||
column={2}
|
||||
size="small"
|
||||
className="xx-template-modal-desc-table"
|
||||
items={[
|
||||
{
|
||||
key: "duration",
|
||||
label: "目标时长",
|
||||
children: formatDuration(template.estimated_duration ?? template.target_duration),
|
||||
},
|
||||
{
|
||||
key: "clips",
|
||||
label: "片段数量",
|
||||
children: `${template.clip_count} 个`,
|
||||
},
|
||||
{
|
||||
key: "ratio",
|
||||
label: "视频比例",
|
||||
children: template.aspect_ratio ?? "16:9",
|
||||
},
|
||||
{
|
||||
key: "usage",
|
||||
label: "使用次数",
|
||||
children: `${template.usage_count ?? 0} 次`,
|
||||
},
|
||||
]}
|
||||
/>
|
||||
|
||||
{/* 素材规则(片段配置) */}
|
||||
{segments.length > 0 && (
|
||||
<div className="xx-template-modal-section">
|
||||
<h4>🎬 素材规则</h4>
|
||||
<div className="xx-template-modal-clip-list">
|
||||
{segments
|
||||
.sort((a, b) => a.segment_order - b.segment_order)
|
||||
.map((seg: TemplateSegment, idx: number) => (
|
||||
<div key={seg.id ?? idx} className="xx-template-modal-clip-item">
|
||||
<span className="xx-template-modal-clip-order">#{seg.segment_order}</span>
|
||||
<span
|
||||
className="xx-template-modal-clip-badge"
|
||||
style={{
|
||||
color: seg.material_type ? getTypeColor(seg.material_type) : "#64748b",
|
||||
background: seg.material_type
|
||||
? `${getTypeColor(seg.material_type)}18`
|
||||
: "#f1f5f9",
|
||||
}}
|
||||
>
|
||||
{getMaterialTypeLabel(seg.material_type)}
|
||||
</span>
|
||||
<span className="xx-template-modal-clip-desc">
|
||||
{seg.description || `片段 ${seg.segment_order}`}
|
||||
</span>
|
||||
<Tooltip title={`时长范围: ${seg.duration_min}秒 - ${seg.duration_max}秒`}>
|
||||
<span className="xx-template-modal-clip-duration">
|
||||
{seg.duration_min}-{seg.duration_max}秒
|
||||
</span>
|
||||
</Tooltip>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
<div className="xx-template-modal-total-duration">
|
||||
预估总时长:{formatDuration(Math.round(totalSegmentDuration))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 样式配置 */}
|
||||
<div className="xx-template-modal-section">
|
||||
<h4>🎨 样式配置</h4>
|
||||
<div className="xx-template-modal-style-grid">
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">字幕样式</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{formatConfig(template.subtitle_config)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">标题样式</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{formatConfig(template.title_config)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">BGM 配置</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{formatConfig(template.bgm_config)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-template-modal-style-item">
|
||||
<span className="xx-template-modal-style-label">视频比例</span>
|
||||
<span className="xx-template-modal-style-value">
|
||||
{template.aspect_ratio ?? "16:9"}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 统计信息 */}
|
||||
<div className="xx-template-modal-stats">
|
||||
<span>已使用 {template.usage_count ?? 0} 次</span>
|
||||
<button
|
||||
className={`xx-template-modal-fav-btn${isFavorite ? " is-favorite" : ""}`}
|
||||
onClick={() => onToggleFavorite(template.id)}
|
||||
>
|
||||
{isFavorite ? "★ 已收藏" : "☆ 收藏"}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* 操作按钮 */}
|
||||
<div className="xx-template-modal-actions">
|
||||
<Button icon={<CopyOutlined />} onClick={() => onCopy(template)}>
|
||||
复制模板
|
||||
</Button>
|
||||
<Button type="primary" icon={<ThunderboltOutlined />} onClick={() => onUse(template)}>
|
||||
使用此模板生成
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
import React from "react"
|
||||
import { Pagination } from "antd"
|
||||
import { InboxOutlined, LoadingOutlined, ExclamationCircleOutlined } from "@ant-design/icons"
|
||||
import { TemplateCard } from "./TemplateCard"
|
||||
import type { TemplateItem } from "@/api/templates"
|
||||
|
||||
interface TemplateGridProps {
|
||||
templates: TemplateItem[]
|
||||
total: number
|
||||
page: number
|
||||
pageSize: number
|
||||
isLoading: boolean
|
||||
isError: boolean
|
||||
errorMessage?: string
|
||||
searchText: string
|
||||
activeType: string
|
||||
durationRange: string
|
||||
onPageChange: (page: number) => void
|
||||
onPreview: (template: TemplateItem) => void
|
||||
onToggleFavorite: (id: string, e: React.MouseEvent) => void
|
||||
onUse: (template: TemplateItem) => void
|
||||
}
|
||||
|
||||
export const TemplateGrid: React.FC<TemplateGridProps> = ({
|
||||
templates,
|
||||
total,
|
||||
page,
|
||||
pageSize,
|
||||
isLoading,
|
||||
isError,
|
||||
errorMessage,
|
||||
searchText,
|
||||
activeType,
|
||||
durationRange,
|
||||
onPageChange,
|
||||
onPreview,
|
||||
onToggleFavorite,
|
||||
onUse,
|
||||
}) => {
|
||||
// Loading 状态
|
||||
if (isLoading) {
|
||||
return (
|
||||
<div className="xx-templates-empty">
|
||||
<div className="xx-templates-empty-icon">
|
||||
<LoadingOutlined />
|
||||
</div>
|
||||
<h3>加载模板中...</h3>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
// Error 状态
|
||||
if (isError) {
|
||||
return (
|
||||
<div className="xx-templates-empty">
|
||||
<div className="xx-templates-empty-icon">
|
||||
<ExclamationCircleOutlined />
|
||||
</div>
|
||||
<h3>加载失败</h3>
|
||||
<p>{errorMessage || "网络异常,请稍后重试"}</p>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
// 空状态
|
||||
if (templates.length === 0) {
|
||||
const hasFilter = !!searchText || activeType !== "全部" || !!durationRange
|
||||
return (
|
||||
<div className="xx-templates-empty">
|
||||
<div className="xx-templates-empty-icon">
|
||||
<InboxOutlined />
|
||||
</div>
|
||||
<h3>{hasFilter ? "未找到匹配的模板" : "暂无模板"}</h3>
|
||||
<p>{hasFilter ? "试试调整搜索条件或切换类型" : "点击上方「创建模板」开始创作"}</p>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
return (
|
||||
<>
|
||||
<div className="xx-templates-grid">
|
||||
{templates.map((tpl) => (
|
||||
<TemplateCard
|
||||
key={tpl.id}
|
||||
template={tpl}
|
||||
isFavorite={tpl.is_favorite ?? false}
|
||||
onPreview={onPreview}
|
||||
onToggleFavorite={onToggleFavorite}
|
||||
onUse={onUse}
|
||||
/>
|
||||
))}
|
||||
</div>
|
||||
|
||||
{/* 分页 */}
|
||||
{total > pageSize && (
|
||||
<div className="xx-templates-pagination">
|
||||
<Pagination
|
||||
current={page}
|
||||
pageSize={pageSize}
|
||||
total={total}
|
||||
showSizeChanger={false}
|
||||
showQuickJumper
|
||||
showTotal={(t) => `共 ${t} 个模板`}
|
||||
onChange={onPageChange}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
import React from "react"
|
||||
import { Button } from "antd"
|
||||
|
||||
interface TemplateHeaderProps {
|
||||
onCreateClick: () => void
|
||||
}
|
||||
|
||||
export const TemplateHeader: React.FC<TemplateHeaderProps> = ({ onCreateClick }) => {
|
||||
return (
|
||||
<div className="xx-templates-header">
|
||||
<div className="xx-templates-header-text">
|
||||
<h2>模板库</h2>
|
||||
<p>选择模板快速创建,支持自定义修改</p>
|
||||
</div>
|
||||
<Button type="primary" onClick={onCreateClick}>
|
||||
+ 创建模板
|
||||
</Button>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
import React from "react"
|
||||
import { SearchOutlined } from "@ant-design/icons"
|
||||
import type { EditTemplateType, DurationRange } from "../../types/templateLibrary"
|
||||
import { TEMPLATE_TYPES, DURATION_OPTIONS } from "../../constants/templateLibrary"
|
||||
|
||||
interface TemplateToolbarProps {
|
||||
searchText: string
|
||||
onSearchChange: (e: React.ChangeEvent<HTMLInputElement>) => void
|
||||
activeType: EditTemplateType | "全部"
|
||||
onTypeChange: (type: EditTemplateType | "全部") => void
|
||||
durationRange: DurationRange
|
||||
onDurationChange: (value: DurationRange) => void
|
||||
}
|
||||
|
||||
export const TemplateToolbar: React.FC<TemplateToolbarProps> = ({
|
||||
searchText,
|
||||
onSearchChange,
|
||||
activeType,
|
||||
onTypeChange,
|
||||
durationRange,
|
||||
onDurationChange,
|
||||
}) => {
|
||||
return (
|
||||
<div className="xx-templates-toolbar">
|
||||
<div className="xx-templates-search">
|
||||
<span className="xx-templates-search-icon">
|
||||
<SearchOutlined />
|
||||
</span>
|
||||
<input
|
||||
className="xx-templates-search-input"
|
||||
type="text"
|
||||
placeholder="搜索模板名称、描述或标签..."
|
||||
value={searchText}
|
||||
onChange={onSearchChange}
|
||||
/>
|
||||
</div>
|
||||
<div className="xx-templates-categories">
|
||||
{TEMPLATE_TYPES.map((cat) => (
|
||||
<button
|
||||
key={cat.type}
|
||||
className={`xx-templates-cat-btn${activeType === cat.type ? " active" : ""}`}
|
||||
onClick={() => onTypeChange(cat.type)}
|
||||
>
|
||||
<span className="xx-templates-cat-icon">{cat.icon}</span>
|
||||
{cat.label}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
{/* 时长筛选 */}
|
||||
<div className="xx-templates-duration-filter">
|
||||
{DURATION_OPTIONS.map((opt) => (
|
||||
<button
|
||||
key={opt.value}
|
||||
className={`xx-templates-duration-btn${durationRange === opt.value ? " active" : ""}`}
|
||||
onClick={() => onDurationChange(opt.value)}
|
||||
>
|
||||
{opt.label}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
import type { EditTemplateType, DurationRange } from "../types/templateLibrary"
|
||||
|
||||
export const TEMPLATE_TYPES: Array<{
|
||||
type: EditTemplateType | "全部"
|
||||
label: string
|
||||
icon: string
|
||||
color: string
|
||||
}> = [
|
||||
{ type: "全部", label: "全部", icon: "📋", color: "#6366f1" },
|
||||
{ type: "口播", label: "口播", icon: "🎙️", color: "#6366f1" },
|
||||
{ type: "种草", label: "种草", icon: "🌱", color: "#10b981" },
|
||||
{ type: "产品", label: "产品", icon: "📦", color: "#0ea5e9" },
|
||||
{ type: "品牌", label: "品牌", icon: "🏷️", color: "#f59e0b" },
|
||||
{ type: "混剪", label: "混剪", icon: "🎬", color: "#8b5cf6" },
|
||||
{ type: "Vlog", label: "Vlog", icon: "📹", color: "#ec4899" },
|
||||
]
|
||||
|
||||
export const DURATION_OPTIONS: Array<{
|
||||
value: DurationRange
|
||||
label: string
|
||||
}> = [
|
||||
{ value: "", label: "全部时长" },
|
||||
{ value: "short", label: "30秒以内" },
|
||||
{ value: "medium", label: "30秒-2分钟" },
|
||||
{ value: "long", label: "2分钟以上" },
|
||||
]
|
||||
|
||||
export const MATERIAL_TYPE_LABELS: Record<string, string> = {
|
||||
video: "视频",
|
||||
image: "图片",
|
||||
audio: "音频",
|
||||
voiceover: "配音",
|
||||
subtitle: "字幕",
|
||||
null: "不限",
|
||||
}
|
||||
|
||||
export const DEFAULT_PAGE_SIZE = 12
|
||||
export const CATEGORY_GRADIENT_MAP: Record<string, string> = {
|
||||
口播: "linear-gradient(135deg, #6366f1, #8b5cf6)",
|
||||
种草: "linear-gradient(135deg, #10b981, #059669)",
|
||||
产品: "linear-gradient(135deg, #0ea5e9, #0284c7)",
|
||||
品牌: "linear-gradient(135deg, #f59e0b, #d97706)",
|
||||
混剪: "linear-gradient(135deg, #8b5cf6, #6d28d9)",
|
||||
Vlog: "linear-gradient(135deg, #ec4899, #db2777)",
|
||||
}
|
||||
export const DEFAULT_GRADIENT = "linear-gradient(135deg, #6366f1, #8b5cf6)"
|
||||
@@ -0,0 +1,62 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { getTemplate, type TemplateItem } from "@/api/templates"
|
||||
|
||||
interface UseTemplateDetailProps {
|
||||
onToggleFavorite: (id: string) => void
|
||||
onUse: (template: TemplateItem) => void
|
||||
onCopy: (template: TemplateItem) => void
|
||||
}
|
||||
|
||||
export const useTemplateDetail = ({ onToggleFavorite, onUse, onCopy }: UseTemplateDetailProps) => {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
/* 弹窗状态 */
|
||||
const [previewTemplate, setPreviewTemplate] = useState<TemplateItem | null>(null)
|
||||
const [detailLoading, setDetailLoading] = useState(false)
|
||||
|
||||
/* 点击卡片 → 获取详情并展示弹窗 */
|
||||
const handlePreview = useCallback(async (template: TemplateItem) => {
|
||||
setDetailLoading(true)
|
||||
setPreviewTemplate(template)
|
||||
try {
|
||||
const detail = await getTemplate(template.id)
|
||||
setPreviewTemplate(detail)
|
||||
} catch {
|
||||
message.warning("模板详情加载失败,显示摘要信息")
|
||||
} finally {
|
||||
setDetailLoading(false)
|
||||
}
|
||||
}, [])
|
||||
|
||||
/* 关闭弹窗 */
|
||||
const handleClose = useCallback(() => {
|
||||
setPreviewTemplate(null)
|
||||
setDetailLoading(false)
|
||||
}, [])
|
||||
|
||||
/* 收藏切换(同时更新预览模板的状态) */
|
||||
const handleToggleFavorite = useCallback(
|
||||
(id: string) => {
|
||||
onToggleFavorite(id)
|
||||
/* 乐观更新详情弹窗的收藏状态 */
|
||||
setPreviewTemplate((prev) =>
|
||||
prev && prev.id === id ? { ...prev, is_favorite: !prev.is_favorite } : prev,
|
||||
)
|
||||
/* 刷新列表缓存 */
|
||||
queryClient.invalidateQueries({ queryKey: ["templates"] })
|
||||
},
|
||||
[onToggleFavorite, queryClient],
|
||||
)
|
||||
|
||||
return {
|
||||
previewTemplate,
|
||||
detailLoading,
|
||||
handlePreview,
|
||||
handleClose,
|
||||
handleToggleFavorite,
|
||||
handleUse: onUse,
|
||||
handleCopy: onCopy,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
import { useState, useMemo, useCallback } from "react"
|
||||
import { useNavigate } from "react-router-dom"
|
||||
import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import {
|
||||
getTemplates,
|
||||
toggleFavoriteTemplate,
|
||||
copyTemplate,
|
||||
type TemplateItem,
|
||||
type TemplateListParams,
|
||||
} from "@/api/templates"
|
||||
import type { EditTemplateType, DurationRange } from "../types/templateLibrary"
|
||||
import { DEFAULT_PAGE_SIZE } from "../constants/templateLibrary"
|
||||
|
||||
export const useTemplateLibrary = () => {
|
||||
const navigate = useNavigate()
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
/* 筛选状态 */
|
||||
const [searchText, setSearchText] = useState("")
|
||||
const [activeType, setActiveType] = useState<EditTemplateType | "全部">("全部")
|
||||
const [durationRange, setDurationRange] = useState<DurationRange>("")
|
||||
const [page, setPage] = useState(1)
|
||||
const [pageSize] = useState(DEFAULT_PAGE_SIZE)
|
||||
|
||||
/* 构建查询参数 */
|
||||
const queryParams: TemplateListParams = useMemo(() => {
|
||||
const params: TemplateListParams = {
|
||||
page,
|
||||
page_size: pageSize,
|
||||
}
|
||||
if (activeType !== "全部") params.category = activeType
|
||||
if (searchText.trim()) params.keyword = searchText.trim()
|
||||
if (durationRange) params.duration_range = durationRange
|
||||
return params
|
||||
}, [page, pageSize, activeType, searchText, durationRange])
|
||||
|
||||
/* 获取模板列表 */
|
||||
const {
|
||||
data: templateData,
|
||||
isLoading,
|
||||
isError,
|
||||
error,
|
||||
} = useQuery({
|
||||
queryKey: ["templates", queryParams],
|
||||
queryFn: () => getTemplates(queryParams),
|
||||
staleTime: 30_000,
|
||||
})
|
||||
|
||||
const templates = templateData?.items ?? []
|
||||
const totalTemplates = templateData?.total ?? 0
|
||||
|
||||
/* 收藏 mutation */
|
||||
const favMutation = useMutation({
|
||||
mutationFn: toggleFavoriteTemplate,
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["templates"] })
|
||||
},
|
||||
})
|
||||
|
||||
/* 复制模板 mutation */
|
||||
const copyMutation = useMutation({
|
||||
mutationFn: copyTemplate,
|
||||
onSuccess: (data) => {
|
||||
message.success(`模板「${data.name}」已复制到「我的模板」`)
|
||||
queryClient.invalidateQueries({ queryKey: ["templates"] })
|
||||
},
|
||||
onError: () => {
|
||||
message.error("复制模板失败,请稍后重试")
|
||||
},
|
||||
})
|
||||
|
||||
/* 操作:切换收藏 */
|
||||
const toggleFavorite = useCallback(
|
||||
(id: string, e?: React.MouseEvent) => {
|
||||
e?.stopPropagation()
|
||||
favMutation.mutate(id)
|
||||
},
|
||||
[favMutation],
|
||||
)
|
||||
|
||||
/* 操作:复制模板 */
|
||||
const handleCopy = useCallback(
|
||||
(template: TemplateItem) => {
|
||||
copyMutation.mutate(template.id)
|
||||
},
|
||||
[copyMutation],
|
||||
)
|
||||
|
||||
/* 操作:使用模板 → 跳转剪辑编辑器 */
|
||||
const handleUse = useCallback(
|
||||
(template: TemplateItem) => {
|
||||
navigate(`/app/editing-planner?templateId=${template.id}`)
|
||||
},
|
||||
[navigate],
|
||||
)
|
||||
|
||||
/* 操作:创建模板 */
|
||||
const handleCreate = useCallback(() => {
|
||||
navigate("/app/editing-planner")
|
||||
}, [navigate])
|
||||
|
||||
/* 搜索 */
|
||||
const handleSearchChange = useCallback((e: React.ChangeEvent<HTMLInputElement>) => {
|
||||
setSearchText(e.target.value)
|
||||
setPage(1)
|
||||
}, [])
|
||||
|
||||
/* 切换分类 */
|
||||
const handleCategoryChange = useCallback((type: EditTemplateType | "全部") => {
|
||||
setActiveType(type)
|
||||
setPage(1)
|
||||
}, [])
|
||||
|
||||
/* 切换时长筛选 */
|
||||
const handleDurationChange = useCallback((value: DurationRange) => {
|
||||
setDurationRange(value)
|
||||
setPage(1)
|
||||
}, [])
|
||||
|
||||
return {
|
||||
/* 状态 */
|
||||
templates,
|
||||
totalTemplates,
|
||||
isLoading,
|
||||
isError,
|
||||
error,
|
||||
searchText,
|
||||
activeType,
|
||||
durationRange,
|
||||
page,
|
||||
pageSize,
|
||||
/* mutations */
|
||||
favMutation,
|
||||
copyMutation,
|
||||
/* setters */
|
||||
setPage,
|
||||
/* handlers */
|
||||
toggleFavorite,
|
||||
handleCopy,
|
||||
handleUse,
|
||||
handleCreate,
|
||||
handleSearchChange,
|
||||
handleCategoryChange,
|
||||
handleDurationChange,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
/** 模板类型 */
|
||||
export type EditTemplateType = "口播" | "种草" | "产品" | "品牌" | "混剪" | "Vlog"
|
||||
|
||||
/** 时长筛选值 */
|
||||
export type DurationRange = "" | "short" | "medium" | "long"
|
||||
|
||||
/** 配置展示字段 */
|
||||
export interface ConfigDisplayFields {
|
||||
font_size?: string | number
|
||||
font_family?: string
|
||||
color?: string
|
||||
position?: string
|
||||
volume?: string | number
|
||||
name?: string
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
import type { ConfigDisplayFields } from "../types/templateLibrary"
|
||||
import {
|
||||
TEMPLATE_TYPES,
|
||||
CATEGORY_GRADIENT_MAP,
|
||||
DEFAULT_GRADIENT,
|
||||
MATERIAL_TYPE_LABELS,
|
||||
} from "../constants/templateLibrary"
|
||||
import type { TemplateSegment } from "@/api/templates"
|
||||
|
||||
/** 获取类型对应颜色 */
|
||||
export const getTypeColor = (type: string): string => {
|
||||
const found = TEMPLATE_TYPES.find((t) => t.type === type)
|
||||
return found?.color ?? "#6366f1"
|
||||
}
|
||||
|
||||
/** 根据 category 生成占位渐变色 */
|
||||
export const gradientForCategory = (category: string): string => {
|
||||
return CATEGORY_GRADIENT_MAP[category] ?? DEFAULT_GRADIENT
|
||||
}
|
||||
|
||||
/** 格式化时长 */
|
||||
export const formatDuration = (seconds: number | undefined | null): string => {
|
||||
if (!seconds || seconds <= 0) return "0秒"
|
||||
const totalSec = Math.round(seconds)
|
||||
const m = Math.floor(totalSec / 60)
|
||||
const s = totalSec % 60
|
||||
if (m === 0) return `${s}秒`
|
||||
return `${m}分${s > 0 ? `${s}秒` : ""}`
|
||||
}
|
||||
|
||||
/** 格式化配置对象为可读文本 */
|
||||
export const formatConfig = (config?: object): string => {
|
||||
if (!config || Object.keys(config).length === 0) return "默认"
|
||||
const c = config as ConfigDisplayFields
|
||||
const parts: string[] = []
|
||||
if (c.font_size) parts.push(`字号: ${c.font_size}`)
|
||||
if (c.font_family) parts.push(`字体: ${c.font_family}`)
|
||||
if (c.color) parts.push(`颜色: ${c.color}`)
|
||||
if (c.position) parts.push(`位置: ${c.position}`)
|
||||
if (c.volume !== undefined) parts.push(`音量: ${c.volume}%`)
|
||||
if (c.name) parts.push(String(c.name))
|
||||
return parts.length > 0 ? parts.join(" / ") : JSON.stringify(config)
|
||||
}
|
||||
|
||||
/** 获取素材类型标签文本 */
|
||||
export const getMaterialTypeLabel = (materialType: string | null | undefined): string => {
|
||||
if (!materialType) return "不限"
|
||||
return MATERIAL_TYPE_LABELS[materialType] ?? materialType
|
||||
}
|
||||
|
||||
/** 计算片段总时长(取每个片段 min/max 的平均值) */
|
||||
export const calcTotalSegmentDuration = (segments: TemplateSegment[]): number => {
|
||||
return segments.reduce((sum, s) => sum + (s.duration_min + s.duration_max) / 2, 0)
|
||||
}
|
||||
@@ -115,6 +115,16 @@ vi.mock("@/api/templates", () => ({
|
||||
vi.mock("@/pages/templates/TemplateLibrary.css", () => ({}))
|
||||
|
||||
import TemplateLibrary from "@/pages/templates/TemplateLibrary"
|
||||
import "@/pages/templates/types/templateLibrary"
|
||||
import "@/pages/templates/constants/templateLibrary"
|
||||
import "@/pages/templates/utils/templateLibrary"
|
||||
import "@/pages/templates/hooks/useTemplateLibrary"
|
||||
import "@/pages/templates/hooks/useTemplateDetail"
|
||||
import "@/pages/templates/components/template-library/TemplateCard"
|
||||
import "@/pages/templates/components/template-library/TemplateDetailModal"
|
||||
import "@/pages/templates/components/template-library/TemplateHeader"
|
||||
import "@/pages/templates/components/template-library/TemplateToolbar"
|
||||
import "@/pages/templates/components/template-library/TemplateGrid"
|
||||
|
||||
describe("TemplateLibrary Page", () => {
|
||||
it("should render without crashing", () => {
|
||||
|
||||
@@ -18,135 +18,22 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, probe_video_info, run_ffmpeg
|
||||
from video_processing.path_security import PathSecurityError, is_in_allowed_dirs, safe_resolve_path
|
||||
|
||||
from packages.domain.video_concat import ( # noqa: F401 向后兼容导出
|
||||
ALLOWED_VIDEO_EXTENSIONS,
|
||||
CONCAT_DEMUXER_REQUIRED_PARAMS,
|
||||
MAX_CONCAT_SEGMENTS,
|
||||
ConcatConfig,
|
||||
ConcatSegment,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
MAX_CONCAT_SEGMENTS = 50 # 最大拼接段数(安全上限,防止OOM)
|
||||
|
||||
ALLOWED_VIDEO_EXTENSIONS = {".mp4", ".mov", ".avi", ".mkv", ".webm", ".flv", ".wmv"}
|
||||
|
||||
# concat demuxer 要求一致的参数列表
|
||||
CONCAT_DEMUXER_REQUIRED_PARAMS = [
|
||||
"codec_name", # 视频编码
|
||||
"width", # 宽度
|
||||
"height", # 高度
|
||||
"r_frame_rate", # 帧率
|
||||
"pix_fmt", # 像素格式
|
||||
"sample_rate", # 音频采样率
|
||||
"channels", # 音频声道数
|
||||
"audio_codec", # 音频编码
|
||||
]
|
||||
|
||||
|
||||
# ── 拼接片段配置 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConcatSegment:
|
||||
"""单个拼接片段."""
|
||||
|
||||
video_path: str # 视频文件路径
|
||||
start_time: float = 0.0 # 开始时间(秒),从视频的哪个位置开始取
|
||||
duration: float = 0.0 # 持续时长(秒),0表示取到末尾
|
||||
has_audio: bool = True # 是否包含音频
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, seg: dict) -> "ConcatSegment":
|
||||
"""从字典创建拼接片段,带安全类型转换."""
|
||||
try:
|
||||
start_time = max(0.0, float(seg.get("start_time", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
start_time = 0.0
|
||||
|
||||
try:
|
||||
duration = max(0.0, float(seg.get("duration", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
duration = 0.0
|
||||
|
||||
return cls(
|
||||
video_path=str(seg.get("video_path", "")),
|
||||
start_time=start_time,
|
||||
duration=duration,
|
||||
has_audio=bool(seg.get("has_audio", True)),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConcatConfig:
|
||||
"""视频拼接配置."""
|
||||
|
||||
segments: list[ConcatSegment] = field(default_factory=list)
|
||||
output_width: int = 0 # 输出宽度(0=自动取第一段)
|
||||
output_height: int = 0 # 输出高度(0=自动取第一段)
|
||||
output_fps: float = 0.0 # 输出帧率(0=自动取第一段)
|
||||
force_reencode: bool = False # 强制重新编码(不用 stream copy)
|
||||
transition: str = "none" # 转场效果(none/crossfade)- 预留
|
||||
transition_duration: float = 0.3 # 转场时长
|
||||
|
||||
@classmethod
|
||||
def from_config_dict(cls, config: dict | None) -> "ConcatConfig":
|
||||
"""从配置字典创建 ConcatConfig."""
|
||||
if not config or not isinstance(config, dict):
|
||||
return cls()
|
||||
|
||||
segments_raw = config.get("segments", [])
|
||||
segments: list[ConcatSegment] = []
|
||||
|
||||
if isinstance(segments_raw, list):
|
||||
for s in segments_raw:
|
||||
if isinstance(s, dict) and s.get("video_path"):
|
||||
try:
|
||||
seg = ConcatSegment.from_dict(s)
|
||||
if seg.video_path:
|
||||
segments.append(seg)
|
||||
except Exception:
|
||||
logger.warning("[concat] skip invalid segment: %s", s)
|
||||
continue
|
||||
|
||||
try:
|
||||
output_width = max(0, int(config.get("output_width", 0)))
|
||||
except (TypeError, ValueError):
|
||||
output_width = 0
|
||||
|
||||
try:
|
||||
output_height = max(0, int(config.get("output_height", 0)))
|
||||
except (TypeError, ValueError):
|
||||
output_height = 0
|
||||
|
||||
try:
|
||||
output_fps = max(0.0, float(config.get("output_fps", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
output_fps = 0.0
|
||||
|
||||
return cls(
|
||||
segments=segments,
|
||||
output_width=output_width,
|
||||
output_height=output_height,
|
||||
output_fps=output_fps,
|
||||
force_reencode=bool(config.get("force_reencode", False)),
|
||||
transition=str(config.get("transition", "none")),
|
||||
transition_duration=max(0.1, float(config.get("transition_duration", 0.3))),
|
||||
)
|
||||
|
||||
@property
|
||||
def has_effect(self) -> bool:
|
||||
"""是否有有效片段需要拼接."""
|
||||
return len([s for s in self.segments if s.video_path]) >= 2
|
||||
|
||||
@property
|
||||
def total_segments(self) -> int:
|
||||
"""有效片段数量."""
|
||||
return len([s for s in self.segments if s.video_path])
|
||||
|
||||
|
||||
# ── 路径安全校验 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
@@ -24,264 +24,34 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from video_processing.path_security import PathSecurityError, is_in_allowed_dirs, safe_resolve_path
|
||||
|
||||
from packages.domain.subtitle_style import (
|
||||
ALLOWED_SUBTITLE_EXTENSIONS,
|
||||
DEFAULT_COLOR,
|
||||
DEFAULT_FONT,
|
||||
DEFAULT_FONT_SIZE,
|
||||
DEFAULT_MAX_CHARS_PER_LINE,
|
||||
DEFAULT_POSITION,
|
||||
DEFAULT_STROKE_COLOR,
|
||||
DEFAULT_STROKE_WIDTH,
|
||||
POSITION_ALIASES,
|
||||
POSITION_ALIGNMENT,
|
||||
SubtitleSegment,
|
||||
SubtitleStyle,
|
||||
)
|
||||
from packages.domain.subtitle_style import escape_ass_text as _escape_ass_text # noqa: F401 向后兼容导出
|
||||
from packages.domain.subtitle_style import format_ass_time as _format_ass_time
|
||||
from packages.domain.subtitle_style import hex_to_ass_bgr as _hex_to_ass_bgr
|
||||
from packages.domain.subtitle_style import hex_to_ass_color as _hex_to_ass_color
|
||||
from packages.domain.subtitle_style import opacity_to_ass_alpha as _opacity_to_ass_alpha
|
||||
from packages.domain.subtitle_style import wrap_text as _wrap_text
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
ALLOWED_SUBTITLE_EXTENSIONS = {".srt", ".ass", ".vtt", ".sub"}
|
||||
|
||||
# 9宫格位置映射(ASS alignment 编号)
|
||||
POSITION_ALIGNMENT = {
|
||||
"top_left": 7,
|
||||
"top_center": 8,
|
||||
"top_right": 9,
|
||||
"middle_left": 4,
|
||||
"center": 5,
|
||||
"middle_right": 6,
|
||||
"bottom_left": 1,
|
||||
"bottom_center": 2,
|
||||
"bottom_right": 3,
|
||||
}
|
||||
|
||||
# 位置简称兼容
|
||||
POSITION_ALIASES = {
|
||||
"top": "top_center",
|
||||
"bottom": "bottom_center",
|
||||
"middle": "center",
|
||||
"left": "middle_left",
|
||||
"right": "middle_right",
|
||||
}
|
||||
|
||||
DEFAULT_FONT = "思源黑体"
|
||||
DEFAULT_FONT_SIZE = 24
|
||||
DEFAULT_COLOR = "#FFFFFF"
|
||||
DEFAULT_STROKE_COLOR = "#000000"
|
||||
DEFAULT_STROKE_WIDTH = 1.5
|
||||
DEFAULT_POSITION = "bottom_center"
|
||||
DEFAULT_MAX_CHARS_PER_LINE = 20
|
||||
|
||||
|
||||
# ── 字幕样式配置 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class SubtitleStyle:
|
||||
"""字幕样式配置."""
|
||||
|
||||
font_name: str = DEFAULT_FONT
|
||||
font_size: int = DEFAULT_FONT_SIZE
|
||||
font_color: str = DEFAULT_COLOR
|
||||
bold: bool = False
|
||||
italic: bool = False
|
||||
|
||||
# 描边
|
||||
stroke_enabled: bool = True
|
||||
stroke_color: str = DEFAULT_STROKE_COLOR
|
||||
stroke_width: float = DEFAULT_STROKE_WIDTH
|
||||
|
||||
# 阴影
|
||||
shadow_enabled: bool = False
|
||||
shadow_color: str = "#000000"
|
||||
shadow_offset_x: int = 2
|
||||
shadow_offset_y: int = 2
|
||||
shadow_blur: float = 0.0
|
||||
|
||||
# 背景框
|
||||
background_enabled: bool = False
|
||||
background_color: str = "#000000"
|
||||
background_opacity: float = 0.5 # 0.0 ~ 1.0
|
||||
background_padding: int = 8
|
||||
background_radius: int = 4
|
||||
|
||||
# 位置
|
||||
position: str = DEFAULT_POSITION # 9宫格位置名
|
||||
margin_v: int = 60 # 垂直边距
|
||||
margin_l: int = 40 # 左边距
|
||||
margin_r: int = 40 # 右边距
|
||||
|
||||
# 多行
|
||||
max_chars_per_line: int = DEFAULT_MAX_CHARS_PER_LINE
|
||||
line_spacing: int = 0 # 行间距
|
||||
|
||||
# 动画
|
||||
fade_in: float = 0.0 # 淡入时长(秒)
|
||||
fade_out: float = 0.0 # 淡出时长(秒)
|
||||
animation_type: str = "none" # none/fade/slide/typewriter
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config: dict[str, Any] | None) -> "SubtitleStyle":
|
||||
"""从字典创建样式配置,带安全类型转换."""
|
||||
if not config or not isinstance(config, dict):
|
||||
return cls()
|
||||
|
||||
def safe_str(key: str, default: str) -> str:
|
||||
val = config.get(key, default)
|
||||
return str(val) if val is not None else default
|
||||
|
||||
def safe_int(key: str, default: int) -> int:
|
||||
try:
|
||||
return int(config.get(key, default))
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
def safe_float(key: str, default: float) -> float:
|
||||
try:
|
||||
return float(config.get(key, default))
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
def safe_bool(key: str, default: bool) -> bool:
|
||||
return bool(config.get(key, default))
|
||||
|
||||
position = safe_str("position", DEFAULT_POSITION)
|
||||
position = POSITION_ALIASES.get(position, position)
|
||||
if position not in POSITION_ALIGNMENT:
|
||||
position = DEFAULT_POSITION
|
||||
|
||||
return cls(
|
||||
font_name=safe_str("font", DEFAULT_FONT),
|
||||
font_size=safe_int("size", DEFAULT_FONT_SIZE),
|
||||
font_color=safe_str("color", DEFAULT_COLOR),
|
||||
bold=safe_bool("bold", False),
|
||||
italic=safe_bool("italic", False),
|
||||
stroke_enabled=safe_bool("stroke_enabled", True),
|
||||
stroke_color=safe_str("stroke_color", DEFAULT_STROKE_COLOR),
|
||||
stroke_width=safe_float("stroke_width", DEFAULT_STROKE_WIDTH),
|
||||
shadow_enabled=safe_bool("shadow_enabled", False),
|
||||
shadow_color=safe_str("shadow_color", "#000000"),
|
||||
shadow_offset_x=safe_int("shadow_offset_x", 2),
|
||||
shadow_offset_y=safe_int("shadow_offset_y", 2),
|
||||
shadow_blur=safe_float("shadow_blur", 0.0),
|
||||
background_enabled=safe_bool("background_enabled", False),
|
||||
background_color=safe_str("background_color", "#000000"),
|
||||
background_opacity=max(0.0, min(1.0, safe_float("background_opacity", 0.5))),
|
||||
background_padding=safe_int("background_padding", 8),
|
||||
background_radius=safe_int("background_radius", 4),
|
||||
position=position,
|
||||
margin_v=safe_int("margin_v", 60),
|
||||
margin_l=safe_int("margin_l", 40),
|
||||
margin_r=safe_int("margin_r", 40),
|
||||
max_chars_per_line=safe_int("max_chars_per_line", DEFAULT_MAX_CHARS_PER_LINE),
|
||||
line_spacing=safe_int("line_spacing", 0),
|
||||
fade_in=max(0.0, safe_float("fade_in", 0.0)),
|
||||
fade_out=max(0.0, safe_float("fade_out", 0.0)),
|
||||
animation_type=safe_str("animation_type", "none"),
|
||||
)
|
||||
|
||||
@property
|
||||
def alignment(self) -> int:
|
||||
"""获取 ASS alignment 编号."""
|
||||
return POSITION_ALIGNMENT.get(self.position, 2)
|
||||
|
||||
@property
|
||||
def ass_font_color(self) -> str:
|
||||
"""ASS 格式颜色 &HAABBGGRR."""
|
||||
return _hex_to_ass_color(self.font_color)
|
||||
|
||||
@property
|
||||
def ass_stroke_color(self) -> str:
|
||||
return _hex_to_ass_color(self.stroke_color)
|
||||
|
||||
@property
|
||||
def ass_shadow_color(self) -> str:
|
||||
return _hex_to_ass_color(self.shadow_color)
|
||||
|
||||
@property
|
||||
def ass_background_color(self) -> str:
|
||||
"""背景框颜色(ASS BackColour),带透明度."""
|
||||
alpha_hex = _opacity_to_ass_alpha(self.background_opacity)
|
||||
color_bgr = _hex_to_ass_bgr(self.background_color)
|
||||
return f"&H{alpha_hex}{color_bgr}"
|
||||
|
||||
|
||||
# ── 工具函数 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _hex_to_ass_color(hex_color: str) -> str:
|
||||
"""HEX → ASS 颜色 &HAABBGGRR(默认不透明)."""
|
||||
hex_color = hex_color.lstrip("#")
|
||||
if len(hex_color) != 6:
|
||||
return "&H00FFFFFF"
|
||||
r, g, b = hex_color[0:2], hex_color[2:4], hex_color[4:6]
|
||||
return f"&H00{b.upper()}{g.upper()}{r.upper()}"
|
||||
|
||||
|
||||
def _hex_to_ass_bgr(hex_color: str) -> str:
|
||||
"""HEX → ASS BGR 部分(不含 alpha)."""
|
||||
hex_color = hex_color.lstrip("#")
|
||||
if len(hex_color) != 6:
|
||||
return "FFFFFF"
|
||||
r, g, b = hex_color[0:2], hex_color[2:4], hex_color[4:6]
|
||||
return f"{b.upper()}{g.upper()}{r.upper()}"
|
||||
|
||||
|
||||
def _opacity_to_ass_alpha(opacity: float) -> str:
|
||||
"""不透明度 → ASS alpha(00=不透明,FF=完全透明)."""
|
||||
alpha = 255 - int(opacity * 255)
|
||||
return f"{alpha:02X}"
|
||||
|
||||
|
||||
def _escape_ass_text(text: str) -> str:
|
||||
"""转义 ASS 文本特殊字符."""
|
||||
text = text.replace("\r\n", "\\N").replace("\n", "\\N").replace("\r", "\\N")
|
||||
text = text.replace("{", "(").replace("}", ")")
|
||||
return text
|
||||
|
||||
|
||||
def _format_ass_time(seconds: float) -> str:
|
||||
"""秒 → ASS 时间格式 H:MM:SS.cc."""
|
||||
hours = int(seconds // 3600)
|
||||
minutes = int((seconds % 3600) // 60)
|
||||
secs = seconds % 60
|
||||
return f"{hours}:{minutes:02d}:{secs:05.2f}"
|
||||
|
||||
|
||||
def _wrap_text(text: str, max_chars: int) -> list[str]:
|
||||
"""按字数换行,优先标点断开."""
|
||||
if len(text) <= max_chars:
|
||||
return [text]
|
||||
|
||||
lines: list[str] = []
|
||||
remaining = text
|
||||
|
||||
while len(remaining) > max_chars:
|
||||
break_point = max_chars
|
||||
punctuations = ",。!?、;:,.;:!?"
|
||||
|
||||
for i in range(max_chars, max_chars // 2, -1):
|
||||
if i < len(remaining) and remaining[i] in punctuations:
|
||||
break_point = i + 1
|
||||
break
|
||||
|
||||
lines.append(remaining[:break_point])
|
||||
remaining = remaining[break_point:]
|
||||
|
||||
if remaining:
|
||||
lines.append(remaining)
|
||||
|
||||
return lines
|
||||
|
||||
|
||||
# ── 字幕片段 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class SubtitleSegment:
|
||||
"""单个字幕片段."""
|
||||
|
||||
start: float # 开始时间(秒)
|
||||
end: float # 结束时间(秒)
|
||||
text: str # 字幕文本
|
||||
style_name: str = "Default" # 使用的样式名
|
||||
|
||||
|
||||
# ── 字幕渲染引擎 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
@@ -52,6 +52,13 @@ from video_processing.trim_engine import TrimConfig, TrimEngine, extract_trim_fr
|
||||
from video_processing.tts_engine import TtsEngine
|
||||
from video_processing.watermark_engine import WatermarkConfig, WatermarkEngine
|
||||
|
||||
from packages.domain.render_layer_utils import LAYER_Z_INDEX as _IMPORTED_LAYER_Z_INDEX
|
||||
from packages.domain.render_layer_utils import can_pass_through as _can_pass_through_pure
|
||||
from packages.domain.render_layer_utils import clip_adjusted_duration as _clip_adjusted_duration_pure
|
||||
from packages.domain.render_layer_utils import clip_effective_duration as _clip_effective_duration_pure
|
||||
from packages.domain.render_layer_utils import clip_playback_speed as _clip_playback_speed_pure
|
||||
from packages.domain.render_layer_utils import estimate_total_duration as _estimate_total_duration_pure
|
||||
from packages.domain.render_layer_utils import resolve_layer_role as _resolve_layer_role_pure
|
||||
from packages.domain.tts_config import TtsConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -107,47 +114,16 @@ class RenderResult:
|
||||
|
||||
|
||||
def _resolve_layer_role(clip_type: str, config: dict[str, Any]) -> str:
|
||||
"""根据 clip_type 和 config.role 确定图层角色。
|
||||
"""根据 clip_type 和 config.role 确定图层角色(向后兼容别名)。
|
||||
|
||||
映射规则:
|
||||
intro / outro → "main"(按 order 排在首/尾)
|
||||
overlay → "overlay"(画中画叠加,z=1)
|
||||
corner_voice → "corner_voice"(右上角小窗,z=1)
|
||||
background → "background"(全屏底图,z=0)
|
||||
b_roll → "broll"(z=0)
|
||||
main + config.role=b_roll → "broll"
|
||||
main (default) → "main"
|
||||
实际实现移至 packages.domain.render_layer_utils.resolve_layer_role。
|
||||
"""
|
||||
role = config.get("role", "")
|
||||
|
||||
if clip_type in ("intro", "outro"):
|
||||
return "main"
|
||||
if clip_type == "overlay":
|
||||
return "overlay"
|
||||
if clip_type == "corner_voice":
|
||||
return "corner_voice"
|
||||
if clip_type == "background":
|
||||
return "background"
|
||||
if clip_type == "b_roll":
|
||||
return "broll"
|
||||
# main type
|
||||
if role == "b_roll":
|
||||
return "broll"
|
||||
if role == "audio":
|
||||
return "audio"
|
||||
return "main"
|
||||
return _resolve_layer_role_pure(clip_type, config)
|
||||
|
||||
|
||||
# ── 图层默认 z_index ─────────────────────────────────────────────────────────
|
||||
|
||||
_LAYER_Z_INDEX: dict[str, int] = {
|
||||
"background": -1,
|
||||
"broll": 0,
|
||||
"main": 0,
|
||||
"overlay": 1,
|
||||
"corner_voice": 1,
|
||||
"audio": 2,
|
||||
}
|
||||
_LAYER_Z_INDEX: dict[str, int] = _IMPORTED_LAYER_Z_INDEX
|
||||
|
||||
# 图层默认 PiP 位置(相对输出画布的偏移)
|
||||
_PIP_SCALE = 0.25 # PiP 占主画面的比例
|
||||
@@ -489,29 +465,9 @@ class UnifiedRenderService:
|
||||
def _estimate_total_duration(self, layers: list[RenderLayer]) -> float:
|
||||
"""估算视频总时长(用于字幕等需要)。
|
||||
|
||||
取主图层(main/broll/background)的总时长,转场重叠按 transition_duration 估算。
|
||||
实际实现移至 packages.domain.render_layer_utils.estimate_total_duration。
|
||||
"""
|
||||
# 找主图层(第一个有视频内容的图层)
|
||||
main_layer = None
|
||||
for role in ("main", "broll", "background"):
|
||||
for layer in layers:
|
||||
if layer.role == role:
|
||||
main_layer = layer
|
||||
break
|
||||
if main_layer:
|
||||
break
|
||||
|
||||
if not main_layer or not main_layer.clips:
|
||||
return 0.0
|
||||
|
||||
total = sum(UnifiedRenderService._clip_adjusted_duration(c) for c in main_layer.clips)
|
||||
|
||||
# 减去转场重叠时间(粗略估算)
|
||||
n_clips = len(main_layer.clips)
|
||||
if n_clips > 1:
|
||||
total -= (n_clips - 1) * self.transition_duration
|
||||
|
||||
return max(0.1, total)
|
||||
return _estimate_total_duration_pure(layers, self.transition_duration)
|
||||
|
||||
def _maybe_generate_ass(self, video_duration: float) -> Path | None:
|
||||
"""根据 plan.config 生成 ASS 字幕文件。
|
||||
@@ -1869,10 +1825,11 @@ class UnifiedRenderService:
|
||||
|
||||
@staticmethod
|
||||
def _clip_effective_duration(clip: ResolvedClip) -> float:
|
||||
"""计算 clip 的有效时长(原速 trim 后时长)."""
|
||||
if clip.duration > 0:
|
||||
return min(clip.duration, clip.actual_duration) if clip.actual_duration > 0 else clip.duration
|
||||
return clip.actual_duration if clip.actual_duration > 0 else 0.0
|
||||
"""计算 clip 的有效时长(原速 trim 后时长)。
|
||||
|
||||
实际实现移至 packages.domain.render_layer_utils.clip_effective_duration。
|
||||
"""
|
||||
return _clip_effective_duration_pure(clip.duration, clip.actual_duration)
|
||||
|
||||
# ── 画中画(PiP)相关方法 ──────────────────────────────────────────────────
|
||||
|
||||
@@ -1969,17 +1926,20 @@ class UnifiedRenderService:
|
||||
|
||||
@staticmethod
|
||||
def _clip_speed(clip: ResolvedClip) -> float:
|
||||
"""获取 clip 的播放速度,无效值回退到 1.0."""
|
||||
speed = getattr(clip, "playback_speed", 1.0)
|
||||
if not isinstance(speed, (int, float)) or speed <= 0:
|
||||
return 1.0
|
||||
return float(speed)
|
||||
"""获取 clip 的播放速度,无效值回退到 1.0。
|
||||
|
||||
实际实现移至 packages.domain.render_layer_utils.clip_playback_speed。
|
||||
"""
|
||||
return _clip_playback_speed_pure(getattr(clip, "playback_speed", 1.0))
|
||||
|
||||
@staticmethod
|
||||
def _clip_adjusted_duration(clip: ResolvedClip) -> float:
|
||||
"""计算调速后的 clip 实际时长(用于拼接计算)."""
|
||||
base = UnifiedRenderService._clip_effective_duration(clip)
|
||||
speed = UnifiedRenderService._clip_speed(clip)
|
||||
if abs(speed - 1.0) < 1e-6:
|
||||
return base
|
||||
return base / speed
|
||||
"""计算调速后的 clip 实际时长(用于拼接计算)。
|
||||
|
||||
实际实现移至 packages.domain.render_layer_utils.clip_adjusted_duration。
|
||||
"""
|
||||
return _clip_adjusted_duration_pure(
|
||||
clip.duration,
|
||||
clip.actual_duration,
|
||||
getattr(clip, "playback_speed", 1.0),
|
||||
)
|
||||
|
||||
@@ -20,19 +20,21 @@ import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from video_processing.ffmpeg_utils import probe_duration
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.domain.bgm_utils import merge_bgm_config
|
||||
from video_processing.ffmpeg_utils import probe_duration
|
||||
from worker_app.tasks.generation_plan_builder import VirtualClip as _VirtualClip
|
||||
from worker_app.tasks.generation_plan_builder import VirtualPlan as _VirtualPlan
|
||||
from worker_app.tasks.generation_plan_builder import apply_template_clip_effects as _apply_template_clip_effects
|
||||
from worker_app.tasks.generation_plan_builder import (
|
||||
VirtualPlan as _VirtualPlan,
|
||||
VirtualClip as _VirtualClip,
|
||||
build_error_info as _build_error_info,
|
||||
extract_intro_outro_from_clip_configs as _extract_intro_outro_from_clip_configs,
|
||||
apply_template_clip_effects as _apply_template_clip_effects,
|
||||
build_clips_by_mode,
|
||||
)
|
||||
from worker_app.tasks.generation_plan_builder import build_error_info as _build_error_info
|
||||
from worker_app.tasks.generation_plan_builder import (
|
||||
extract_intro_outro_from_clip_configs as _extract_intro_outro_from_clip_configs,
|
||||
)
|
||||
|
||||
from packages.domain.bgm_utils import merge_bgm_config
|
||||
|
||||
OUTPUT_WIDTH = 1280
|
||||
OUTPUT_HEIGHT = 720
|
||||
|
||||
@@ -14,70 +14,18 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
SQLAlchemyIngestJobRepository,
|
||||
)
|
||||
from packages.domain import Asset, AssetStatus, IngestJobStatus
|
||||
from packages.domain.media_validation import (
|
||||
MIN_AUDIO_FILE_SIZE,
|
||||
MIN_IMAGE_FILE_SIZE,
|
||||
MIN_VIDEO_FILE_SIZE,
|
||||
SUPPORTED_VIDEO_CODECS,
|
||||
)
|
||||
from packages.domain.media_validation import is_valid_media as _is_valid_media
|
||||
from packages.domain.media_validation import safe_parse_fps as _safe_parse_fps
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
# 最小有效文件大小(字节):小于此值的直接判为无效,避免文本/空文件伪装成媒体
|
||||
MIN_VIDEO_FILE_SIZE = 1024 # 1KB
|
||||
MIN_AUDIO_FILE_SIZE = 100 # 100B
|
||||
MIN_IMAGE_FILE_SIZE = 100 # 100B
|
||||
|
||||
# 支持的视频编码格式(白名单,尽可能放宽)
|
||||
# 渲染引擎会在 concat 前统一转码为 h264,因此只要 ffprobe 能识别的视频编码都允许 ingested
|
||||
SUPPORTED_VIDEO_CODECS = {
|
||||
"h264",
|
||||
"avc1",
|
||||
"avc", # H.264 / AVC
|
||||
"hevc",
|
||||
"h265",
|
||||
"hev1",
|
||||
"hvc1", # H.265 / HEVC
|
||||
"vp9",
|
||||
"vp09", # VP9
|
||||
"av1",
|
||||
"av01", # AV1
|
||||
"vp8",
|
||||
"vp08", # VP8
|
||||
"mpeg4",
|
||||
"mp4v", # MPEG-4
|
||||
"mpeg2video",
|
||||
"mpg2", # MPEG-2
|
||||
"wmv2",
|
||||
"wmv1",
|
||||
"vc1", # WMV / VC-1
|
||||
"flv1",
|
||||
"flv",
|
||||
"vp6f", # Flash / FLV
|
||||
"theora",
|
||||
"ogg", # Theora
|
||||
"prores",
|
||||
"prores_ks",
|
||||
"apcn",
|
||||
"apch",
|
||||
"apco",
|
||||
"apcs",
|
||||
"ap4h",
|
||||
"ap4x", # Apple ProRes
|
||||
"dnxhd",
|
||||
"dnxhr", # DNxHD / DNxHR
|
||||
}
|
||||
|
||||
|
||||
def _safe_parse_fps(fps_str: str) -> float:
|
||||
"""Safely parse fps from a fraction string like \"30/1\" or \"30000/1001\"."""
|
||||
try:
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/", 1)
|
||||
den_val = float(den)
|
||||
if den_val == 0:
|
||||
return 0.0
|
||||
return float(num) / den_val
|
||||
return float(fps_str)
|
||||
except (ValueError, ZeroDivisionError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def extract_media_metadata(file_url: str, media_type: str) -> tuple[dict, bool]:
|
||||
"""
|
||||
提取媒体文件的元数据。
|
||||
@@ -207,38 +155,6 @@ def extract_media_metadata(file_url: str, media_type: str) -> tuple[dict, bool]:
|
||||
return metadata, success
|
||||
|
||||
|
||||
def _is_valid_media(metadata: dict, media_type: str) -> bool:
|
||||
"""根据元数据判断文件是否为有效媒体文件。
|
||||
|
||||
Args:
|
||||
metadata: extract_media_metadata 返回的元数据
|
||||
media_type: 媒体类型
|
||||
|
||||
Returns:
|
||||
True 表示文件有效
|
||||
"""
|
||||
size = int(metadata.get("size_bytes", 0))
|
||||
|
||||
if media_type == "video":
|
||||
duration = float(metadata.get("duration", 0))
|
||||
if size < MIN_VIDEO_FILE_SIZE or duration <= 0:
|
||||
return False
|
||||
# 编码格式校验:只排除明确非视频的编码格式,只要 ffprobe 能识别的视频编码都允许
|
||||
# 渲染引擎会在 concat 前统一转码为 h264 yuv420p,ingest 层不再做严格的编码拦截
|
||||
codec = str(metadata.get("codec", "")).lower()
|
||||
if codec and codec not in SUPPORTED_VIDEO_CODECS:
|
||||
logger.info("检测到非白名单视频编码 %s,仍允许 ingested,渲染层会统一转码", codec)
|
||||
return True
|
||||
if media_type == "audio":
|
||||
duration = float(metadata.get("duration", 0))
|
||||
return size >= MIN_AUDIO_FILE_SIZE and duration > 0
|
||||
if media_type == "image":
|
||||
width = int(metadata.get("width", 0))
|
||||
height = int(metadata.get("height", 0))
|
||||
return size >= MIN_IMAGE_FILE_SIZE and width > 0 and height > 0
|
||||
return False
|
||||
|
||||
|
||||
@celery_app.task(name="worker.ingest_asset")
|
||||
def ingest_asset(job_id: str) -> dict:
|
||||
"""
|
||||
|
||||
Executable
+267
@@ -0,0 +1,267 @@
|
||||
"""片段操作工具 — EditPlanClip 分割/合并等纯逻辑操作。
|
||||
|
||||
从 edit_plan_service.py 抽离的纯函数集合,专门负责:
|
||||
- 片段分割:将一个片段从指定位置拆分为两个
|
||||
- 片段合并:将多个连续片段合并为一个
|
||||
- Order 重排计算
|
||||
|
||||
所有函数均为纯函数,不依赖数据库或外部 IO。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
DEFAULT_SPLIT_DURATION = 5.0
|
||||
ROUND_PRECISION = 3
|
||||
|
||||
|
||||
# ── 数据结构 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SplitResult:
|
||||
"""片段分割结果。"""
|
||||
|
||||
left_duration: float
|
||||
right_duration: float
|
||||
right_start_time: float
|
||||
left_trim_end: float
|
||||
right_trim_start: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MergeResult:
|
||||
"""片段合并结果。"""
|
||||
|
||||
total_duration: float
|
||||
merged_text: str
|
||||
merged_config: dict[str, Any]
|
||||
first_order: int
|
||||
shift_amount: int
|
||||
|
||||
|
||||
# ── 分割 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def validate_split_time(split_time: float, duration: float) -> None:
|
||||
"""校验分割时间是否合法。
|
||||
|
||||
Args:
|
||||
split_time: 分割点(秒)
|
||||
duration: 原片段时长(秒)
|
||||
|
||||
Raises:
|
||||
ValueError: 分割时间不在 (0, duration) 范围内
|
||||
"""
|
||||
if split_time <= 0 or split_time >= duration:
|
||||
raise ValueError(f"分割时间必须在 (0, {duration:.3f}) 范围内,当前: {split_time}")
|
||||
|
||||
|
||||
def calculate_split(
|
||||
duration: float,
|
||||
split_time: float,
|
||||
start_time: float = 0.0,
|
||||
*,
|
||||
precision: int = ROUND_PRECISION,
|
||||
) -> SplitResult:
|
||||
"""计算片段分割后的各项参数。
|
||||
|
||||
左半部分:从 0 到 split_time
|
||||
右半部分:从 split_time 到 duration
|
||||
|
||||
Args:
|
||||
duration: 原片段时长(秒)
|
||||
split_time: 分割点(秒)
|
||||
start_time: 原片段起始时间(秒),右半部分 start_time 需要加上 left_duration
|
||||
precision: 小数精度(默认 3 位,即毫秒)
|
||||
|
||||
Returns:
|
||||
SplitResult 包含左右部分的时长、右半部分 start_time、trim 信息
|
||||
"""
|
||||
validate_split_time(split_time, duration)
|
||||
|
||||
left_duration = round(split_time, precision)
|
||||
right_duration = round(duration - split_time, precision)
|
||||
right_start_time = round(start_time + left_duration, precision)
|
||||
|
||||
return SplitResult(
|
||||
left_duration=left_duration,
|
||||
right_duration=right_duration,
|
||||
right_start_time=right_start_time,
|
||||
left_trim_end=right_duration,
|
||||
right_trim_start=left_duration,
|
||||
)
|
||||
|
||||
|
||||
# ── 合并 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def validate_merge_clips(clips: list[Any]) -> tuple[str, int]:
|
||||
"""校验待合并的片段列表。
|
||||
|
||||
校验项:
|
||||
1. 至少 2 个片段
|
||||
2. 属于同一计划
|
||||
3. order 连续
|
||||
4. 类型一致
|
||||
|
||||
Args:
|
||||
clips: 按任意顺序排列的片段列表(会自动按 order 排序)
|
||||
|
||||
Returns:
|
||||
(plan_id, first_order) 元组
|
||||
|
||||
Raises:
|
||||
ValueError: 校验失败
|
||||
"""
|
||||
if len(clips) < 2:
|
||||
raise ValueError("至少需要 2 个片段才能合并")
|
||||
|
||||
# 校验:同一计划
|
||||
plan_id = clips[0].plan_id
|
||||
for c in clips[1:]:
|
||||
if c.plan_id != plan_id:
|
||||
raise ValueError("只能合并同一计划下的片段")
|
||||
|
||||
# 按 order 排序
|
||||
sorted_clips = sorted(clips, key=lambda c: c.order)
|
||||
|
||||
# 校验:order 连续
|
||||
for i in range(1, len(sorted_clips)):
|
||||
if sorted_clips[i].order != sorted_clips[i - 1].order + 1:
|
||||
raise ValueError(f"片段不连续:order {sorted_clips[i-1].order} → {sorted_clips[i].order}")
|
||||
|
||||
# 校验:类型一致
|
||||
clip_type = sorted_clips[0].clip_type
|
||||
for c in sorted_clips[1:]:
|
||||
if c.clip_type != clip_type:
|
||||
raise ValueError("只能合并相同类型的片段")
|
||||
|
||||
return plan_id, sorted_clips[0].order
|
||||
|
||||
|
||||
def calculate_merge(
|
||||
clips: list[Any],
|
||||
*,
|
||||
precision: int = ROUND_PRECISION,
|
||||
) -> MergeResult:
|
||||
"""计算多个片段合并后的参数。
|
||||
|
||||
合并规则:
|
||||
- 时长:所有片段时长之和
|
||||
- 文案:用换行连接非空文案
|
||||
- config:后面的覆盖前面的,移除 trim_start/trim_end
|
||||
- first_order:第一个片段的 order
|
||||
- shift_amount:合并后 order 前移位数(n-1)
|
||||
|
||||
Args:
|
||||
clips: 待合并片段列表(会自动按 order 排序)
|
||||
precision: 时长精度(默认 3 位)
|
||||
|
||||
Returns:
|
||||
MergeResult 合并结果
|
||||
"""
|
||||
if not clips:
|
||||
raise ValueError("合并的片段列表不能为空")
|
||||
|
||||
# 按 order 排序
|
||||
sorted_clips = sorted(clips, key=lambda c: c.order)
|
||||
|
||||
# 总时长
|
||||
total_duration = round(sum(c.duration for c in sorted_clips), precision)
|
||||
|
||||
# 合并文案
|
||||
merged_text = "\n".join(c.text_content for c in sorted_clips if c.text_content and c.text_content.strip())
|
||||
|
||||
# 合并 config(后面的覆盖前面的)
|
||||
merged_config: dict[str, Any] = {}
|
||||
for c in sorted_clips:
|
||||
if c.config:
|
||||
merged_config.update(c.config)
|
||||
# 清理 trim 相关字段(合并后就是完整片段了)
|
||||
merged_config.pop("trim_start", None)
|
||||
merged_config.pop("trim_end", None)
|
||||
|
||||
first_order = sorted_clips[0].order
|
||||
shift_amount = len(sorted_clips) - 1
|
||||
|
||||
return MergeResult(
|
||||
total_duration=total_duration,
|
||||
merged_text=merged_text,
|
||||
merged_config=merged_config,
|
||||
first_order=first_order,
|
||||
shift_amount=shift_amount,
|
||||
)
|
||||
|
||||
|
||||
# ── Order 重排 ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def calculate_reorder_new_orders(
|
||||
ordered_ids: list[str],
|
||||
current_items: list[Any],
|
||||
*,
|
||||
id_attr: str = "id",
|
||||
order_attr: str = "order",
|
||||
) -> dict[str, int]:
|
||||
"""根据新顺序计算每个 item 的新 order 值。
|
||||
|
||||
Args:
|
||||
ordered_ids: 按新顺序排列的 ID 列表
|
||||
current_items: 当前所有 item 列表
|
||||
id_attr: ID 属性名
|
||||
order_attr: order 属性名
|
||||
|
||||
Returns:
|
||||
{item_id: new_order} 映射
|
||||
|
||||
Raises:
|
||||
ValueError: ID 列表与当前 items 不匹配
|
||||
"""
|
||||
current_ids = {getattr(c, id_attr) for c in current_items}
|
||||
ordered_id_set = set(ordered_ids)
|
||||
|
||||
if ordered_id_set != current_ids:
|
||||
raise ValueError("ID 列表与当前 items 不匹配")
|
||||
|
||||
return {item_id: idx for idx, item_id in enumerate(ordered_ids)}
|
||||
|
||||
|
||||
def calculate_shift_orders(
|
||||
items: list[Any],
|
||||
threshold_order: int,
|
||||
shift: int,
|
||||
*,
|
||||
excluded_ids: set[str] | None = None,
|
||||
order_attr: str = "order",
|
||||
id_attr: str = "id",
|
||||
) -> list[tuple[Any, int]]:
|
||||
"""计算 order 需要偏移的 items 及新 order 值。
|
||||
|
||||
Args:
|
||||
items: 所有 item 列表
|
||||
threshold_order: 只处理 order > threshold_order 的 item
|
||||
shift: 偏移量(正数加,负数减)
|
||||
excluded_ids: 排除的 ID 集合
|
||||
order_attr: order 属性名
|
||||
id_attr: ID 属性名
|
||||
|
||||
Returns:
|
||||
[(item, new_order), ...] 列表
|
||||
"""
|
||||
excluded = excluded_ids or set()
|
||||
result: list[tuple[Any, int]] = []
|
||||
|
||||
for item in items:
|
||||
item_id = getattr(item, id_attr)
|
||||
if item_id in excluded:
|
||||
continue
|
||||
current_order = getattr(item, order_attr)
|
||||
if current_order > threshold_order:
|
||||
result.append((item, current_order + shift))
|
||||
|
||||
return result
|
||||
Executable
+110
@@ -0,0 +1,110 @@
|
||||
"""媒体文件有效性校验与元数据解析工具。
|
||||
|
||||
从 worker ingest 任务中抽取的纯逻辑模块,包含:
|
||||
- FPS 解析:从分数格式字符串(如 30000/1001)安全解析帧率
|
||||
- 媒体有效性校验:根据元数据判断视频/音频/图片文件是否有效
|
||||
- 常量定义:最小文件大小、支持的视频编码白名单
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
# 最小有效文件大小(字节):小于此值的直接判为无效,避免文本/空文件伪装成媒体
|
||||
MIN_VIDEO_FILE_SIZE = 1024 # 1KB
|
||||
MIN_AUDIO_FILE_SIZE = 100 # 100B
|
||||
MIN_IMAGE_FILE_SIZE = 100 # 100B
|
||||
|
||||
# 支持的视频编码格式(白名单,尽可能放宽)
|
||||
# 渲染引擎会在 concat 前统一转码为 h264,因此只要 ffprobe 能识别的视频编码都允许 ingested
|
||||
SUPPORTED_VIDEO_CODECS: frozenset[str] = frozenset(
|
||||
{
|
||||
"h264",
|
||||
"avc1",
|
||||
"avc", # H.264 / AVC
|
||||
"hevc",
|
||||
"h265",
|
||||
"hev1",
|
||||
"hvc1", # H.265 / HEVC
|
||||
"vp9",
|
||||
"vp09", # VP9
|
||||
"av1",
|
||||
"av01", # AV1
|
||||
"vp8",
|
||||
"vp08", # VP8
|
||||
"mpeg4",
|
||||
"mp4v", # MPEG-4
|
||||
"mpeg2video",
|
||||
"mpg2", # MPEG-2
|
||||
"wmv2",
|
||||
"wmv1",
|
||||
"vc1", # WMV / VC-1
|
||||
"flv1",
|
||||
"flv",
|
||||
"vp6f", # Flash / FLV
|
||||
"theora",
|
||||
"ogg", # Theora
|
||||
"prores",
|
||||
"prores_ks",
|
||||
"apcn",
|
||||
"apch",
|
||||
"apco",
|
||||
"apcs",
|
||||
"ap4h",
|
||||
"ap4x", # Apple ProRes
|
||||
"dnxhd",
|
||||
"dnxhr", # DNxHD / DNxHR
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def safe_parse_fps(fps_str: str) -> float:
|
||||
"""Safely parse fps from a fraction string like "30/1" or "30000/1001".
|
||||
|
||||
Args:
|
||||
fps_str: FPS 字符串,支持小数格式("30.0")或分数格式("30000/1001")
|
||||
|
||||
Returns:
|
||||
解析得到的帧率浮点数;解析失败或分母为0时返回 0.0
|
||||
"""
|
||||
try:
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/", 1)
|
||||
den_val = float(den)
|
||||
if den_val == 0:
|
||||
return 0.0
|
||||
return float(num) / den_val
|
||||
return float(fps_str)
|
||||
except (ValueError, ZeroDivisionError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def is_valid_media(metadata: dict, media_type: str) -> bool:
|
||||
"""根据元数据判断文件是否为有效媒体文件。
|
||||
|
||||
Args:
|
||||
metadata: 媒体元数据字典,可能包含 size_bytes / duration / codec / width / height 等
|
||||
media_type: 媒体类型(video / audio / image)
|
||||
|
||||
Returns:
|
||||
True 表示文件有效
|
||||
"""
|
||||
size = int(metadata.get("size_bytes", 0))
|
||||
|
||||
if media_type == "video":
|
||||
duration = float(metadata.get("duration", 0))
|
||||
if size < MIN_VIDEO_FILE_SIZE or duration <= 0:
|
||||
return False
|
||||
# 编码格式校验:只排除明确非视频的编码格式,只要 ffprobe 能识别的视频编码都允许
|
||||
# 渲染引擎会在 concat 前统一转码为 h264 yuv420p,ingest 层不再做严格的编码拦截
|
||||
codec = str(metadata.get("codec", "")).lower()
|
||||
if codec and codec not in SUPPORTED_VIDEO_CODECS:
|
||||
# 非白名单编码仍允许通过,仅记录日志(调用方负责日志)
|
||||
pass
|
||||
return True
|
||||
if media_type == "audio":
|
||||
duration = float(metadata.get("duration", 0))
|
||||
return size >= MIN_AUDIO_FILE_SIZE and duration > 0
|
||||
if media_type == "image":
|
||||
width = int(metadata.get("width", 0))
|
||||
height = int(metadata.get("height", 0))
|
||||
return size >= MIN_IMAGE_FILE_SIZE and width > 0 and height > 0
|
||||
return False
|
||||
@@ -321,11 +321,7 @@ def create_clips_from_configs(
|
||||
clip_type = cfg.clip_type.value if hasattr(cfg.clip_type, "value") else cfg.clip_type
|
||||
|
||||
# transition_effect 可能是枚举或字符串
|
||||
transition = (
|
||||
cfg.transition_effect.value
|
||||
if hasattr(cfg.transition_effect, "value")
|
||||
else cfg.transition_effect
|
||||
)
|
||||
transition = cfg.transition_effect.value if hasattr(cfg.transition_effect, "value") else cfg.transition_effect
|
||||
|
||||
# 从 clip config 中解析 playback_speed(兼容 speed_ratio 字段名)
|
||||
clip_cfg = cfg.config or {}
|
||||
|
||||
Executable
+241
@@ -0,0 +1,241 @@
|
||||
"""渲染图层工具函数 — 纯函数集合.
|
||||
|
||||
从 unified_render_service.py 抽离的纯逻辑,负责:
|
||||
- clip 时长计算(有效时长、调速后时长)
|
||||
- clip_type → layer_role 映射
|
||||
- 总时长估算
|
||||
- 图层默认属性(z_index 等)
|
||||
|
||||
所有函数均为纯函数,不依赖 FFmpeg、数据库或外部 IO。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
# ── 图层角色定义 ─────────────────────────────────────────────────────────────
|
||||
|
||||
# 图层默认 z_index 映射
|
||||
LAYER_Z_INDEX: dict[str, int] = {
|
||||
"background": -1,
|
||||
"broll": 0,
|
||||
"main": 0,
|
||||
"overlay": 1,
|
||||
"corner_voice": 1,
|
||||
"audio": 2,
|
||||
}
|
||||
|
||||
# 图层默认 PiP 缩放比例(相对于主画面)
|
||||
PIP_DEFAULT_SCALE = 0.25
|
||||
|
||||
# 主视频图层角色(用于总时长计算、直通判断等)
|
||||
MAIN_LAYER_ROLES = frozenset({"main", "broll", "background"})
|
||||
|
||||
|
||||
# ── clip_type → layer_role 映射 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def resolve_layer_role(clip_type: str, config: dict[str, Any] | None = None) -> str:
|
||||
"""根据 clip_type 和 config.role 确定图层角色。
|
||||
|
||||
映射规则:
|
||||
intro / outro → "main"(按 order 排在首/尾)
|
||||
overlay → "overlay"(画中画叠加,z=1)
|
||||
corner_voice → "corner_voice"(右上角小窗,z=1)
|
||||
background → "background"(全屏底图,z=0)
|
||||
b_roll → "broll"(z=0)
|
||||
main + config.role=b_roll → "broll"
|
||||
main + config.role=audio → "audio"
|
||||
main (default) → "main"
|
||||
|
||||
Args:
|
||||
clip_type: 片段类型字符串
|
||||
config: 片段配置字典(可选)
|
||||
|
||||
Returns:
|
||||
图层角色字符串
|
||||
"""
|
||||
role = (config or {}).get("role", "") if config else ""
|
||||
|
||||
if clip_type in ("intro", "outro"):
|
||||
return "main"
|
||||
if clip_type == "overlay":
|
||||
return "overlay"
|
||||
if clip_type == "corner_voice":
|
||||
return "corner_voice"
|
||||
if clip_type == "background":
|
||||
return "background"
|
||||
if clip_type == "b_roll":
|
||||
return "broll"
|
||||
# main type
|
||||
if role == "b_roll":
|
||||
return "broll"
|
||||
if role == "audio":
|
||||
return "audio"
|
||||
return "main"
|
||||
|
||||
|
||||
def get_layer_z_index(role: str) -> int:
|
||||
"""获取图层角色的默认 z_index。
|
||||
|
||||
Args:
|
||||
role: 图层角色
|
||||
|
||||
Returns:
|
||||
z_index 值,未知角色返回 0
|
||||
"""
|
||||
return LAYER_Z_INDEX.get(role, 0)
|
||||
|
||||
|
||||
# ── clip 时长计算 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def clip_effective_duration(
|
||||
duration: float,
|
||||
actual_duration: float = 0.0,
|
||||
) -> float:
|
||||
"""计算 clip 的有效时长(原速 trim 后时长)。
|
||||
|
||||
规则:
|
||||
- duration > 0: min(duration, actual_duration),actual=0 时用 duration
|
||||
- duration <= 0: actual_duration,actual=0 时返回 0
|
||||
|
||||
Args:
|
||||
duration: 配置的时长(0 表示使用完整素材)
|
||||
actual_duration: 素材实际时长(probe 后的结果)
|
||||
|
||||
Returns:
|
||||
有效时长(秒)
|
||||
"""
|
||||
if duration > 0:
|
||||
return min(duration, actual_duration) if actual_duration > 0 else duration
|
||||
return actual_duration if actual_duration > 0 else 0.0
|
||||
|
||||
|
||||
def clip_playback_speed(playback_speed: Any) -> float:
|
||||
"""获取 clip 的播放速度,无效值回退到 1.0。
|
||||
|
||||
Args:
|
||||
playback_speed: 播放速度(可为任意类型
|
||||
|
||||
Returns:
|
||||
有效的播放速度(正数)
|
||||
"""
|
||||
if not isinstance(playback_speed, (int, float)):
|
||||
return 1.0
|
||||
if playback_speed <= 0:
|
||||
return 1.0
|
||||
return float(playback_speed)
|
||||
|
||||
|
||||
def clip_adjusted_duration(
|
||||
duration: float,
|
||||
actual_duration: float = 0.0,
|
||||
playback_speed: Any = 1.0,
|
||||
) -> float:
|
||||
"""计算调速后的 clip 实际时长(用于拼接计算)。
|
||||
|
||||
Args:
|
||||
duration: 配置的时长
|
||||
actual_duration: 素材实际时长
|
||||
playback_speed: 播放速度
|
||||
|
||||
Returns:
|
||||
调速后的时长
|
||||
"""
|
||||
base = clip_effective_duration(duration, actual_duration)
|
||||
speed = clip_playback_speed(playback_speed)
|
||||
if abs(speed - 1.0) < 1e-6:
|
||||
return base
|
||||
return base / speed
|
||||
|
||||
|
||||
# ── 总时长估算 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def estimate_total_duration(
|
||||
layers: list[Any],
|
||||
transition_duration: float = 0.0,
|
||||
) -> float:
|
||||
"""估算视频总时长。
|
||||
|
||||
取主图层(main/broll/background)的总调整后时长,减去转场重叠时间。
|
||||
|
||||
Args:
|
||||
layers: 图层列表(每个元素需有 role 和 clips 属性,
|
||||
clips 中元素需有 duration/actual_duration/playback_speed 属性)
|
||||
transition_duration: 转场时长(秒),用于估算重叠时间
|
||||
|
||||
Returns:
|
||||
估算的总时长(秒),最小 0.1
|
||||
"""
|
||||
# 找主图层(第一个有视频内容的图层)
|
||||
main_layer = None
|
||||
for role in ("main", "broll", "background"):
|
||||
for layer in layers:
|
||||
if getattr(layer, "role", None) == role and getattr(layer, "clips", None):
|
||||
main_layer = layer
|
||||
break
|
||||
if main_layer:
|
||||
break
|
||||
|
||||
if not main_layer or not getattr(main_layer, "clips", None):
|
||||
return 0.0
|
||||
|
||||
clips = getattr(main_layer, "clips", [])
|
||||
total = sum(
|
||||
clip_adjusted_duration(
|
||||
duration=getattr(c, "duration", 0),
|
||||
actual_duration=getattr(c, "actual_duration", 0.0),
|
||||
playback_speed=getattr(c, "playback_speed", 1.0),
|
||||
)
|
||||
for c in clips
|
||||
)
|
||||
|
||||
# 减去转场重叠时间(粗略估算)
|
||||
n_clips = len(clips)
|
||||
if n_clips > 1 and transition_duration > 0:
|
||||
total -= (n_clips - 1) * transition_duration
|
||||
|
||||
return max(0.1, total)
|
||||
|
||||
|
||||
# ── 直通 / Stream Copy 判断辅助 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def can_pass_through(
|
||||
layers: list[Any],
|
||||
has_stickers: bool = False,
|
||||
has_watermark: bool = False,
|
||||
) -> bool:
|
||||
"""判断是否可以走直通优化路径(单 clip 简单场景)。
|
||||
|
||||
条件:
|
||||
1. 只有 1 个图层
|
||||
2. 该图层是视频图层(main/broll/background)
|
||||
3. 该图层只有 1 个 clip(无转场需求)
|
||||
4. 没有贴纸
|
||||
5. 没有水印
|
||||
|
||||
Args:
|
||||
layers: 图层列表
|
||||
has_stickers: 是否有贴纸
|
||||
has_watermark: 是否有水印
|
||||
|
||||
Returns:
|
||||
是否可以走直通
|
||||
"""
|
||||
if len(layers) != 1:
|
||||
return False
|
||||
layer = layers[0]
|
||||
role = getattr(layer, "role", "")
|
||||
if role not in MAIN_LAYER_ROLES:
|
||||
return False
|
||||
clips = getattr(layer, "clips", [])
|
||||
if len(clips) != 1:
|
||||
return False
|
||||
if has_stickers:
|
||||
return False
|
||||
if has_watermark:
|
||||
return False
|
||||
return True
|
||||
Executable
+272
@@ -0,0 +1,272 @@
|
||||
"""字幕样式领域模型 — 纯逻辑,无外部依赖.
|
||||
|
||||
抽离自 subtitle_render_engine.py 的数据类和工具函数,
|
||||
方便单测覆盖,同时保持向后兼容。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
# 9宫格位置映射(ASS alignment 编号)
|
||||
POSITION_ALIGNMENT: dict[str, int] = {
|
||||
"top_left": 7,
|
||||
"top_center": 8,
|
||||
"top_right": 9,
|
||||
"middle_left": 4,
|
||||
"center": 5,
|
||||
"middle_right": 6,
|
||||
"bottom_left": 1,
|
||||
"bottom_center": 2,
|
||||
"bottom_right": 3,
|
||||
}
|
||||
|
||||
# 位置简称兼容
|
||||
POSITION_ALIASES: dict[str, str] = {
|
||||
"top": "top_center",
|
||||
"bottom": "bottom_center",
|
||||
"middle": "center",
|
||||
"left": "middle_left",
|
||||
"right": "middle_right",
|
||||
}
|
||||
|
||||
DEFAULT_FONT = "思源黑体"
|
||||
DEFAULT_FONT_SIZE = 24
|
||||
DEFAULT_COLOR = "#FFFFFF"
|
||||
DEFAULT_STROKE_COLOR = "#000000"
|
||||
DEFAULT_STROKE_WIDTH = 1.5
|
||||
DEFAULT_POSITION = "bottom_center"
|
||||
DEFAULT_MAX_CHARS_PER_LINE = 20
|
||||
|
||||
ALLOWED_SUBTITLE_EXTENSIONS = {".srt", ".ass", ".vtt", ".sub"}
|
||||
|
||||
|
||||
# ── 工具函数 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def hex_to_ass_color(hex_color: str) -> str:
|
||||
"""HEX → ASS 颜色 &HAABBGGRR(默认不透明)."""
|
||||
hex_color = hex_color.lstrip("#")
|
||||
if len(hex_color) != 6:
|
||||
return "&H00FFFFFF"
|
||||
r, g, b = hex_color[0:2], hex_color[2:4], hex_color[4:6]
|
||||
return f"&H00{b.upper()}{g.upper()}{r.upper()}"
|
||||
|
||||
|
||||
def hex_to_ass_bgr(hex_color: str) -> str:
|
||||
"""HEX → ASS BGR 部分(不含 alpha)."""
|
||||
hex_color = hex_color.lstrip("#")
|
||||
if len(hex_color) != 6:
|
||||
return "FFFFFF"
|
||||
r, g, b = hex_color[0:2], hex_color[2:4], hex_color[4:6]
|
||||
return f"{b.upper()}{g.upper()}{r.upper()}"
|
||||
|
||||
|
||||
def opacity_to_ass_alpha(opacity: float) -> str:
|
||||
"""不透明度 → ASS alpha(00=不透明,FF=完全透明)."""
|
||||
alpha = 255 - int(max(0.0, min(1.0, opacity)) * 255)
|
||||
return f"{alpha:02X}"
|
||||
|
||||
|
||||
def escape_ass_text(text: str) -> str:
|
||||
"""转义 ASS 文本特殊字符."""
|
||||
text = text.replace("\r\n", "\\N").replace("\n", "\\N").replace("\r", "\\N")
|
||||
text = text.replace("{", "(").replace("}", ")")
|
||||
return text
|
||||
|
||||
|
||||
def format_ass_time(seconds: float) -> str:
|
||||
"""秒 → ASS 时间格式 H:MM:SS.cc."""
|
||||
if seconds < 0:
|
||||
seconds = 0.0
|
||||
hours = int(seconds // 3600)
|
||||
minutes = int((seconds % 3600) // 60)
|
||||
secs = seconds % 60
|
||||
return f"{hours}:{minutes:02d}:{secs:05.2f}"
|
||||
|
||||
|
||||
def wrap_text(text: str, max_chars: int) -> list[str]:
|
||||
"""按字数换行,优先标点断开."""
|
||||
if max_chars <= 0:
|
||||
return [text]
|
||||
if not text or len(text) <= max_chars:
|
||||
return [text]
|
||||
|
||||
lines: list[str] = []
|
||||
remaining = text
|
||||
punctuations = ",。!?、;:,.;:!?"
|
||||
|
||||
while len(remaining) > max_chars:
|
||||
break_point = max_chars
|
||||
# 在 max_chars 到 max_chars//2 之间寻找标点断点
|
||||
for i in range(max_chars, max_chars // 2, -1):
|
||||
if i < len(remaining) and remaining[i] in punctuations:
|
||||
break_point = i + 1
|
||||
break
|
||||
|
||||
lines.append(remaining[:break_point])
|
||||
remaining = remaining[break_point:]
|
||||
|
||||
if remaining:
|
||||
lines.append(remaining)
|
||||
|
||||
return lines
|
||||
|
||||
|
||||
# ── 字幕样式配置 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class SubtitleStyle:
|
||||
"""字幕样式配置."""
|
||||
|
||||
font_name: str = DEFAULT_FONT
|
||||
font_size: int = DEFAULT_FONT_SIZE
|
||||
font_color: str = DEFAULT_COLOR
|
||||
bold: bool = False
|
||||
italic: bool = False
|
||||
|
||||
# 描边
|
||||
stroke_enabled: bool = True
|
||||
stroke_color: str = DEFAULT_STROKE_COLOR
|
||||
stroke_width: float = DEFAULT_STROKE_WIDTH
|
||||
|
||||
# 阴影
|
||||
shadow_enabled: bool = False
|
||||
shadow_color: str = "#000000"
|
||||
shadow_offset_x: int = 2
|
||||
shadow_offset_y: int = 2
|
||||
shadow_blur: float = 0.0
|
||||
|
||||
# 背景框
|
||||
background_enabled: bool = False
|
||||
background_color: str = "#000000"
|
||||
background_opacity: float = 0.5
|
||||
background_padding: int = 8
|
||||
background_radius: int = 4
|
||||
|
||||
# 位置
|
||||
position: str = DEFAULT_POSITION
|
||||
margin_v: int = 60
|
||||
margin_l: int = 40
|
||||
margin_r: int = 40
|
||||
|
||||
# 多行
|
||||
max_chars_per_line: int = DEFAULT_MAX_CHARS_PER_LINE
|
||||
line_spacing: int = 0
|
||||
|
||||
# 动画
|
||||
fade_in: float = 0.0
|
||||
fade_out: float = 0.0
|
||||
animation_type: str = "none"
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config: dict[str, Any] | None) -> "SubtitleStyle":
|
||||
"""从字典创建样式配置,带安全类型转换."""
|
||||
if not config or not isinstance(config, dict):
|
||||
return cls()
|
||||
|
||||
def safe_str(key: str, default: str) -> str:
|
||||
val = config.get(key, default)
|
||||
return str(val) if val is not None else default
|
||||
|
||||
def safe_int(key: str, default: int) -> int:
|
||||
try:
|
||||
return int(config.get(key, default))
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
def safe_float(key: str, default: float) -> float:
|
||||
try:
|
||||
return float(config.get(key, default))
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
def safe_bool(key: str, default: bool) -> bool:
|
||||
return bool(config.get(key, default))
|
||||
|
||||
position = safe_str("position", DEFAULT_POSITION)
|
||||
position = POSITION_ALIASES.get(position, position)
|
||||
if position not in POSITION_ALIGNMENT:
|
||||
position = DEFAULT_POSITION
|
||||
|
||||
return cls(
|
||||
font_name=safe_str("font", DEFAULT_FONT),
|
||||
font_size=safe_int("size", DEFAULT_FONT_SIZE),
|
||||
font_color=safe_str("color", DEFAULT_COLOR),
|
||||
bold=safe_bool("bold", False),
|
||||
italic=safe_bool("italic", False),
|
||||
stroke_enabled=safe_bool("stroke_enabled", True),
|
||||
stroke_color=safe_str("stroke_color", DEFAULT_STROKE_COLOR),
|
||||
stroke_width=safe_float("stroke_width", DEFAULT_STROKE_WIDTH),
|
||||
shadow_enabled=safe_bool("shadow_enabled", False),
|
||||
shadow_color=safe_str("shadow_color", "#000000"),
|
||||
shadow_offset_x=safe_int("shadow_offset_x", 2),
|
||||
shadow_offset_y=safe_int("shadow_offset_y", 2),
|
||||
shadow_blur=safe_float("shadow_blur", 0.0),
|
||||
background_enabled=safe_bool("background_enabled", False),
|
||||
background_color=safe_str("background_color", "#000000"),
|
||||
background_opacity=max(0.0, min(1.0, safe_float("background_opacity", 0.5))),
|
||||
background_padding=safe_int("background_padding", 8),
|
||||
background_radius=safe_int("background_radius", 4),
|
||||
position=position,
|
||||
margin_v=safe_int("margin_v", 60),
|
||||
margin_l=safe_int("margin_l", 40),
|
||||
margin_r=safe_int("margin_r", 40),
|
||||
max_chars_per_line=safe_int("max_chars_per_line", DEFAULT_MAX_CHARS_PER_LINE),
|
||||
line_spacing=safe_int("line_spacing", 0),
|
||||
fade_in=max(0.0, safe_float("fade_in", 0.0)),
|
||||
fade_out=max(0.0, safe_float("fade_out", 0.0)),
|
||||
animation_type=safe_str("animation_type", "none"),
|
||||
)
|
||||
|
||||
@property
|
||||
def alignment(self) -> int:
|
||||
"""获取 ASS alignment 编号."""
|
||||
return POSITION_ALIGNMENT.get(self.position, 2)
|
||||
|
||||
@property
|
||||
def ass_font_color(self) -> str:
|
||||
"""ASS 格式颜色 &HAABBGGRR."""
|
||||
return hex_to_ass_color(self.font_color)
|
||||
|
||||
@property
|
||||
def ass_stroke_color(self) -> str:
|
||||
return hex_to_ass_color(self.stroke_color)
|
||||
|
||||
@property
|
||||
def ass_shadow_color(self) -> str:
|
||||
return hex_to_ass_color(self.shadow_color)
|
||||
|
||||
@property
|
||||
def ass_background_color(self) -> str:
|
||||
"""背景框颜色(ASS BackColour),带透明度."""
|
||||
alpha_hex = opacity_to_ass_alpha(self.background_opacity)
|
||||
color_bgr = hex_to_ass_bgr(self.background_color)
|
||||
return f"&H{alpha_hex}{color_bgr}"
|
||||
|
||||
|
||||
# ── 字幕片段 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class SubtitleSegment:
|
||||
"""单个字幕片段."""
|
||||
|
||||
start: float # 开始时间(秒)
|
||||
end: float # 结束时间(秒)
|
||||
text: str # 字幕文本
|
||||
style_name: str = "Default" # 使用的样式名
|
||||
|
||||
@property
|
||||
def duration(self) -> float:
|
||||
"""字幕时长."""
|
||||
return max(0.0, self.end - self.start)
|
||||
|
||||
@property
|
||||
def is_valid(self) -> bool:
|
||||
"""是否有效(有文本且时长>0)."""
|
||||
return bool(self.text) and self.end > self.start
|
||||
Executable
+293
@@ -0,0 +1,293 @@
|
||||
"""模板片段转换器 — 纯函数集合.
|
||||
|
||||
从 edit_template_service.py 抽离的纯逻辑,负责在不同数据形态间转换:
|
||||
- 剪辑计划片段 (EditPlanClip) → 模板片段配置 (TemplateClipConfig)
|
||||
- 模板片段配置 → 版本快照 dict
|
||||
- 版本快照 dict → 模板片段配置
|
||||
- 计划 config → 模板 config(过滤运行时字段)
|
||||
|
||||
所有函数均为纯函数,不依赖数据库或外部 IO。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from packages.domain.template_clip_config import (
|
||||
ClipType,
|
||||
TemplateClipConfig,
|
||||
TransitionEffect,
|
||||
)
|
||||
|
||||
# ── 安全枚举解析 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def safe_parse_transition_effect(value: Any, default: TransitionEffect = TransitionEffect.CUT) -> TransitionEffect:
|
||||
"""安全解析转场效果枚举,解析失败返回默认值。
|
||||
|
||||
Args:
|
||||
value: 待解析的值(枚举、字符串或其他)
|
||||
default: 解析失败时的默认值
|
||||
|
||||
Returns:
|
||||
TransitionEffect 枚举值
|
||||
"""
|
||||
if isinstance(value, TransitionEffect):
|
||||
return value
|
||||
try:
|
||||
return TransitionEffect(value)
|
||||
except (ValueError, TypeError):
|
||||
return default
|
||||
|
||||
|
||||
def safe_parse_clip_type(value: Any, default: ClipType = ClipType.MAIN) -> ClipType:
|
||||
"""安全解析片段类型枚举,解析失败返回默认值。
|
||||
|
||||
Args:
|
||||
value: 待解析的值(枚举、字符串或其他)
|
||||
default: 解析失败时的默认值
|
||||
|
||||
Returns:
|
||||
ClipType 枚举值
|
||||
"""
|
||||
if isinstance(value, ClipType):
|
||||
return value
|
||||
try:
|
||||
return ClipType(value)
|
||||
except (ValueError, TypeError):
|
||||
return default
|
||||
|
||||
|
||||
# ── Config 字段过滤 ─────────────────────────────────────────────────────────
|
||||
|
||||
# 默认需要从 clip config 中移除的素材/运行时字段
|
||||
_DEFAULT_CLIP_CONFIG_SKIP_KEYS = frozenset(
|
||||
{
|
||||
"asset_info",
|
||||
"source_asset_id",
|
||||
}
|
||||
)
|
||||
|
||||
# 默认需要从 plan config 中移除的运行时/实例字段
|
||||
_DEFAULT_PLAN_CONFIG_SKIP_KEYS = frozenset(
|
||||
{
|
||||
"is_template_draft",
|
||||
"asset_ids",
|
||||
"source_edit_plan_id",
|
||||
"generation_task_id",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def filter_clip_config(
|
||||
clip_config: dict[str, Any] | None,
|
||||
playback_speed: float | None = None,
|
||||
skip_keys: frozenset[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建模板片段的 config 字典。
|
||||
|
||||
处理逻辑:
|
||||
1. 如果 playback_speed 存在且不等于 1.0,加入 config
|
||||
2. 合并 clip 自身的 config
|
||||
3. 移除素材相关字段
|
||||
|
||||
Args:
|
||||
clip_config: 原始片段 config(可为 None)
|
||||
playback_speed: 播放速度(可选,1.0 时不写入)
|
||||
skip_keys: 需要跳过的字段集合(None 时用默认)
|
||||
|
||||
Returns:
|
||||
过滤后的 config 字典
|
||||
"""
|
||||
skip = skip_keys if skip_keys is not None else _DEFAULT_CLIP_CONFIG_SKIP_KEYS
|
||||
result: dict[str, Any] = {}
|
||||
|
||||
if playback_speed is not None and playback_speed != 1.0:
|
||||
result["playback_speed"] = playback_speed
|
||||
|
||||
if clip_config:
|
||||
result.update(clip_config)
|
||||
|
||||
for key in skip:
|
||||
result.pop(key, None)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def filter_plan_config_to_template(
|
||||
plan_config: dict[str, Any] | None,
|
||||
skip_keys: frozenset[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""从计划 config 中提取模板 config(过滤运行时/实例字段)。
|
||||
|
||||
Args:
|
||||
plan_config: 原始计划 config(可为 None)
|
||||
skip_keys: 需要跳过的字段集合(None 时用默认)
|
||||
|
||||
Returns:
|
||||
过滤后的模板 config
|
||||
"""
|
||||
skip = skip_keys if skip_keys is not None else _DEFAULT_PLAN_CONFIG_SKIP_KEYS
|
||||
if not plan_config:
|
||||
return {}
|
||||
return {k: v for k, v in plan_config.items() if k not in skip}
|
||||
|
||||
|
||||
# ── Clip → TemplateClipConfig 转换 ────────────────────────────────────────
|
||||
|
||||
|
||||
def clip_to_template_clip_config(
|
||||
template_id: str,
|
||||
clip: Any,
|
||||
) -> TemplateClipConfig:
|
||||
"""将剪辑计划片段转换为模板片段配置。
|
||||
|
||||
转换规则:
|
||||
- clip_type → 安全解析后映射
|
||||
- order → 保持不变
|
||||
- duration → min_duration = max_duration = duration(固定时长)
|
||||
- text_content → text_template
|
||||
- transition_effect → 安全解析后映射
|
||||
- playback_speed → 存入 config(非 1.0 时)
|
||||
- clip.config → 合并入 config(过滤素材字段)
|
||||
|
||||
Args:
|
||||
template_id: 目标模板 ID
|
||||
clip: 源片段对象(需有 clip_type/order/duration/text_content/
|
||||
transition_effect/playback_speed/config 属性)
|
||||
|
||||
Returns:
|
||||
新创建的 TemplateClipConfig 实例
|
||||
"""
|
||||
clip_type = safe_parse_clip_type(getattr(clip, "clip_type", None))
|
||||
transition = safe_parse_transition_effect(getattr(clip, "transition_effect", None))
|
||||
|
||||
config = filter_clip_config(
|
||||
getattr(clip, "config", None),
|
||||
playback_speed=getattr(clip, "playback_speed", None),
|
||||
)
|
||||
|
||||
duration = getattr(clip, "duration", 0.0) or 0.0
|
||||
|
||||
return TemplateClipConfig.create(
|
||||
template_id=template_id,
|
||||
clip_type=clip_type,
|
||||
order=getattr(clip, "order", 0),
|
||||
min_duration=duration,
|
||||
max_duration=duration,
|
||||
text_template=getattr(clip, "text_content", "") or "",
|
||||
transition_effect=transition,
|
||||
config=config,
|
||||
)
|
||||
|
||||
|
||||
def clips_to_template_clip_configs(
|
||||
template_id: str,
|
||||
clips: list[Any],
|
||||
) -> list[TemplateClipConfig]:
|
||||
"""批量将剪辑计划片段转换为模板片段配置列表。
|
||||
|
||||
Args:
|
||||
template_id: 目标模板 ID
|
||||
clips: 源片段对象列表
|
||||
|
||||
Returns:
|
||||
TemplateClipConfig 实例列表
|
||||
"""
|
||||
return [clip_to_template_clip_config(template_id, c) for c in clips]
|
||||
|
||||
|
||||
# ── TemplateClipConfig → Snapshot 转换 ────────────────────────────────────
|
||||
|
||||
|
||||
def _enum_value(value: Any) -> Any:
|
||||
"""获取枚举的 value 值(兼容枚举和字符串)。"""
|
||||
if hasattr(value, "value"):
|
||||
return value.value
|
||||
return value
|
||||
|
||||
|
||||
def clip_config_to_snapshot(cfg: Any) -> dict[str, Any]:
|
||||
"""将模板片段配置转换为版本快照 dict。
|
||||
|
||||
Args:
|
||||
cfg: TemplateClipConfig 对象(或有对应属性的对象)
|
||||
|
||||
Returns:
|
||||
快照字典,包含 clip_type/order/min_duration/max_duration/
|
||||
text_template/transition_effect/config
|
||||
"""
|
||||
return {
|
||||
"clip_type": _enum_value(getattr(cfg, "clip_type", None)),
|
||||
"order": getattr(cfg, "order", 0),
|
||||
"min_duration": getattr(cfg, "min_duration", 0.0),
|
||||
"max_duration": getattr(cfg, "max_duration", 0.0),
|
||||
"text_template": getattr(cfg, "text_template", "") or "",
|
||||
"transition_effect": _enum_value(getattr(cfg, "transition_effect", None)),
|
||||
"config": dict(getattr(cfg, "config", {}) or {}),
|
||||
}
|
||||
|
||||
|
||||
def clip_configs_to_snapshots(configs: list[Any]) -> list[dict[str, Any]]:
|
||||
"""批量将模板片段配置转换为版本快照列表。"""
|
||||
return [clip_config_to_snapshot(c) for c in configs]
|
||||
|
||||
|
||||
# ── Snapshot → TemplateClipConfig 转换 ────────────────────────────────────
|
||||
|
||||
|
||||
def snapshot_to_template_clip_config(
|
||||
template_id: str,
|
||||
snapshot: dict[str, Any],
|
||||
) -> TemplateClipConfig:
|
||||
"""将版本快照 dict 转换为模板片段配置。
|
||||
|
||||
Args:
|
||||
template_id: 目标模板 ID
|
||||
snapshot: 快照字典
|
||||
|
||||
Returns:
|
||||
新创建的 TemplateClipConfig 实例
|
||||
"""
|
||||
clip_type = safe_parse_clip_type(snapshot.get("clip_type", "main"))
|
||||
transition = safe_parse_transition_effect(snapshot.get("transition_effect", "cut"))
|
||||
|
||||
return TemplateClipConfig.create(
|
||||
template_id=template_id,
|
||||
clip_type=clip_type,
|
||||
order=snapshot.get("order", 0),
|
||||
min_duration=snapshot.get("min_duration", 0.0),
|
||||
max_duration=snapshot.get("max_duration", 0.0),
|
||||
text_template=snapshot.get("text_template", ""),
|
||||
transition_effect=transition,
|
||||
config=dict(snapshot.get("config", {}) or {}),
|
||||
)
|
||||
|
||||
|
||||
def snapshots_to_template_clip_configs(
|
||||
template_id: str,
|
||||
snapshots: list[dict[str, Any]],
|
||||
) -> list[TemplateClipConfig]:
|
||||
"""批量将版本快照转换为模板片段配置列表。"""
|
||||
return [snapshot_to_template_clip_config(template_id, s) for s in snapshots]
|
||||
|
||||
|
||||
# ── 名称校验工具 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def validate_template_name(name: str | None) -> str:
|
||||
"""校验并清洗模板名称。
|
||||
|
||||
Args:
|
||||
name: 原始名称
|
||||
|
||||
Returns:
|
||||
清洗后的名称(去除首尾空格)
|
||||
|
||||
Raises:
|
||||
ValueError: 名称为空
|
||||
"""
|
||||
clean_name = name.strip() if name else ""
|
||||
if not clean_name:
|
||||
raise ValueError("模板名称不能为空")
|
||||
return clean_name
|
||||
Executable
+177
@@ -0,0 +1,177 @@
|
||||
"""视频拼接领域模型 — 纯逻辑,无外部依赖.
|
||||
|
||||
抽离自 concat_engine.py 的数据类和配置解析逻辑,
|
||||
方便单测覆盖,同时保持向后兼容。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
MAX_CONCAT_SEGMENTS = 50 # 最大拼接段数(安全上限,防止OOM)
|
||||
|
||||
ALLOWED_VIDEO_EXTENSIONS = {".mp4", ".mov", ".avi", ".mkv", ".webm", ".flv", ".wmv"}
|
||||
|
||||
# concat demuxer 要求一致的参数列表
|
||||
CONCAT_DEMUXER_REQUIRED_PARAMS = [
|
||||
"codec_name",
|
||||
"width",
|
||||
"height",
|
||||
"r_frame_rate",
|
||||
"pix_fmt",
|
||||
"sample_rate",
|
||||
"channels",
|
||||
"audio_codec",
|
||||
]
|
||||
|
||||
|
||||
# ── 拼接片段配置 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConcatSegment:
|
||||
"""单个拼接片段."""
|
||||
|
||||
video_path: str # 视频文件路径
|
||||
start_time: float = 0.0 # 开始时间(秒)
|
||||
duration: float = 0.0 # 持续时长(秒),0表示取到末尾
|
||||
has_audio: bool = True # 是否包含音频
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, seg: dict[str, Any] | None) -> "ConcatSegment":
|
||||
"""从字典创建拼接片段,带安全类型转换."""
|
||||
if not seg or not isinstance(seg, dict):
|
||||
return cls(video_path="")
|
||||
|
||||
try:
|
||||
start_time = max(0.0, float(seg.get("start_time", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
start_time = 0.0
|
||||
|
||||
try:
|
||||
duration = max(0.0, float(seg.get("duration", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
duration = 0.0
|
||||
|
||||
return cls(
|
||||
video_path=str(seg.get("video_path", "")),
|
||||
start_time=start_time,
|
||||
duration=duration,
|
||||
has_audio=bool(seg.get("has_audio", True)),
|
||||
)
|
||||
|
||||
@property
|
||||
def is_valid(self) -> bool:
|
||||
"""是否为有效片段(有视频路径)."""
|
||||
return bool(self.video_path)
|
||||
|
||||
@property
|
||||
def effective_duration(self) -> float:
|
||||
"""有效时长(duration > 0 时取 duration,否则 0)."""
|
||||
return max(0.0, self.duration)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConcatConfig:
|
||||
"""视频拼接配置."""
|
||||
|
||||
segments: list[ConcatSegment] = field(default_factory=list)
|
||||
output_width: int = 0 # 输出宽度(0=自动取第一段)
|
||||
output_height: int = 0 # 输出高度(0=自动取第一段)
|
||||
output_fps: float = 0.0 # 输出帧率(0=自动取第一段)
|
||||
force_reencode: bool = False # 强制重新编码
|
||||
transition: str = "none" # 转场效果(none/crossfade)
|
||||
transition_duration: float = 0.3 # 转场时长
|
||||
|
||||
@classmethod
|
||||
def from_config_dict(cls, config: dict[str, Any] | None) -> "ConcatConfig":
|
||||
"""从配置字典创建 ConcatConfig."""
|
||||
if not config or not isinstance(config, dict):
|
||||
return cls()
|
||||
|
||||
segments_raw = config.get("segments", [])
|
||||
segments: list[ConcatSegment] = []
|
||||
|
||||
if isinstance(segments_raw, list):
|
||||
for s in segments_raw:
|
||||
if isinstance(s, dict) and s.get("video_path"):
|
||||
try:
|
||||
seg = ConcatSegment.from_dict(s)
|
||||
if seg.is_valid:
|
||||
segments.append(seg)
|
||||
except Exception:
|
||||
logger.warning("[concat] skip invalid segment: %s", s)
|
||||
continue
|
||||
|
||||
try:
|
||||
output_width = max(0, int(config.get("output_width", 0)))
|
||||
except (TypeError, ValueError):
|
||||
output_width = 0
|
||||
|
||||
try:
|
||||
output_height = max(0, int(config.get("output_height", 0)))
|
||||
except (TypeError, ValueError):
|
||||
output_height = 0
|
||||
|
||||
try:
|
||||
output_fps = max(0.0, float(config.get("output_fps", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
output_fps = 0.0
|
||||
|
||||
try:
|
||||
transition_duration = max(0.1, float(config.get("transition_duration", 0.3)))
|
||||
except (TypeError, ValueError):
|
||||
transition_duration = 0.3
|
||||
|
||||
return cls(
|
||||
segments=segments,
|
||||
output_width=output_width,
|
||||
output_height=output_height,
|
||||
output_fps=output_fps,
|
||||
force_reencode=bool(config.get("force_reencode", False)),
|
||||
transition=str(config.get("transition", "none")),
|
||||
transition_duration=transition_duration,
|
||||
)
|
||||
|
||||
@property
|
||||
def has_effect(self) -> bool:
|
||||
"""是否有有效片段需要拼接(至少2段)."""
|
||||
return self.valid_segment_count >= 2
|
||||
|
||||
@property
|
||||
def valid_segment_count(self) -> int:
|
||||
"""有效片段数量."""
|
||||
return sum(1 for s in self.segments if s.is_valid)
|
||||
|
||||
@property
|
||||
def total_segments(self) -> int:
|
||||
"""有效片段数量(向后兼容别名)."""
|
||||
return self.valid_segment_count
|
||||
|
||||
@property
|
||||
def first_valid_segment(self) -> ConcatSegment | None:
|
||||
"""第一个有效片段."""
|
||||
for s in self.segments:
|
||||
if s.is_valid:
|
||||
return s
|
||||
return None
|
||||
|
||||
@property
|
||||
def estimated_total_duration(self) -> float:
|
||||
"""估算总时长(只统计有明确duration的片段)."""
|
||||
total = 0.0
|
||||
for s in self.segments:
|
||||
if s.is_valid and s.duration > 0:
|
||||
total += s.duration
|
||||
return total
|
||||
|
||||
def clamp_segments(self, max_segments: int = MAX_CONCAT_SEGMENTS) -> None:
|
||||
"""截断片段数量,防止OOM."""
|
||||
if len(self.segments) > max_segments:
|
||||
self.segments = self.segments[:max_segments]
|
||||
Executable
+375
@@ -0,0 +1,375 @@
|
||||
"""视频滤镜构建器 — FFmpeg filter_complex 纯逻辑层。
|
||||
|
||||
从 video_compose_service.py 抽离的纯函数集合,专门负责 FFmpeg 滤镜链的构建,
|
||||
不依赖数据库、不做 IO,便于单元测试。
|
||||
|
||||
主要职责:
|
||||
- 单片段滤镜链构建(scale / pad / format / fps / setpts / trim)
|
||||
- concat 滤镜构建(无转场高效拼接)
|
||||
- xfade 转场滤镜链构建(fade / slide / dissolve / wipe)
|
||||
- 完整 filter_complex 策略选择与组装
|
||||
- 音频流判断与音频滤镜归一化
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from packages.domain.edit_plan_clip import EditPlanClip
|
||||
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
DEFAULT_OUTPUT_WIDTH = 1280
|
||||
DEFAULT_OUTPUT_HEIGHT = 720
|
||||
DEFAULT_FPS = 25
|
||||
|
||||
# xfade 转场映射:TransitionEffect → FFmpeg xfade transition 名称
|
||||
XFADE_TRANSITION_MAP: dict[str, str] = {
|
||||
TransitionEffect.FADE: "fade",
|
||||
TransitionEffect.SLIDE_LEFT: "slideleft",
|
||||
TransitionEffect.SLIDE_RIGHT: "slideright",
|
||||
TransitionEffect.DISSOLVE: "dissolve",
|
||||
TransitionEffect.WIPE: "wipeleft",
|
||||
}
|
||||
|
||||
# 转场默认时长(秒)
|
||||
DEFAULT_TRANSITION_DURATION = 0.5
|
||||
|
||||
# 默认片段时长(当 clip.duration <= 0 时使用)
|
||||
DEFAULT_CLIP_DURATION = 5.0
|
||||
|
||||
|
||||
# ── 数据结构 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ClipFilterChain:
|
||||
"""单个片段的滤镜链描述。"""
|
||||
|
||||
clip_id: str
|
||||
input_index: int
|
||||
video_label: str
|
||||
audio_label: str | None
|
||||
filters: list[str]
|
||||
duration: float
|
||||
|
||||
|
||||
# ── 单片段滤镜链 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def build_clip_filter(
|
||||
clip: "EditPlanClip",
|
||||
input_index: int,
|
||||
output_width: int,
|
||||
output_height: int,
|
||||
fps: int,
|
||||
) -> ClipFilterChain:
|
||||
"""为单个片段构建滤镜链。
|
||||
|
||||
滤镜顺序:
|
||||
1. scale — 等比缩放到目标分辨率(保证覆盖,不裁剪内容)
|
||||
2. pad — 居中+留黑边到目标分辨率(保持原始比例)
|
||||
3. format — 统一像素格式为 yuv420p(concat 要求像素格式一致)
|
||||
4. fps — 统一帧率(concat 要求所有输入帧率一致)
|
||||
5. setpts — 重置时间戳 + 起始偏移
|
||||
6. trim — 视频时长裁剪 + 重置 PTS
|
||||
|
||||
Args:
|
||||
clip: 剪辑计划片段
|
||||
input_index: 输入流索引(对应第几个 -i)
|
||||
output_width: 输出宽度(像素)
|
||||
output_height: 输出高度(像素)
|
||||
fps: 输出帧率
|
||||
|
||||
Returns:
|
||||
ClipFilterChain 描述对象
|
||||
"""
|
||||
duration = clip.duration if clip.duration > 0 else DEFAULT_CLIP_DURATION
|
||||
start = clip.start_time
|
||||
|
||||
filters: list[str] = []
|
||||
|
||||
# 1. scale: 等比缩放(保持比例,不裁剪)
|
||||
filters.append(f"scale={output_width}:{output_height}" f":force_original_aspect_ratio=decrease")
|
||||
|
||||
# 2. pad: 居中+留黑边到目标分辨率
|
||||
filters.append(f"pad={output_width}:{output_height}:(ow-iw)/2:(oh-ih)/2:black")
|
||||
|
||||
# 3. format: 统一像素格式为 yuv420p
|
||||
filters.append("format=yuv420p")
|
||||
|
||||
# 4. fps: 统一帧率
|
||||
if fps and fps > 0:
|
||||
filters.append(f"fps={fps}")
|
||||
|
||||
# 5. setpts: 重置时间戳 + 偏移
|
||||
if start > 0:
|
||||
filters.append(f"setpts=PTS-STARTPTS+{start}/TB")
|
||||
else:
|
||||
filters.append("setpts=PTS-STARTPTS")
|
||||
|
||||
# 6. trim: 视频时长 + 重置 PTS
|
||||
filters.append(f"trim=0:{duration}")
|
||||
filters.append("setpts=PTS-STARTPTS")
|
||||
|
||||
video_label = f"v{input_index}"
|
||||
|
||||
# 音频标签:title/subtitle 是纯文字/图片卡片,没有音频流
|
||||
clip_type = clip.clip_type.lower() if clip.clip_type else ""
|
||||
has_audio_stream = clip_type not in ("title", "subtitle")
|
||||
audio_label = f"a{input_index}" if has_audio_stream else None
|
||||
|
||||
return ClipFilterChain(
|
||||
clip_id=clip.id,
|
||||
input_index=input_index,
|
||||
video_label=video_label,
|
||||
audio_label=audio_label,
|
||||
filters=filters,
|
||||
duration=duration,
|
||||
)
|
||||
|
||||
|
||||
# ── 滤镜串联工具 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def chain_filters(filters: list[str], output_label: str, input_label: str = "0:v") -> str:
|
||||
"""将滤镜列表串联为 FFmpeg 滤镜字符串。
|
||||
|
||||
Args:
|
||||
filters: 滤镜表达式列表
|
||||
output_label: 输出标签名(不含方括号)
|
||||
input_label: 输入标签(默认 "0:v")
|
||||
|
||||
Returns:
|
||||
形如 "[0:v]scale=1280:720,fps=25[v0]" 的字符串
|
||||
"""
|
||||
filter_body = ",".join(filters)
|
||||
return f"[{input_label}]{filter_body}[{output_label}]"
|
||||
|
||||
|
||||
# ── 音频判断 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def has_audio(clip_chains: list[ClipFilterChain]) -> bool:
|
||||
"""是否有任何片段包含音频流。"""
|
||||
return any(c.audio_label is not None for c in clip_chains)
|
||||
|
||||
|
||||
# ── concat 滤镜 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def build_concat_filter(
|
||||
clip_chains: list[ClipFilterChain],
|
||||
) -> tuple[str, float]:
|
||||
"""构建 concat 滤镜(无转场,高效拼接)。
|
||||
|
||||
视频:每个片段先应用各自滤镜链,再用 concat 滤镜拼接
|
||||
音频:先 aformat 归一化(48000Hz/stereo/fltp)再 concat,
|
||||
避免不同采样率/声道导致 concat 失败
|
||||
|
||||
Args:
|
||||
clip_chains: 各片段的滤镜链描述
|
||||
|
||||
Returns:
|
||||
(filter_complex_string, estimated_total_duration)
|
||||
"""
|
||||
n = len(clip_chains)
|
||||
if n == 0:
|
||||
return "", 0.0
|
||||
|
||||
parts: list[str] = []
|
||||
total_duration = 0.0
|
||||
|
||||
# 每个片段的视频滤镜链
|
||||
for idx, chain in enumerate(clip_chains):
|
||||
filter_body = ",".join(chain.filters)
|
||||
parts.append(f"[{idx}:v]{filter_body}[{chain.video_label}]")
|
||||
total_duration += chain.duration
|
||||
|
||||
# 视频 concat 滤镜
|
||||
concat_inputs = "".join(f"[{c.video_label}]" for c in clip_chains)
|
||||
parts.append(f"{concat_inputs}concat=n={n}:v=1:a=0[outv]")
|
||||
|
||||
# 音频:归一化 + concat
|
||||
_append_audio_concat(parts, clip_chains)
|
||||
|
||||
return ";".join(parts), total_duration
|
||||
|
||||
|
||||
# ── xfade 转场滤镜 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def build_xfade_filter(
|
||||
clip_chains: list[ClipFilterChain],
|
||||
transition_duration: float,
|
||||
transitions: list[str],
|
||||
) -> tuple[str, float]:
|
||||
"""构建 xfade 转场滤镜链。
|
||||
|
||||
每两个相邻片段之间插入 xfade 转场。
|
||||
offset = 前一个片段的累积时长 - 转场时长。
|
||||
|
||||
视频转场支持:fade / slideleft / slideright / dissolve / wipeleft
|
||||
|
||||
Args:
|
||||
clip_chains: 各片段的滤镜链描述
|
||||
transition_duration: 转场时长(秒)
|
||||
transitions: 每个片段对应的转场效果列表(索引对应片段)
|
||||
|
||||
Returns:
|
||||
(filter_complex_string, estimated_total_duration)
|
||||
"""
|
||||
n = len(clip_chains)
|
||||
if n == 0:
|
||||
return "", 0.0
|
||||
|
||||
parts: list[str] = []
|
||||
total_duration = 0.0
|
||||
|
||||
# 每个片段的视频滤镜链
|
||||
for idx, chain in enumerate(clip_chains):
|
||||
filter_body = ",".join(chain.filters)
|
||||
parts.append(f"[{idx}:v]{filter_body}[{chain.video_label}]")
|
||||
total_duration += chain.duration
|
||||
|
||||
# 单片段:直接 copy 输出(无音频,与原实现保持一致)
|
||||
if n == 1:
|
||||
parts.append(f"[{clip_chains[0].video_label}]copy[outv]")
|
||||
return ";".join(parts), total_duration
|
||||
|
||||
# xfade 链式转场
|
||||
cumulative = 0.0
|
||||
prev_label = clip_chains[0].video_label
|
||||
|
||||
for i in range(1, n):
|
||||
cumulative += clip_chains[i - 1].duration
|
||||
offset = max(0.0, cumulative - transition_duration * i)
|
||||
|
||||
# 获取转场类型
|
||||
transition = transitions[i] if i < len(transitions) else "cut"
|
||||
xfade_transition = XFADE_TRANSITION_MAP.get(transition, "fade")
|
||||
|
||||
out_label = "outv" if i == n - 1 else f"xf{i}"
|
||||
|
||||
parts.append(
|
||||
f"[{prev_label}][{clip_chains[i].video_label}]"
|
||||
f"xfade=transition={xfade_transition}"
|
||||
f":duration={transition_duration}"
|
||||
f":offset={offset:.3f}"
|
||||
f"[{out_label}]"
|
||||
)
|
||||
prev_label = out_label
|
||||
|
||||
# 总时长减去转场重叠部分
|
||||
total_duration -= transition_duration * (n - 1)
|
||||
total_duration = max(0.0, total_duration)
|
||||
|
||||
# 音频:xfade 路径下的音频处理
|
||||
# 注意:使用 chain.audio_label 作为输入标签(与原实现保持一致)
|
||||
audio_chains = [c for c in clip_chains if c.audio_label]
|
||||
if len(audio_chains) >= 2:
|
||||
normalized_labels: list[str] = []
|
||||
for chain in audio_chains:
|
||||
norm_label = f"anorm_{chain.video_label}"
|
||||
audio_filters = [
|
||||
"aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp",
|
||||
f"atrim=0:{chain.duration}",
|
||||
"asetpts=PTS-STARTPTS",
|
||||
]
|
||||
parts.append(f"[{chain.audio_label}]{','.join(audio_filters)}[{norm_label}]")
|
||||
normalized_labels.append(norm_label)
|
||||
audio_inputs = "".join(f"[{label}]" for label in normalized_labels)
|
||||
parts.append(f"{audio_inputs}concat=n={len(normalized_labels)}:v=0:a=1[outa]")
|
||||
elif len(audio_chains) == 1:
|
||||
parts.append(f"[{audio_chains[0].audio_label}]acopy[outa]")
|
||||
|
||||
return ";".join(parts), total_duration
|
||||
|
||||
|
||||
# ── 完整 filter_complex 构建(策略选择) ────────────────────────────────────
|
||||
|
||||
|
||||
def build_filter_complex(
|
||||
clip_chains: list[ClipFilterChain],
|
||||
output_width: int,
|
||||
output_height: int,
|
||||
transition_duration: float,
|
||||
transitions: list[str],
|
||||
) -> tuple[str, float]:
|
||||
"""构建完整的 filter_complex 字符串(策略自动选择)。
|
||||
|
||||
策略:
|
||||
- 空列表:返回空字符串 + 0 时长
|
||||
- 单片段:直接输出(scale+pad+fps+trim 单链)
|
||||
- 多片段 + 全 cut:使用 concat 滤镜(高效)
|
||||
- 多片段 + 有转场:使用 xfade 滤镜链
|
||||
|
||||
Args:
|
||||
clip_chains: 各片段的滤镜链描述
|
||||
output_width: 输出宽度(目前单片段策略不使用,保留参数一致性)
|
||||
output_height: 输出高度(同上)
|
||||
transition_duration: 转场时长(秒)
|
||||
transitions: 每个片段对应的转场效果列表
|
||||
|
||||
Returns:
|
||||
(filter_complex_string, estimated_total_duration)
|
||||
"""
|
||||
n = len(clip_chains)
|
||||
|
||||
if n == 0:
|
||||
return "", 0.0
|
||||
|
||||
# 单片段
|
||||
if n == 1:
|
||||
chain = clip_chains[0]
|
||||
filter_str = chain_filters(chain.filters, chain.video_label)
|
||||
# 音频直通
|
||||
if chain.audio_label:
|
||||
filter_str += f";[0:a]{chain.audio_label}"
|
||||
total_duration = chain.duration
|
||||
return filter_str, total_duration
|
||||
|
||||
# 检查是否有转场
|
||||
has_transitions = any(t != TransitionEffect.CUT and t != "cut" for t in transitions)
|
||||
|
||||
if not has_transitions:
|
||||
return build_concat_filter(clip_chains)
|
||||
|
||||
# 有转场:使用 xfade
|
||||
return build_xfade_filter(
|
||||
clip_chains=clip_chains,
|
||||
transition_duration=transition_duration,
|
||||
transitions=transitions,
|
||||
)
|
||||
|
||||
|
||||
# ── 内部辅助:音频处理 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _append_audio_concat(parts: list[str], clip_chains: list[ClipFilterChain]) -> None:
|
||||
"""追加音频归一化 + concat 滤镜链到 parts(concat 路径)。
|
||||
|
||||
与原实现保持一致:归一化输出标签复用 chain.audio_label,
|
||||
concat 直接使用 audio_label 作为输入。
|
||||
"""
|
||||
audio_chains = [c for c in clip_chains if c.audio_label]
|
||||
if not audio_chains:
|
||||
return
|
||||
|
||||
# 先 aformat 归一化,输出到 chain.audio_label
|
||||
for chain in audio_chains:
|
||||
audio_filters = [
|
||||
"aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp",
|
||||
f"atrim=0:{chain.duration}",
|
||||
"asetpts=PTS-STARTPTS",
|
||||
]
|
||||
parts.append(f"[{chain.input_index}:a]{','.join(audio_filters)}[{chain.audio_label}]")
|
||||
|
||||
# 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]")
|
||||
Executable
+499
@@ -0,0 +1,499 @@
|
||||
"""Application 层零测试模块合集 — 第100波里程碑。
|
||||
|
||||
覆盖:
|
||||
- packages/application/generated_videos.py (8个UseCase)
|
||||
- packages/application/assets.py (ListAssets + CreateAsset)
|
||||
- packages/application/asset_libraries.py (ListLibraries + CreateLibrary)
|
||||
|
||||
策略: Mock repository,测参数校验 + 委托行为
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.asset_libraries import (
|
||||
CreateAssetLibraryCommand,
|
||||
CreateAssetLibraryUseCase,
|
||||
ListAssetLibrariesUseCase,
|
||||
)
|
||||
from packages.application.assets import (
|
||||
CreateAssetCommand,
|
||||
CreateAssetUseCase,
|
||||
ListAssetsUseCase,
|
||||
)
|
||||
from packages.application.generated_videos import (
|
||||
GetGeneratedVideoDownloadUrlUseCase,
|
||||
GetGeneratedVideoUseCase,
|
||||
GetVideosByIdsUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
ListGeneratedVideosPaginatedUseCase,
|
||||
ListGeneratedVideosUseCase,
|
||||
UpdateVideoReviewStatusUseCase,
|
||||
)
|
||||
from packages.domain import AssetLibraryKind, AssetStatus, ClassificationStatus, GeneratedVideo
|
||||
|
||||
# ── generated_videos.py ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListGeneratedVideosUseCase:
|
||||
def test_success(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_project.return_value = [MagicMock(spec=GeneratedVideo)]
|
||||
use_case = ListGeneratedVideosUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj1")
|
||||
|
||||
assert len(result) == 1
|
||||
mock_repo.list_by_project.assert_called_once_with("proj1")
|
||||
|
||||
def test_strips_project_id(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListGeneratedVideosUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" proj1 ")
|
||||
|
||||
mock_repo.list_by_project.assert_called_once_with("proj1")
|
||||
|
||||
def test_empty_project_id_raises(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListGeneratedVideosUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
def test_whitespace_project_id_raises(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListGeneratedVideosUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute(" \t ")
|
||||
|
||||
|
||||
class TestListGeneratedVideosPaginatedUseCase:
|
||||
def test_default_params(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_paginated.return_value = ([], 0)
|
||||
use_case = ListGeneratedVideosPaginatedUseCase(mock_repo)
|
||||
|
||||
result, total = use_case.execute()
|
||||
|
||||
assert total == 0
|
||||
assert result == []
|
||||
mock_repo.list_paginated.assert_called_once_with(
|
||||
user_id=None,
|
||||
project_id=None,
|
||||
status=None,
|
||||
review_status=None,
|
||||
page=1,
|
||||
page_size=20,
|
||||
)
|
||||
|
||||
def test_page_below_1_clamps_to_1(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_paginated.return_value = ([], 0)
|
||||
use_case = ListGeneratedVideosPaginatedUseCase(mock_repo)
|
||||
|
||||
use_case.execute(page=0)
|
||||
|
||||
mock_repo.list_paginated.assert_called_once()
|
||||
call_kwargs = mock_repo.list_paginated.call_args.kwargs
|
||||
assert call_kwargs["page"] == 1
|
||||
|
||||
def test_negative_page_clamps(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_paginated.return_value = ([], 0)
|
||||
use_case = ListGeneratedVideosPaginatedUseCase(mock_repo)
|
||||
|
||||
use_case.execute(page=-5)
|
||||
|
||||
assert mock_repo.list_paginated.call_args.kwargs["page"] == 1
|
||||
|
||||
def test_page_size_zero_clamps(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_paginated.return_value = ([], 0)
|
||||
use_case = ListGeneratedVideosPaginatedUseCase(mock_repo)
|
||||
|
||||
use_case.execute(page_size=0)
|
||||
|
||||
assert mock_repo.list_paginated.call_args.kwargs["page_size"] == 20
|
||||
|
||||
def test_page_size_over_100_clamps(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_paginated.return_value = ([], 0)
|
||||
use_case = ListGeneratedVideosPaginatedUseCase(mock_repo)
|
||||
|
||||
use_case.execute(page_size=200)
|
||||
|
||||
assert mock_repo.list_paginated.call_args.kwargs["page_size"] == 20
|
||||
|
||||
def test_page_size_50_ok(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_paginated.return_value = ([], 0)
|
||||
use_case = ListGeneratedVideosPaginatedUseCase(mock_repo)
|
||||
|
||||
use_case.execute(page_size=50)
|
||||
|
||||
assert mock_repo.list_paginated.call_args.kwargs["page_size"] == 50
|
||||
|
||||
def test_with_all_filters(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_paginated.return_value = ([], 0)
|
||||
use_case = ListGeneratedVideosPaginatedUseCase(mock_repo)
|
||||
|
||||
use_case.execute(
|
||||
user_id="u1",
|
||||
project_id="p1",
|
||||
status="completed",
|
||||
review_status="approved",
|
||||
page=2,
|
||||
page_size=10,
|
||||
)
|
||||
|
||||
mock_repo.list_paginated.assert_called_once_with(
|
||||
user_id="u1",
|
||||
project_id="p1",
|
||||
status="completed",
|
||||
review_status="approved",
|
||||
page=2,
|
||||
page_size=10,
|
||||
)
|
||||
|
||||
|
||||
class TestGetGeneratedVideoUseCase:
|
||||
def test_found(self):
|
||||
mock_repo = MagicMock()
|
||||
expected = MagicMock(spec=GeneratedVideo)
|
||||
mock_repo.get.return_value = expected
|
||||
use_case = GetGeneratedVideoUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vid1")
|
||||
|
||||
assert result == expected
|
||||
mock_repo.get.assert_called_once_with("vid1")
|
||||
|
||||
def test_not_found(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetGeneratedVideoUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vid1")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestListGeneratedVideosByTaskUseCase:
|
||||
def test_success(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_generation_task.return_value = [MagicMock()]
|
||||
use_case = ListGeneratedVideosByTaskUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("task1")
|
||||
|
||||
assert len(result) == 1
|
||||
mock_repo.list_by_generation_task.assert_called_once_with("task1")
|
||||
|
||||
def test_strips_task_id(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListGeneratedVideosByTaskUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" task1 ")
|
||||
|
||||
mock_repo.list_by_generation_task.assert_called_once_with("task1")
|
||||
|
||||
def test_empty_task_id_raises(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListGeneratedVideosByTaskUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="generation_task_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
|
||||
class TestGetGeneratedVideoDownloadUrlUseCase:
|
||||
def test_found(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_item = MagicMock()
|
||||
mock_item.file_url = "https://cdn/v.mp4"
|
||||
mock_repo.get.return_value = mock_item
|
||||
use_case = GetGeneratedVideoDownloadUrlUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vid1")
|
||||
|
||||
assert result == "https://cdn/v.mp4"
|
||||
|
||||
def test_not_found_returns_none(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetGeneratedVideoDownloadUrlUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vid1")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestUpdateVideoReviewStatusUseCase:
|
||||
def test_pending_review(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.update_review_status.return_value = MagicMock()
|
||||
use_case = UpdateVideoReviewStatusUseCase(mock_repo)
|
||||
|
||||
use_case.execute("vid1", "pending_review")
|
||||
|
||||
mock_repo.update_review_status.assert_called_once_with("vid1", "pending_review")
|
||||
|
||||
def test_approved(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = UpdateVideoReviewStatusUseCase(mock_repo)
|
||||
|
||||
use_case.execute("vid1", "approved")
|
||||
|
||||
mock_repo.update_review_status.assert_called_once_with("vid1", "approved")
|
||||
|
||||
def test_rejected(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = UpdateVideoReviewStatusUseCase(mock_repo)
|
||||
|
||||
use_case.execute("vid1", "rejected")
|
||||
|
||||
mock_repo.update_review_status.assert_called_once_with("vid1", "rejected")
|
||||
|
||||
def test_strips_video_id(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = UpdateVideoReviewStatusUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" vid1 ", "approved")
|
||||
|
||||
mock_repo.update_review_status.assert_called_once_with("vid1", "approved")
|
||||
|
||||
def test_empty_video_id_raises(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = UpdateVideoReviewStatusUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="video_id 不能为空"):
|
||||
use_case.execute("", "approved")
|
||||
|
||||
def test_invalid_status_raises(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = UpdateVideoReviewStatusUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="无效的 review_status"):
|
||||
use_case.execute("vid1", "invalid_status")
|
||||
|
||||
def test_not_found_returns_none(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.update_review_status.return_value = None
|
||||
use_case = UpdateVideoReviewStatusUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vid1", "approved")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestGetVideosByIdsUseCase:
|
||||
def test_success(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.get_by_ids.return_value = [MagicMock(), MagicMock()]
|
||||
use_case = GetVideosByIdsUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(["id1", "id2", "id3"])
|
||||
|
||||
assert len(result) == 2
|
||||
mock_repo.get_by_ids.assert_called_once_with(["id1", "id2", "id3"])
|
||||
|
||||
def test_empty_list(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.get_by_ids.return_value = []
|
||||
use_case = GetVideosByIdsUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute([])
|
||||
|
||||
assert result == []
|
||||
mock_repo.get_by_ids.assert_called_once_with([])
|
||||
|
||||
|
||||
# ── assets.py ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateAssetCommand:
|
||||
def test_minimal(self):
|
||||
cmd = CreateAssetCommand(
|
||||
project_id="p1",
|
||||
library_id="l1",
|
||||
name="test.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
assert cmd.project_id == "p1"
|
||||
assert cmd.library_id == "l1"
|
||||
assert cmd.name == "test.mp4"
|
||||
assert cmd.storage_key == "k"
|
||||
assert cmd.mime_type == "video/mp4"
|
||||
assert cmd.file_size == 0
|
||||
assert cmd.status == AssetStatus.UPLOADING
|
||||
assert cmd.classification_status == ClassificationStatus.PENDING
|
||||
|
||||
def test_full(self):
|
||||
cmd = CreateAssetCommand(
|
||||
project_id="p1",
|
||||
library_id="l1",
|
||||
name="test.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
metadata={"k": "v"},
|
||||
file_size=1024,
|
||||
duration=10.0,
|
||||
width=1920,
|
||||
height=1080,
|
||||
fps=30.0,
|
||||
codec="h264",
|
||||
status=AssetStatus.READY,
|
||||
quality_score=0.9,
|
||||
uploaded_by_user_id="u1",
|
||||
)
|
||||
assert cmd.file_size == 1024
|
||||
assert cmd.duration == 10.0
|
||||
assert cmd.status == AssetStatus.READY
|
||||
assert cmd.quality_score == 0.9
|
||||
|
||||
|
||||
class TestListAssetsUseCase:
|
||||
def test_success(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_library.return_value = []
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("lib1")
|
||||
|
||||
assert result == []
|
||||
mock_repo.find_by_library.assert_called_once_with("lib1")
|
||||
|
||||
def test_strips_library_id(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" lib1 ")
|
||||
|
||||
mock_repo.find_by_library.assert_called_once_with("lib1")
|
||||
|
||||
def test_empty_library_id_raises(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="library_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
|
||||
class TestCreateAssetUseCase:
|
||||
def test_creates_asset_via_repo(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.return_value = MagicMock()
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
|
||||
cmd = CreateAssetCommand(
|
||||
project_id="p1",
|
||||
library_id="l1",
|
||||
name="test.mp4",
|
||||
storage_key="videos/t.mp4",
|
||||
mime_type="video/mp4",
|
||||
file_size=1024,
|
||||
)
|
||||
result = use_case.execute(cmd)
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
created_asset = mock_repo.create.call_args[0][0]
|
||||
assert created_asset.project_id == "p1"
|
||||
assert created_asset.name == "test.mp4"
|
||||
assert created_asset.file_size == 1024
|
||||
assert created_asset.status == AssetStatus.UPLOADING
|
||||
|
||||
def test_asset_create_validation_propagates(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
|
||||
cmd = CreateAssetCommand(
|
||||
project_id="p1",
|
||||
library_id="l1",
|
||||
name="",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="素材名称不能为空"):
|
||||
use_case.execute(cmd)
|
||||
|
||||
|
||||
# ── asset_libraries.py ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateAssetLibraryCommand:
|
||||
def test_creation(self):
|
||||
cmd = CreateAssetLibraryCommand(
|
||||
project_id="p1",
|
||||
name="我的库",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
assert cmd.project_id == "p1"
|
||||
assert cmd.name == "我的库"
|
||||
assert cmd.kind == AssetLibraryKind.VIDEO
|
||||
|
||||
|
||||
class TestListAssetLibrariesUseCase:
|
||||
def test_success(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_project.return_value = []
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("p1")
|
||||
|
||||
assert result == []
|
||||
mock_repo.find_by_project.assert_called_once_with("p1")
|
||||
|
||||
def test_strips_project_id(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" p1 ")
|
||||
|
||||
mock_repo.find_by_project.assert_called_once_with("p1")
|
||||
|
||||
def test_empty_project_id_raises(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
|
||||
class TestCreateAssetLibraryUseCase:
|
||||
def test_creates_library_via_repo(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.return_value = MagicMock()
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
|
||||
cmd = CreateAssetLibraryCommand(
|
||||
project_id="p1",
|
||||
name="视频库",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
result = use_case.execute(cmd)
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
created = mock_repo.create.call_args[0][0]
|
||||
assert created.project_id == "p1"
|
||||
assert created.name == "视频库"
|
||||
assert created.kind == AssetLibraryKind.VIDEO
|
||||
|
||||
def test_validation_propagates(self):
|
||||
mock_repo = MagicMock()
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
|
||||
cmd = CreateAssetLibraryCommand(
|
||||
project_id="p1",
|
||||
name="",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="素材库名称不能为空"):
|
||||
use_case.execute(cmd)
|
||||
@@ -19,19 +19,19 @@ from typing import Optional
|
||||
import pytest
|
||||
|
||||
from packages.domain.asset_scoring import (
|
||||
AssetScoreDetail,
|
||||
MEDIUM_BUCKET_MAX,
|
||||
MIN_QUALITY_SCORE,
|
||||
OPTIMAL_DURATION_MAX,
|
||||
OPTIMAL_DURATION_MIN,
|
||||
SHORT_BUCKET_MAX,
|
||||
SmartSelectResult,
|
||||
TARGET_HEIGHT,
|
||||
TARGET_WIDTH,
|
||||
WEIGHT_BITRATE,
|
||||
WEIGHT_DURATION,
|
||||
WEIGHT_QUALITY,
|
||||
WEIGHT_RESOLUTION,
|
||||
AssetScoreDetail,
|
||||
SmartSelectResult,
|
||||
_bucket_by_duration,
|
||||
calculate_total_score,
|
||||
diverse_selection,
|
||||
|
||||
Executable
+69
@@ -0,0 +1,69 @@
|
||||
"""infer_mime_type_from_storage_key 纯逻辑单测.
|
||||
|
||||
Worker core 工具函数,从 storage_key 推断 MIME 类型。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from worker_app.core.asset_types import infer_mime_type_from_storage_key
|
||||
|
||||
|
||||
class TestInferMimeTypeFromStorageKey:
|
||||
"""infer_mime_type_from_storage_key 测试."""
|
||||
|
||||
def test_mp4_returns_video_mp4(self):
|
||||
"""mp4 后缀返回 video/mp4."""
|
||||
assert infer_mime_type_from_storage_key("projects/abc/video.mp4") == "video/mp4"
|
||||
|
||||
def test_mov_returns_quicktime(self):
|
||||
"""mov 后缀返回 video/quicktime."""
|
||||
assert infer_mime_type_from_storage_key("uploads/test.mov") == "video/quicktime"
|
||||
|
||||
def test_m4v_returns_video_mp4(self):
|
||||
"""m4v 后缀返回 video/mp4."""
|
||||
assert infer_mime_type_from_storage_key("clip.m4v") == "video/mp4"
|
||||
|
||||
def test_avi_returns_video_mp4(self):
|
||||
"""avi 后缀返回 video/mp4."""
|
||||
assert infer_mime_type_from_storage_key("movie.avi") == "video/mp4"
|
||||
|
||||
def test_mkv_returns_video_mp4(self):
|
||||
"""mkv 后缀返回 video/mp4."""
|
||||
assert infer_mime_type_from_storage_key("video.mkv") == "video/mp4"
|
||||
|
||||
def test_webm_returns_video_mp4(self):
|
||||
"""webm 后缀返回 video/mp4."""
|
||||
assert infer_mime_type_from_storage_key("output.webm") == "video/mp4"
|
||||
|
||||
def test_jpg_default_returns_image_jpeg(self):
|
||||
"""jpg 等非视频后缀默认返回 image/jpeg."""
|
||||
assert infer_mime_type_from_storage_key("thumb.jpg") == "image/jpeg"
|
||||
|
||||
def test_png_default_returns_image_jpeg(self):
|
||||
"""png 也返回 image/jpeg(当前实现的默认值)."""
|
||||
assert infer_mime_type_from_storage_key("image.png") == "image/jpeg"
|
||||
|
||||
def test_no_extension_returns_jpeg(self):
|
||||
"""无扩展名返回 image/jpeg."""
|
||||
assert infer_mime_type_from_storage_key("random_file") == "image/jpeg"
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""大小写不敏感."""
|
||||
assert infer_mime_type_from_storage_key("VIDEO.MP4") == "video/mp4"
|
||||
assert infer_mime_type_from_storage_key("Clip.MOV") == "video/quicktime"
|
||||
|
||||
def test_deep_path(self):
|
||||
"""多级路径正常推断."""
|
||||
assert infer_mime_type_from_storage_key("generated/projects/abc/def/output.mp4") == "video/mp4"
|
||||
|
||||
def test_empty_string(self):
|
||||
"""空字符串返回 image/jpeg(默认值)."""
|
||||
assert infer_mime_type_from_storage_key("") == "image/jpeg"
|
||||
|
||||
def test_filename_with_multiple_dots(self):
|
||||
"""文件名含多个点时取最后一个扩展名."""
|
||||
assert infer_mime_type_from_storage_key("my.video.file.mp4") == "video/mp4"
|
||||
|
||||
def test_mov_case_insensitive_upper(self):
|
||||
"""MOV 大写也识别为 quicktime."""
|
||||
assert infer_mime_type_from_storage_key("video.MOV") == "video/quicktime"
|
||||
Executable
+139
@@ -0,0 +1,139 @@
|
||||
"""mark_asset_used_for_generation 深度补充单测.
|
||||
|
||||
补全边界场景:空 metadata、None metadata、last_used_at 格式、
|
||||
review_status 已有值不覆盖、多次调用递增。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from worker_app.core.asset_usage import mark_asset_used_for_generation
|
||||
|
||||
from packages.domain import Asset, AssetStatus
|
||||
|
||||
|
||||
def _asset() -> Asset:
|
||||
return Asset.create(
|
||||
project_id="project-1",
|
||||
library_id="library-1",
|
||||
name="video.mp4",
|
||||
storage_key="uploads/video.mp4",
|
||||
mime_type="video/mp4",
|
||||
file_size=1024,
|
||||
status=AssetStatus.READY,
|
||||
)
|
||||
|
||||
|
||||
class TestMarkAssetUsedForGeneration:
|
||||
"""mark_asset_used_for_generation 深度测试."""
|
||||
|
||||
def test_first_use_sets_count_to_1(self):
|
||||
"""首次使用,use_count 从 0 变 1."""
|
||||
asset = _asset()
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == 1
|
||||
|
||||
def test_increments_existing_count(self):
|
||||
"""已有计数时递增."""
|
||||
asset = _asset()
|
||||
asset.metadata = {"generation_use_count": 5}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == 6
|
||||
|
||||
def test_zero_count_increments_to_1(self):
|
||||
"""计数为 0 时递增到 1."""
|
||||
asset = _asset()
|
||||
asset.metadata = {"generation_use_count": 0}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == 1
|
||||
|
||||
def test_empty_metadata_still_works(self):
|
||||
"""空 dict metadata 也能正常工作."""
|
||||
asset = _asset()
|
||||
asset.metadata = {}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == 1
|
||||
assert asset.metadata["review_status"] == "pending_review"
|
||||
assert "last_used_at" in asset.metadata
|
||||
|
||||
def test_none_metadata_field_defaults_to_0(self):
|
||||
"""metadata 中 generation_use_count 为 None 时按 0 处理."""
|
||||
asset = _asset()
|
||||
asset.metadata = {"generation_use_count": None}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == 1
|
||||
|
||||
def test_string_count_gets_casted(self):
|
||||
"""字符串类型的 use_count 通过 int() 转换."""
|
||||
asset = _asset()
|
||||
asset.metadata = {"generation_use_count": "3"}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == 4
|
||||
|
||||
def test_preserves_other_metadata_fields(self):
|
||||
"""不覆盖 metadata 中的其他字段."""
|
||||
asset = _asset()
|
||||
asset.metadata = {
|
||||
"generation_use_count": 1,
|
||||
"custom_field": "value",
|
||||
"tags": ["a", "b"],
|
||||
}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["custom_field"] == "value"
|
||||
assert asset.metadata["tags"] == ["a", "b"]
|
||||
assert asset.metadata["generation_use_count"] == 2
|
||||
|
||||
def test_review_status_pending_when_not_set(self):
|
||||
"""review_status 未设置时设为 pending_review."""
|
||||
asset = _asset()
|
||||
asset.metadata = {}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["review_status"] == "pending_review"
|
||||
|
||||
def test_review_status_not_overwritten_if_present(self):
|
||||
"""review_status 已有值时不覆盖."""
|
||||
asset = _asset()
|
||||
asset.metadata = {"review_status": "approved"}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["review_status"] == "approved"
|
||||
|
||||
def test_review_status_empty_string_considered_falsy(self):
|
||||
"""review_status 为空字符串时视为 falsy,设置为 pending_review."""
|
||||
asset = _asset()
|
||||
asset.metadata = {"review_status": ""}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["review_status"] == "pending_review"
|
||||
|
||||
def test_last_used_at_is_iso_format(self):
|
||||
"""last_used_at 是 ISO 格式时间字符串."""
|
||||
asset = _asset()
|
||||
mark_asset_used_for_generation(asset)
|
||||
ts = asset.metadata["last_used_at"]
|
||||
# 可以被解析为 ISO 格式
|
||||
parsed = datetime.fromisoformat(ts)
|
||||
assert parsed.tzinfo is not None # 带时区
|
||||
|
||||
def test_last_used_at_is_utc(self):
|
||||
"""last_used_at 是 UTC 时间."""
|
||||
asset = _asset()
|
||||
before = datetime.now(timezone.utc)
|
||||
mark_asset_used_for_generation(asset)
|
||||
after = datetime.now(timezone.utc)
|
||||
ts = datetime.fromisoformat(asset.metadata["last_used_at"])
|
||||
assert before <= ts <= after
|
||||
|
||||
def test_multiple_calls_increment_count(self):
|
||||
"""多次调用持续递增."""
|
||||
asset = _asset()
|
||||
mark_asset_used_for_generation(asset)
|
||||
mark_asset_used_for_generation(asset)
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == 3
|
||||
|
||||
def test_negative_count_still_increments(self):
|
||||
"""负数计数(异常数据)也能递增."""
|
||||
asset = _asset()
|
||||
asset.metadata = {"generation_use_count": -5}
|
||||
mark_asset_used_for_generation(asset)
|
||||
assert asset.metadata["generation_use_count"] == -4
|
||||
Executable
+649
@@ -0,0 +1,649 @@
|
||||
"""clip_operations 单元测试 — 片段分割/合并纯逻辑层。
|
||||
|
||||
覆盖:
|
||||
- validate_split_time: 分割时间校验
|
||||
- calculate_split: 分割参数计算
|
||||
- validate_merge_clips: 合并校验
|
||||
- calculate_merge: 合并参数计算
|
||||
- calculate_reorder_new_orders: 重排 order 计算
|
||||
- calculate_shift_orders: order 偏移计算
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from dataclasses import dataclass
|
||||
|
||||
from packages.domain.clip_operations import (
|
||||
ROUND_PRECISION,
|
||||
MergeResult,
|
||||
SplitResult,
|
||||
calculate_merge,
|
||||
calculate_reorder_new_orders,
|
||||
calculate_shift_orders,
|
||||
calculate_split,
|
||||
validate_merge_clips,
|
||||
validate_split_time,
|
||||
)
|
||||
|
||||
# ── Mock Clip ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class _MockClip:
|
||||
id: str = "c1"
|
||||
plan_id: str = "p1"
|
||||
order: int = 0
|
||||
clip_type: str = "main"
|
||||
duration: float = 5.0
|
||||
start_time: float = 0.0
|
||||
text_content: str = ""
|
||||
config: dict | None = None
|
||||
|
||||
|
||||
# ── validate_split_time 测试 ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateSplitTime(unittest.TestCase):
|
||||
"""validate_split_time 分割时间校验测试。"""
|
||||
|
||||
def test_valid_split(self):
|
||||
"""合法分割时间不报错。"""
|
||||
validate_split_time(2.5, 5.0) # 不抛异常
|
||||
|
||||
def test_split_at_zero(self):
|
||||
"""分割时间为 0 时报错。"""
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
validate_split_time(0.0, 5.0)
|
||||
self.assertIn("分割时间", str(ctx.exception))
|
||||
|
||||
def test_split_negative(self):
|
||||
"""分割时间为负数时报错。"""
|
||||
with self.assertRaises(ValueError):
|
||||
validate_split_time(-1.0, 5.0)
|
||||
|
||||
def test_split_at_duration(self):
|
||||
"""分割时间等于 duration 时报错。"""
|
||||
with self.assertRaises(ValueError):
|
||||
validate_split_time(5.0, 5.0)
|
||||
|
||||
def test_split_over_duration(self):
|
||||
"""分割时间超过 duration 时报错。"""
|
||||
with self.assertRaises(ValueError):
|
||||
validate_split_time(6.0, 5.0)
|
||||
|
||||
def test_split_very_small(self):
|
||||
"""很小的正数是合法的。"""
|
||||
validate_split_time(0.001, 5.0) # 不抛异常
|
||||
|
||||
def test_split_just_below_duration(self):
|
||||
"""略小于 duration 是合法的。"""
|
||||
validate_split_time(4.999, 5.0) # 不抛异常
|
||||
|
||||
def test_error_message_contains_duration(self):
|
||||
"""错误消息包含 duration 值。"""
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
validate_split_time(6.0, 5.0)
|
||||
self.assertIn("5.000", str(ctx.exception))
|
||||
|
||||
|
||||
# ── calculate_split 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateSplit(unittest.TestCase):
|
||||
"""calculate_split 分割计算测试。"""
|
||||
|
||||
def test_middle_split(self):
|
||||
"""从中间分割。"""
|
||||
result = calculate_split(duration=10.0, split_time=5.0)
|
||||
self.assertIsInstance(result, SplitResult)
|
||||
self.assertEqual(result.left_duration, 5.0)
|
||||
self.assertEqual(result.right_duration, 5.0)
|
||||
self.assertEqual(result.right_start_time, 5.0)
|
||||
self.assertEqual(result.left_trim_end, 5.0)
|
||||
self.assertEqual(result.right_trim_start, 5.0)
|
||||
|
||||
def test_early_split(self):
|
||||
"""从开头附近分割。"""
|
||||
result = calculate_split(duration=10.0, split_time=2.0)
|
||||
self.assertEqual(result.left_duration, 2.0)
|
||||
self.assertEqual(result.right_duration, 8.0)
|
||||
self.assertEqual(result.right_start_time, 2.0)
|
||||
|
||||
def test_late_split(self):
|
||||
"""从结尾附近分割。"""
|
||||
result = calculate_split(duration=10.0, split_time=8.0)
|
||||
self.assertEqual(result.left_duration, 8.0)
|
||||
self.assertEqual(result.right_duration, 2.0)
|
||||
|
||||
def test_with_start_time_offset(self):
|
||||
"""带 start_time 偏移。"""
|
||||
result = calculate_split(duration=5.0, split_time=2.0, start_time=10.0)
|
||||
self.assertEqual(result.left_duration, 2.0)
|
||||
self.assertEqual(result.right_duration, 3.0)
|
||||
self.assertEqual(result.right_start_time, 12.0)
|
||||
|
||||
def test_zero_start_time(self):
|
||||
"""start_time 为 0 时 right_start_time 等于 left_duration。"""
|
||||
result = calculate_split(duration=5.0, split_time=2.0, start_time=0.0)
|
||||
self.assertEqual(result.right_start_time, result.left_duration)
|
||||
|
||||
def test_round_to_precision(self):
|
||||
"""结果精度符合设置。"""
|
||||
result = calculate_split(duration=1.0, split_time=1 / 3, precision=3)
|
||||
# 1/3 ≈ 0.333(3位精度)
|
||||
self.assertAlmostEqual(result.left_duration, 0.333, places=3)
|
||||
self.assertAlmostEqual(result.right_duration, 0.667, places=3)
|
||||
|
||||
def test_default_precision_3(self):
|
||||
"""默认精度为 3 位小数。"""
|
||||
self.assertEqual(ROUND_PRECISION, 3)
|
||||
result = calculate_split(duration=1.0, split_time=0.333333)
|
||||
# 默认用 ROUND_PRECISION = 3
|
||||
self.assertEqual(result.left_duration, 0.333)
|
||||
|
||||
def test_custom_precision(self):
|
||||
"""自定义精度。"""
|
||||
result = calculate_split(duration=1.0, split_time=0.123456, precision=5)
|
||||
self.assertEqual(result.left_duration, 0.12346)
|
||||
self.assertEqual(result.right_duration, 0.87654)
|
||||
|
||||
def test_invalid_split_time_raises(self):
|
||||
"""不合法的分割时间抛出 ValueError。"""
|
||||
with self.assertRaises(ValueError):
|
||||
calculate_split(duration=5.0, split_time=0.0)
|
||||
with self.assertRaises(ValueError):
|
||||
calculate_split(duration=5.0, split_time=5.0)
|
||||
with self.assertRaises(ValueError):
|
||||
calculate_split(duration=5.0, split_time=-1.0)
|
||||
|
||||
def test_trim_values_match_durations(self):
|
||||
"""trim 值与对应时长一致。"""
|
||||
result = calculate_split(duration=7.5, split_time=3.0)
|
||||
self.assertEqual(result.right_trim_start, result.left_duration)
|
||||
self.assertEqual(result.left_trim_end, result.right_duration)
|
||||
|
||||
def test_frozen_result(self):
|
||||
"""SplitResult 是 frozen dataclass。"""
|
||||
result = calculate_split(duration=5.0, split_time=2.0)
|
||||
with self.assertRaises(Exception):
|
||||
result.left_duration = 3.0 # type: ignore[misc]
|
||||
|
||||
|
||||
# ── validate_merge_clips 测试 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateMergeClips(unittest.TestCase):
|
||||
"""validate_merge_clips 合并校验测试。"""
|
||||
|
||||
def test_two_consecutive_clips(self):
|
||||
"""两个连续片段:校验通过。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", plan_id="p1", order=0, clip_type="main"),
|
||||
_MockClip(id="c2", plan_id="p1", order=1, clip_type="main"),
|
||||
]
|
||||
plan_id, first_order = validate_merge_clips(clips)
|
||||
self.assertEqual(plan_id, "p1")
|
||||
self.assertEqual(first_order, 0)
|
||||
|
||||
def test_three_consecutive_clips(self):
|
||||
"""三个连续片段:校验通过。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", plan_id="p1", order=2, clip_type="main"),
|
||||
_MockClip(id="c2", plan_id="p1", order=3, clip_type="main"),
|
||||
_MockClip(id="c3", plan_id="p1", order=4, clip_type="main"),
|
||||
]
|
||||
plan_id, first_order = validate_merge_clips(clips)
|
||||
self.assertEqual(plan_id, "p1")
|
||||
self.assertEqual(first_order, 2)
|
||||
|
||||
def test_unordered_input(self):
|
||||
"""输入顺序不影响校验(自动排序)。"""
|
||||
clips = [
|
||||
_MockClip(id="c2", plan_id="p1", order=1, clip_type="main"),
|
||||
_MockClip(id="c1", plan_id="p1", order=0, clip_type="main"),
|
||||
]
|
||||
plan_id, first_order = validate_merge_clips(clips)
|
||||
self.assertEqual(plan_id, "p1")
|
||||
self.assertEqual(first_order, 0)
|
||||
|
||||
def test_single_clip_raises(self):
|
||||
"""只有一个片段报错。"""
|
||||
clips = [_MockClip(id="c1", plan_id="p1", order=0)]
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
validate_merge_clips(clips)
|
||||
self.assertIn("至少需要 2 个", str(ctx.exception))
|
||||
|
||||
def test_empty_clips_raises(self):
|
||||
"""空列表报错。"""
|
||||
with self.assertRaises(ValueError):
|
||||
validate_merge_clips([])
|
||||
|
||||
def test_different_plan_raises(self):
|
||||
"""不同计划的片段报错。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", plan_id="p1", order=0, clip_type="main"),
|
||||
_MockClip(id="c2", plan_id="p2", order=1, clip_type="main"),
|
||||
]
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
validate_merge_clips(clips)
|
||||
self.assertIn("同一计划", str(ctx.exception))
|
||||
|
||||
def test_non_consecutive_raises(self):
|
||||
"""不连续的片段报错。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", plan_id="p1", order=0, clip_type="main"),
|
||||
_MockClip(id="c2", plan_id="p1", order=2, clip_type="main"),
|
||||
]
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
validate_merge_clips(clips)
|
||||
self.assertIn("不连续", str(ctx.exception))
|
||||
|
||||
def test_different_type_raises(self):
|
||||
"""类型不同报错。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", plan_id="p1", order=0, clip_type="main"),
|
||||
_MockClip(id="c2", plan_id="p1", order=1, clip_type="title"),
|
||||
]
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
validate_merge_clips(clips)
|
||||
self.assertIn("相同类型", str(ctx.exception))
|
||||
|
||||
def test_gap_in_middle_raises(self):
|
||||
"""中间有间隔报错(三个片段中间缺一个)。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", plan_id="p1", order=0, clip_type="main"),
|
||||
_MockClip(id="c2", plan_id="p1", order=1, clip_type="main"),
|
||||
_MockClip(id="c3", plan_id="p1", order=3, clip_type="main"),
|
||||
]
|
||||
with self.assertRaises(ValueError):
|
||||
validate_merge_clips(clips)
|
||||
|
||||
|
||||
# ── calculate_merge 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateMerge(unittest.TestCase):
|
||||
"""calculate_merge 合并计算测试。"""
|
||||
|
||||
def test_two_clips_total_duration(self):
|
||||
"""两个片段:总时长相加。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=5.0),
|
||||
_MockClip(id="c2", order=1, duration=3.0),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertIsInstance(result, MergeResult)
|
||||
self.assertEqual(result.total_duration, 8.0)
|
||||
self.assertEqual(result.first_order, 0)
|
||||
self.assertEqual(result.shift_amount, 1)
|
||||
|
||||
def test_three_clips_total_duration(self):
|
||||
"""三个片段:总时长相加。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=2, duration=2.0),
|
||||
_MockClip(id="c2", order=3, duration=3.5),
|
||||
_MockClip(id="c3", order=4, duration=4.5),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.total_duration, 10.0)
|
||||
self.assertEqual(result.first_order, 2)
|
||||
self.assertEqual(result.shift_amount, 2)
|
||||
|
||||
def test_unordered_input(self):
|
||||
"""输入顺序不影响结果(自动按 order 排序)。"""
|
||||
clips = [
|
||||
_MockClip(id="c2", order=1, duration=3.0),
|
||||
_MockClip(id="c1", order=0, duration=5.0),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.total_duration, 8.0)
|
||||
self.assertEqual(result.first_order, 0)
|
||||
|
||||
def test_merge_text_newline_join(self):
|
||||
"""文案用换行连接,跳过空字符串。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=2.0, text_content="第一段"),
|
||||
_MockClip(id="c2", order=1, duration=3.0, text_content="第二段"),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.merged_text, "第一段\n第二段")
|
||||
|
||||
def test_merge_text_skip_empty(self):
|
||||
"""空文案或纯空格被跳过。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=2.0, text_content=""),
|
||||
_MockClip(id="c2", order=1, duration=3.0, text_content="有内容"),
|
||||
_MockClip(id="c3", order=2, duration=1.0, text_content=" "),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.merged_text, "有内容")
|
||||
|
||||
def test_all_empty_text(self):
|
||||
"""所有文案都空时合并结果为空字符串。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=2.0, text_content=""),
|
||||
_MockClip(id="c2", order=1, duration=3.0, text_content=" "),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.merged_text, "")
|
||||
|
||||
def test_merge_config_later_overwrites(self):
|
||||
"""后面的 config 覆盖前面的。"""
|
||||
clips = [
|
||||
_MockClip(
|
||||
id="c1",
|
||||
order=0,
|
||||
duration=2.0,
|
||||
config={"color": "red", "speed": 1.0},
|
||||
),
|
||||
_MockClip(
|
||||
id="c2",
|
||||
order=1,
|
||||
duration=3.0,
|
||||
config={"color": "blue", "filter": "vintage"},
|
||||
),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.merged_config["color"], "blue") # 后面的覆盖
|
||||
self.assertEqual(result.merged_config["speed"], 1.0) # 保留前面的
|
||||
self.assertEqual(result.merged_config["filter"], "vintage") # 新增的
|
||||
|
||||
def test_merge_config_none_handled(self):
|
||||
"""config 为 None 时正常处理。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=2.0, config=None),
|
||||
_MockClip(id="c2", order=1, duration=3.0, config={"key": "value"}),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.merged_config["key"], "value")
|
||||
|
||||
def test_trim_fields_removed(self):
|
||||
"""trim_start 和 trim_end 被移除(合并后是完整片段)。"""
|
||||
clips = [
|
||||
_MockClip(
|
||||
id="c1",
|
||||
order=0,
|
||||
duration=2.0,
|
||||
config={"trim_end": 2.0, "color": "red"},
|
||||
),
|
||||
_MockClip(
|
||||
id="c2",
|
||||
order=1,
|
||||
duration=3.0,
|
||||
config={"trim_start": 1.0, "speed": 1.5},
|
||||
),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
self.assertNotIn("trim_start", result.merged_config)
|
||||
self.assertNotIn("trim_end", result.merged_config)
|
||||
self.assertEqual(result.merged_config["color"], "red")
|
||||
self.assertEqual(result.merged_config["speed"], 1.5)
|
||||
|
||||
def test_empty_clips_raises(self):
|
||||
"""空列表报错。"""
|
||||
with self.assertRaises(ValueError):
|
||||
calculate_merge([])
|
||||
|
||||
def test_single_clip_merge(self):
|
||||
"""单个片段也能计算(虽然 validate 会限制,但函数本身支持)。"""
|
||||
clips = [_MockClip(id="c1", order=5, duration=5.0, text_content="唯一")]
|
||||
result = calculate_merge(clips)
|
||||
self.assertEqual(result.total_duration, 5.0)
|
||||
self.assertEqual(result.first_order, 5)
|
||||
self.assertEqual(result.shift_amount, 0)
|
||||
self.assertEqual(result.merged_text, "唯一")
|
||||
|
||||
def test_round_precision(self):
|
||||
"""时长精度符合设置。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=1 / 3),
|
||||
_MockClip(id="c2", order=1, duration=1 / 3),
|
||||
]
|
||||
result = calculate_merge(clips, precision=3)
|
||||
self.assertEqual(result.total_duration, 0.667)
|
||||
|
||||
def test_frozen_result(self):
|
||||
"""MergeResult 是 frozen dataclass。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=2.0),
|
||||
_MockClip(id="c2", order=1, duration=3.0),
|
||||
]
|
||||
result = calculate_merge(clips)
|
||||
with self.assertRaises(Exception):
|
||||
result.total_duration = 10.0 # type: ignore[misc]
|
||||
|
||||
|
||||
# ── calculate_reorder_new_orders 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateReorderNewOrders(unittest.TestCase):
|
||||
"""calculate_reorder_new_orders 重排计算测试。"""
|
||||
|
||||
def test_two_items_swap(self):
|
||||
"""两个元素交换顺序。"""
|
||||
items = [
|
||||
_MockClip(id="a", order=0),
|
||||
_MockClip(id="b", order=1),
|
||||
]
|
||||
result = calculate_reorder_new_orders(["b", "a"], items)
|
||||
self.assertEqual(result["b"], 0)
|
||||
self.assertEqual(result["a"], 1)
|
||||
|
||||
def test_three_items_reorder(self):
|
||||
"""三个元素重新排序。"""
|
||||
items = [
|
||||
_MockClip(id="a", order=0),
|
||||
_MockClip(id="b", order=1),
|
||||
_MockClip(id="c", order=2),
|
||||
]
|
||||
result = calculate_reorder_new_orders(["c", "a", "b"], items)
|
||||
self.assertEqual(result["c"], 0)
|
||||
self.assertEqual(result["a"], 1)
|
||||
self.assertEqual(result["b"], 2)
|
||||
|
||||
def test_same_order(self):
|
||||
"""顺序不变。"""
|
||||
items = [
|
||||
_MockClip(id="a", order=0),
|
||||
_MockClip(id="b", order=1),
|
||||
]
|
||||
result = calculate_reorder_new_orders(["a", "b"], items)
|
||||
self.assertEqual(result["a"], 0)
|
||||
self.assertEqual(result["b"], 1)
|
||||
|
||||
def test_mismatched_ids_raises(self):
|
||||
"""ID 不匹配报错。"""
|
||||
items = [
|
||||
_MockClip(id="a", order=0),
|
||||
_MockClip(id="b", order=1),
|
||||
]
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
calculate_reorder_new_orders(["a", "c"], items)
|
||||
self.assertIn("不匹配", str(ctx.exception))
|
||||
|
||||
def test_extra_id_in_list_raises(self):
|
||||
"""有序列表多出 ID 报错。"""
|
||||
items = [_MockClip(id="a", order=0)]
|
||||
with self.assertRaises(ValueError):
|
||||
calculate_reorder_new_orders(["a", "b"], items)
|
||||
|
||||
def test_missing_id_raises(self):
|
||||
"""有序列表缺少 ID 报错。"""
|
||||
items = [
|
||||
_MockClip(id="a", order=0),
|
||||
_MockClip(id="b", order=1),
|
||||
]
|
||||
with self.assertRaises(ValueError):
|
||||
calculate_reorder_new_orders(["a"], items)
|
||||
|
||||
def test_custom_attr_names(self):
|
||||
"""自定义属性名。"""
|
||||
|
||||
@dataclass
|
||||
class _Item:
|
||||
key: str
|
||||
pos: int
|
||||
|
||||
items = [_Item(key="x", pos=0), _Item(key="y", pos=1)]
|
||||
result = calculate_reorder_new_orders(["y", "x"], items, id_attr="key", order_attr="pos")
|
||||
self.assertEqual(result["y"], 0)
|
||||
self.assertEqual(result["x"], 1)
|
||||
|
||||
|
||||
# ── calculate_shift_orders 测试 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateShiftOrders(unittest.TestCase):
|
||||
"""calculate_shift_orders order 偏移计算测试。"""
|
||||
|
||||
def test_shift_positive(self):
|
||||
"""正偏移:order 增加。"""
|
||||
items = [
|
||||
_MockClip(id="c1", order=0),
|
||||
_MockClip(id="c2", order=1),
|
||||
_MockClip(id="c3", order=2),
|
||||
]
|
||||
result = calculate_shift_orders(items, threshold_order=0, shift=1)
|
||||
shifted = {c.id: new_order for c, new_order in result}
|
||||
# order > 0 的才会被偏移
|
||||
self.assertEqual(len(shifted), 2)
|
||||
self.assertEqual(shifted["c2"], 2)
|
||||
self.assertEqual(shifted["c3"], 3)
|
||||
self.assertNotIn("c1", shifted)
|
||||
|
||||
def test_shift_negative(self):
|
||||
"""负偏移:order 减少。"""
|
||||
items = [
|
||||
_MockClip(id="c1", order=0),
|
||||
_MockClip(id="c2", order=2),
|
||||
_MockClip(id="c3", order=3),
|
||||
]
|
||||
result = calculate_shift_orders(items, threshold_order=1, shift=-1)
|
||||
shifted = {c.id: new_order for c, new_order in result}
|
||||
self.assertEqual(shifted["c2"], 1)
|
||||
self.assertEqual(shifted["c3"], 2)
|
||||
self.assertNotIn("c1", shifted)
|
||||
|
||||
def test_excluded_ids(self):
|
||||
"""排除指定 ID。"""
|
||||
items = [
|
||||
_MockClip(id="c1", order=0),
|
||||
_MockClip(id="c2", order=1),
|
||||
_MockClip(id="c3", order=2),
|
||||
]
|
||||
result = calculate_shift_orders(
|
||||
items,
|
||||
threshold_order=0,
|
||||
shift=1,
|
||||
excluded_ids={"c2"},
|
||||
)
|
||||
shifted = {c.id: new_order for c, new_order in result}
|
||||
self.assertNotIn("c1", shifted)
|
||||
self.assertNotIn("c2", shifted) # 被排除
|
||||
self.assertEqual(shifted["c3"], 3)
|
||||
|
||||
def test_no_items_above_threshold(self):
|
||||
"""没有元素高于阈值时返回空。"""
|
||||
items = [
|
||||
_MockClip(id="c1", order=0),
|
||||
_MockClip(id="c2", order=1),
|
||||
]
|
||||
result = calculate_shift_orders(items, threshold_order=5, shift=1)
|
||||
self.assertEqual(result, [])
|
||||
|
||||
def test_zero_shift(self):
|
||||
"""偏移量为 0 时仍返回(虽然 order 不变)。"""
|
||||
items = [
|
||||
_MockClip(id="c1", order=0),
|
||||
_MockClip(id="c2", order=1),
|
||||
]
|
||||
result = calculate_shift_orders(items, threshold_order=0, shift=0)
|
||||
shifted = {c.id: new_order for c, new_order in result}
|
||||
self.assertEqual(shifted["c2"], 1) # 1 + 0 = 1
|
||||
|
||||
def test_threshold_exclusive(self):
|
||||
"""阈值是严格大于(不包含等于)。"""
|
||||
items = [
|
||||
_MockClip(id="c1", order=5), # 等于 threshold,不偏移
|
||||
_MockClip(id="c2", order=6), # 大于 threshold,偏移
|
||||
]
|
||||
result = calculate_shift_orders(items, threshold_order=5, shift=2)
|
||||
shifted = {c.id: new_order for c, new_order in result}
|
||||
self.assertNotIn("c1", shifted)
|
||||
self.assertEqual(shifted["c2"], 8)
|
||||
|
||||
def test_empty_items(self):
|
||||
"""空 items 返回空列表。"""
|
||||
result = calculate_shift_orders([], threshold_order=0, shift=1)
|
||||
self.assertEqual(result, [])
|
||||
|
||||
|
||||
# ── 集成测试 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestEndToEndClipOperations(unittest.TestCase):
|
||||
"""端到端集成测试:完整分割+重排流程。"""
|
||||
|
||||
def test_split_then_shift(self):
|
||||
"""分割一个片段后,后面的片段 order +1。"""
|
||||
# 模拟:有3个片段,分割第1个(order=1),后面的+1
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=3.0),
|
||||
_MockClip(id="c2", order=1, duration=5.0),
|
||||
_MockClip(id="c3", order=2, duration=4.0),
|
||||
]
|
||||
|
||||
# 计算分割
|
||||
split = calculate_split(duration=5.0, split_time=2.0)
|
||||
|
||||
# 后面的片段 order +1
|
||||
shift_result = calculate_shift_orders(
|
||||
clips,
|
||||
threshold_order=1,
|
||||
shift=1,
|
||||
excluded_ids={"c2"},
|
||||
)
|
||||
|
||||
shifted = {c.id: new_order for c, new_order in shift_result}
|
||||
self.assertEqual(shifted["c3"], 3) # 2+1
|
||||
self.assertNotIn("c1", shifted) # order <= 1
|
||||
self.assertNotIn("c2", shifted) # 被排除
|
||||
|
||||
# 左半部分保留 order=1,右半部分 order=2
|
||||
self.assertEqual(split.left_duration, 2.0)
|
||||
self.assertEqual(split.right_duration, 3.0)
|
||||
|
||||
def test_merge_then_shift(self):
|
||||
"""合并两个片段后,后面的片段 order -1。"""
|
||||
clips = [
|
||||
_MockClip(id="c1", order=0, duration=3.0, text_content="A"),
|
||||
_MockClip(id="c2", order=1, duration=2.0, text_content="B"),
|
||||
_MockClip(id="c3", order=2, duration=4.0),
|
||||
_MockClip(id="c4", order=3, duration=1.0),
|
||||
]
|
||||
|
||||
# 校验合并
|
||||
validate_merge_clips(clips[:2])
|
||||
|
||||
# 计算合并
|
||||
merge = calculate_merge(clips[:2])
|
||||
self.assertEqual(merge.total_duration, 5.0)
|
||||
self.assertEqual(merge.merged_text, "A\nB")
|
||||
self.assertEqual(merge.shift_amount, 1)
|
||||
|
||||
# 后面的片段前移 1 位
|
||||
shift_result = calculate_shift_orders(
|
||||
clips,
|
||||
threshold_order=0, # order > 0 的
|
||||
shift=-merge.shift_amount,
|
||||
excluded_ids={"c1", "c2"},
|
||||
)
|
||||
|
||||
shifted = {c.id: new_order for c, new_order in shift_result}
|
||||
self.assertEqual(shifted["c3"], 1) # 2-1
|
||||
self.assertEqual(shifted["c4"], 2) # 3-1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Executable
+488
@@ -0,0 +1,488 @@
|
||||
"""Domain entities 单元测试。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.classification import (
|
||||
AssetLibraryKind,
|
||||
ClassificationStatus,
|
||||
IngestJobStatus,
|
||||
)
|
||||
from packages.domain.entities import (
|
||||
Asset,
|
||||
AssetLibrary,
|
||||
AssetStatus,
|
||||
IngestJob,
|
||||
Project,
|
||||
User,
|
||||
)
|
||||
|
||||
|
||||
class TestProjectCreate:
|
||||
def test_create_success(self):
|
||||
project = Project.create(owner_user_id="user1", name="我的项目")
|
||||
assert project.id is not None
|
||||
assert len(project.id) == 32
|
||||
assert project.owner_user_id == "user1"
|
||||
assert project.name == "我的项目"
|
||||
assert project.description == ""
|
||||
assert project.shared_users == []
|
||||
assert isinstance(project.created_at, datetime)
|
||||
|
||||
def test_create_with_description(self):
|
||||
project = Project.create("u1", "Test Project", "A test description")
|
||||
assert project.description == "A test description"
|
||||
|
||||
def test_create_strips_name(self):
|
||||
project = Project.create("u1", " 带空格的项目 ")
|
||||
assert project.name == "带空格的项目"
|
||||
|
||||
def test_create_strips_description(self):
|
||||
project = Project.create("u1", "P1", " desc ")
|
||||
assert project.description == "desc"
|
||||
|
||||
def test_create_empty_name(self):
|
||||
with pytest.raises(ValueError, match="项目名称不能为空"):
|
||||
Project.create("u1", "")
|
||||
|
||||
def test_create_whitespace_name(self):
|
||||
with pytest.raises(ValueError, match="项目名称不能为空"):
|
||||
Project.create("u1", " \t ")
|
||||
|
||||
def test_create_unique_ids(self):
|
||||
p1 = Project.create("u1", "P1")
|
||||
p2 = Project.create("u1", "P2")
|
||||
assert p1.id != p2.id
|
||||
|
||||
|
||||
class TestProjectAccess:
|
||||
def test_is_owner_true(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
assert project.is_owner("owner1") is True
|
||||
|
||||
def test_is_owner_false(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
assert project.is_owner("other") is False
|
||||
|
||||
def test_is_shared_with_true(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
project.shared_users = ["user_a", "user_b"]
|
||||
assert project.is_shared_with("user_a") is True
|
||||
assert project.is_shared_with("user_b") is True
|
||||
|
||||
def test_is_shared_with_false(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
project.shared_users = ["user_a"]
|
||||
assert project.is_shared_with("user_c") is False
|
||||
|
||||
def test_can_access_owner(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
assert project.can_access("owner1") is True
|
||||
|
||||
def test_can_access_shared_user(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
project.shared_users = ["shared_user"]
|
||||
assert project.can_access("shared_user") is True
|
||||
|
||||
def test_cannot_access_other(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
assert project.can_access("stranger") is False
|
||||
|
||||
def test_empty_shared_users(self):
|
||||
project = Project.create("owner1", "P1")
|
||||
assert project.shared_users == []
|
||||
assert project.is_shared_with("anyone") is False
|
||||
|
||||
|
||||
class TestAssetLibraryCreate:
|
||||
def test_create_video_library(self):
|
||||
lib = AssetLibrary.create("proj1", "视频素材库", AssetLibraryKind.VIDEO)
|
||||
assert lib.id is not None
|
||||
assert len(lib.id) == 32
|
||||
assert lib.project_id == "proj1"
|
||||
assert lib.name == "视频素材库"
|
||||
assert lib.kind == AssetLibraryKind.VIDEO
|
||||
assert lib.asset_count == 0
|
||||
assert lib.total_size == 0
|
||||
|
||||
def test_create_voice_library(self):
|
||||
lib = AssetLibrary.create("proj1", "音乐库", AssetLibraryKind.VOICE)
|
||||
assert lib.kind == AssetLibraryKind.VOICE
|
||||
|
||||
def test_create_image_library(self):
|
||||
lib = AssetLibrary.create("proj1", "图片库", AssetLibraryKind.IMAGE)
|
||||
assert lib.kind == AssetLibraryKind.IMAGE
|
||||
|
||||
def test_create_strips_name(self):
|
||||
lib = AssetLibrary.create("p1", " 我的库 ", AssetLibraryKind.VIDEO)
|
||||
assert lib.name == "我的库"
|
||||
|
||||
def test_create_empty_name(self):
|
||||
with pytest.raises(ValueError, match="素材库名称不能为空"):
|
||||
AssetLibrary.create("p1", "", AssetLibraryKind.VIDEO)
|
||||
|
||||
def test_create_whitespace_name(self):
|
||||
with pytest.raises(ValueError, match="素材库名称不能为空"):
|
||||
AssetLibrary.create("p1", " \t ", AssetLibraryKind.VIDEO)
|
||||
|
||||
|
||||
class TestAssetStatusEnum:
|
||||
def test_basic_values(self):
|
||||
assert AssetStatus.UPLOADING.value == "uploading"
|
||||
assert AssetStatus.READY.value == "ready"
|
||||
assert AssetStatus.PROCESSING.value == "processing"
|
||||
assert AssetStatus.ERROR.value == "error"
|
||||
assert AssetStatus.DELETED.value == "deleted"
|
||||
|
||||
def test_missing_uploaded_maps_to_ready(self):
|
||||
assert AssetStatus("uploaded") == AssetStatus.READY
|
||||
|
||||
def test_missing_success_maps_to_ready(self):
|
||||
assert AssetStatus("success") == AssetStatus.READY
|
||||
|
||||
def test_missing_ok_maps_to_ready(self):
|
||||
assert AssetStatus("ok") == AssetStatus.READY
|
||||
|
||||
def test_missing_done_maps_to_ready(self):
|
||||
assert AssetStatus("done") == AssetStatus.READY
|
||||
|
||||
def test_missing_complete_maps_to_ready(self):
|
||||
assert AssetStatus("complete") == AssetStatus.READY
|
||||
|
||||
def test_missing_upload_maps_to_uploading(self):
|
||||
assert AssetStatus("upload") == AssetStatus.UPLOADING
|
||||
|
||||
def test_missing_uploading_start_maps_to_uploading(self):
|
||||
assert AssetStatus("uploading_start") == AssetStatus.UPLOADING
|
||||
|
||||
def test_missing_upload_start_maps_to_uploading(self):
|
||||
assert AssetStatus("upload_start") == AssetStatus.UPLOADING
|
||||
|
||||
def test_missing_failed_maps_to_error(self):
|
||||
assert AssetStatus("failed") == AssetStatus.ERROR
|
||||
|
||||
def test_missing_fail_maps_to_error(self):
|
||||
assert AssetStatus("fail") == AssetStatus.ERROR
|
||||
|
||||
def test_missing_err_maps_to_error(self):
|
||||
assert AssetStatus("err") == AssetStatus.ERROR
|
||||
|
||||
def test_missing_process_maps_to_processing(self):
|
||||
assert AssetStatus("process") == AssetStatus.PROCESSING
|
||||
|
||||
def test_missing_running_maps_to_processing(self):
|
||||
assert AssetStatus("running") == AssetStatus.PROCESSING
|
||||
|
||||
def test_missing_run_maps_to_processing(self):
|
||||
assert AssetStatus("run") == AssetStatus.PROCESSING
|
||||
|
||||
def test_missing_unknown_value_falls_back_to_ready(self):
|
||||
assert AssetStatus("completely_unknown_status") == AssetStatus.READY
|
||||
|
||||
def test_missing_empty_string_falls_back_to_ready(self):
|
||||
assert AssetStatus("") == AssetStatus.READY
|
||||
|
||||
def test_missing_case_insensitive(self):
|
||||
assert AssetStatus("UPLOADED") == AssetStatus.READY
|
||||
assert AssetStatus("Success") == AssetStatus.READY
|
||||
assert AssetStatus("FAILED") == AssetStatus.ERROR
|
||||
|
||||
def test_missing_with_whitespace(self):
|
||||
assert AssetStatus(" uploaded ") == AssetStatus.READY
|
||||
assert AssetStatus("\tfailed\n") == AssetStatus.ERROR
|
||||
|
||||
def test_missing_non_string_value(self):
|
||||
assert AssetStatus(None) == AssetStatus.READY
|
||||
assert AssetStatus(123) == AssetStatus.READY
|
||||
|
||||
def test_known_values_still_work(self):
|
||||
assert AssetStatus("uploading") == AssetStatus.UPLOADING
|
||||
assert AssetStatus("ready") == AssetStatus.READY
|
||||
assert AssetStatus("processing") == AssetStatus.PROCESSING
|
||||
assert AssetStatus("error") == AssetStatus.ERROR
|
||||
assert AssetStatus("deleted") == AssetStatus.DELETED
|
||||
|
||||
|
||||
class TestAssetCreate:
|
||||
def test_create_minimal(self):
|
||||
asset = Asset.create(
|
||||
project_id="proj1",
|
||||
library_id="lib1",
|
||||
name="test.mp4",
|
||||
storage_key="videos/test.mp4",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
assert asset.id is not None
|
||||
assert len(asset.id) == 32
|
||||
assert asset.project_id == "proj1"
|
||||
assert asset.library_id == "lib1"
|
||||
assert asset.name == "test.mp4"
|
||||
assert asset.storage_key == "videos/test.mp4"
|
||||
assert asset.mime_type == "video/mp4"
|
||||
assert asset.file_size == 0
|
||||
assert asset.thumbnail_url is None
|
||||
assert asset.duration is None
|
||||
assert asset.width is None
|
||||
assert asset.height is None
|
||||
assert asset.status == AssetStatus.UPLOADING
|
||||
assert asset.classification_status == ClassificationStatus.PENDING
|
||||
assert asset.quality_score is None
|
||||
assert asset.tag_ids == []
|
||||
assert isinstance(asset.created_at, datetime)
|
||||
assert isinstance(asset.updated_at, datetime)
|
||||
|
||||
def test_create_with_all_fields(self):
|
||||
asset = Asset.create(
|
||||
project_id="proj1",
|
||||
library_id="lib1",
|
||||
name="movie.mp4",
|
||||
storage_key="v/m.mp4",
|
||||
mime_type="video/mp4",
|
||||
metadata={"key": "val"},
|
||||
file_size=1024000,
|
||||
thumbnail_url="http://cdn/thumb.jpg",
|
||||
duration=120.5,
|
||||
width=1920,
|
||||
height=1080,
|
||||
fps=30.0,
|
||||
codec="h264",
|
||||
status=AssetStatus.READY,
|
||||
classification_status=ClassificationStatus.COMPLETED,
|
||||
quality_score=0.85,
|
||||
uploaded_by_user_id="user1",
|
||||
file_hash="abc123",
|
||||
)
|
||||
assert asset.file_size == 1024000
|
||||
assert asset.thumbnail_url == "http://cdn/thumb.jpg"
|
||||
assert asset.duration == 120.5
|
||||
assert asset.width == 1920
|
||||
assert asset.height == 1080
|
||||
assert asset.fps == 30.0
|
||||
assert asset.codec == "h264"
|
||||
assert asset.status == AssetStatus.READY
|
||||
assert asset.classification_status == ClassificationStatus.COMPLETED
|
||||
assert asset.quality_score == 0.85
|
||||
assert asset.uploaded_by_user_id == "user1"
|
||||
assert asset.file_hash == "abc123"
|
||||
assert asset.metadata == {"key": "val"}
|
||||
|
||||
def test_create_strips_name(self):
|
||||
asset = Asset.create("p1", "l1", " test.mp4 ", "k", "video/mp4")
|
||||
assert asset.name == "test.mp4"
|
||||
|
||||
def test_create_strips_storage_key(self):
|
||||
asset = Asset.create("p1", "l1", "n", " key.mp4 ", "video/mp4")
|
||||
assert asset.storage_key == "key.mp4"
|
||||
|
||||
def test_create_strips_mime_type(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", " video/mp4 ")
|
||||
assert asset.mime_type == "video/mp4"
|
||||
|
||||
def test_create_empty_name(self):
|
||||
with pytest.raises(ValueError, match="素材名称不能为空"):
|
||||
Asset.create("p1", "l1", "", "k", "video/mp4")
|
||||
|
||||
def test_create_empty_storage_key(self):
|
||||
with pytest.raises(ValueError, match="storage_key 不能为空"):
|
||||
Asset.create("p1", "l1", "n", "", "video/mp4")
|
||||
|
||||
def test_create_empty_mime_type(self):
|
||||
with pytest.raises(ValueError, match="mime_type 不能为空"):
|
||||
Asset.create("p1", "l1", "n", "k", "")
|
||||
|
||||
def test_create_whitespace_storage_key(self):
|
||||
with pytest.raises(ValueError, match="storage_key 不能为空"):
|
||||
Asset.create("p1", "l1", "n", " \t ", "video/mp4")
|
||||
|
||||
def test_create_none_metadata_defaults_to_empty_dict(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4", metadata=None)
|
||||
assert asset.metadata == {}
|
||||
|
||||
def test_create_unique_ids(self):
|
||||
a1 = Asset.create("p1", "l1", "n1", "k1", "video/mp4")
|
||||
a2 = Asset.create("p1", "l1", "n2", "k2", "video/mp4")
|
||||
assert a1.id != a2.id
|
||||
|
||||
|
||||
class TestAssetFileType:
|
||||
def test_video_mime(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
assert asset.file_type == "video"
|
||||
|
||||
def test_audio_mime(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "audio/mpeg")
|
||||
assert asset.file_type == "audio"
|
||||
|
||||
def test_image_mime(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "image/jpeg")
|
||||
assert asset.file_type == "image"
|
||||
|
||||
def test_simple_mime_no_slash(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "application")
|
||||
assert asset.file_type == "application"
|
||||
|
||||
|
||||
class TestAssetTags:
|
||||
def test_add_tag(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag("tag1")
|
||||
assert "tag1" in asset.tag_ids
|
||||
assert len(asset.tag_ids) == 1
|
||||
|
||||
def test_add_tag_strips(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag(" tag_trim ")
|
||||
assert "tag_trim" in asset.tag_ids
|
||||
|
||||
def test_add_tag_duplicate_prevented(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag("tag1")
|
||||
asset.add_tag("tag1")
|
||||
assert asset.tag_ids.count("tag1") == 1
|
||||
assert len(asset.tag_ids) == 1
|
||||
|
||||
def test_add_tag_empty(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
with pytest.raises(ValueError, match="标签 ID 不能为空"):
|
||||
asset.add_tag("")
|
||||
|
||||
def test_add_tag_whitespace(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
with pytest.raises(ValueError, match="标签 ID 不能为空"):
|
||||
asset.add_tag(" \t ")
|
||||
|
||||
def test_add_multiple_tags(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag("t1")
|
||||
asset.add_tag("t2")
|
||||
asset.add_tag("t3")
|
||||
assert asset.tag_ids == ["t1", "t2", "t3"]
|
||||
|
||||
def test_remove_tag(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag("t1")
|
||||
asset.add_tag("t2")
|
||||
asset.remove_tag("t1")
|
||||
assert asset.tag_ids == ["t2"]
|
||||
|
||||
def test_remove_nonexistent_tag_idempotent(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag("t1")
|
||||
# 删除不存在的标签不报错
|
||||
asset.remove_tag("nonexistent")
|
||||
assert asset.tag_ids == ["t1"]
|
||||
|
||||
def test_remove_tag_strips(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag("t1")
|
||||
asset.remove_tag(" t1 ")
|
||||
assert asset.tag_ids == []
|
||||
|
||||
def test_add_tag_updates_updated_at(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
old_time = asset.updated_at
|
||||
asset.add_tag("t1")
|
||||
assert asset.updated_at >= old_time
|
||||
|
||||
def test_remove_tag_updates_updated_at(self):
|
||||
asset = Asset.create("p1", "l1", "n", "k", "video/mp4")
|
||||
asset.add_tag("t1")
|
||||
old_time = asset.updated_at
|
||||
asset.remove_tag("t1")
|
||||
assert asset.updated_at >= old_time
|
||||
|
||||
|
||||
class TestIngestJobCreate:
|
||||
def test_create_success(self):
|
||||
job = IngestJob.create(
|
||||
project_id="proj1",
|
||||
library_id="lib1",
|
||||
storage_key="videos/test.mp4",
|
||||
)
|
||||
assert job.id is not None
|
||||
assert len(job.id) == 32
|
||||
assert job.project_id == "proj1"
|
||||
assert job.library_id == "lib1"
|
||||
assert job.storage_key == "videos/test.mp4"
|
||||
assert job.status == IngestJobStatus.PENDING
|
||||
assert job.error_message == ""
|
||||
assert job.result_asset_id == ""
|
||||
assert job.file_hash == ""
|
||||
|
||||
def test_create_with_hash(self):
|
||||
job = IngestJob.create("p1", "l1", "k", file_hash="abcdef123456")
|
||||
assert job.file_hash == "abcdef123456"
|
||||
|
||||
def test_create_strips_project_id(self):
|
||||
job = IngestJob.create(" p1 ", "l1", "k")
|
||||
assert job.project_id == "p1"
|
||||
|
||||
def test_create_strips_library_id(self):
|
||||
job = IngestJob.create("p1", " l1 ", "k")
|
||||
assert job.library_id == "l1"
|
||||
|
||||
def test_create_strips_storage_key(self):
|
||||
job = IngestJob.create("p1", "l1", " k ")
|
||||
assert job.storage_key == "k"
|
||||
|
||||
def test_create_strips_file_hash(self):
|
||||
job = IngestJob.create("p1", "l1", "k", file_hash=" hash ")
|
||||
assert job.file_hash == "hash"
|
||||
|
||||
def test_create_empty_project_id(self):
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
IngestJob.create("", "l1", "k")
|
||||
|
||||
def test_create_empty_library_id(self):
|
||||
with pytest.raises(ValueError, match="library_id 不能为空"):
|
||||
IngestJob.create("p1", "", "k")
|
||||
|
||||
def test_create_empty_storage_key(self):
|
||||
with pytest.raises(ValueError, match="storage_key 不能为空"):
|
||||
IngestJob.create("p1", "l1", "")
|
||||
|
||||
def test_create_whitespace_project_id(self):
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
IngestJob.create(" \t ", "l1", "k")
|
||||
|
||||
def test_create_unique_ids(self):
|
||||
j1 = IngestJob.create("p1", "l1", "k1")
|
||||
j2 = IngestJob.create("p1", "l1", "k2")
|
||||
assert j1.id != j2.id
|
||||
|
||||
|
||||
class TestUserDataclass:
|
||||
def test_default_values(self):
|
||||
user = User(id="u1", email="test@example.com", display_name="Test User")
|
||||
assert user.id == "u1"
|
||||
assert user.email == "test@example.com"
|
||||
assert user.display_name == "Test User"
|
||||
assert user.username == ""
|
||||
assert user.password_hash == ""
|
||||
assert user.email_verified is False
|
||||
assert user.subscription_plan == "free"
|
||||
assert user.subscription_status == "active"
|
||||
assert user.max_projects == 3
|
||||
assert user.max_storage_gb == 10
|
||||
assert user.used_storage_gb == 0.0
|
||||
assert user.is_admin is False
|
||||
assert user.wechat_openid is None
|
||||
assert user.phone is None
|
||||
assert user.phone_verified is False
|
||||
assert isinstance(user.created_at, datetime)
|
||||
|
||||
def test_admin_user(self):
|
||||
user = User(id="admin", email="admin@test.com", display_name="Admin", is_admin=True)
|
||||
assert user.is_admin is True
|
||||
|
||||
def test_pro_subscription(self):
|
||||
user = User(
|
||||
id="u1",
|
||||
email="u@t.com",
|
||||
display_name="U",
|
||||
subscription_plan="pro",
|
||||
max_storage_gb=100,
|
||||
)
|
||||
assert user.subscription_plan == "pro"
|
||||
assert user.max_storage_gb == 100
|
||||
Executable
+391
@@ -0,0 +1,391 @@
|
||||
"""Domain 小模块合集单元测试。
|
||||
|
||||
覆盖零测试的小 domain 模块:
|
||||
- EditingMode 枚举
|
||||
- Template / TemplateSegment
|
||||
- TemplateClipConfig + ClipType + TransitionEffect
|
||||
- EditTemplateVersion
|
||||
- VoiceLibraryItem
|
||||
- TitleLibraryItem
|
||||
- Recipe / RecipeItem
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.editing_mode import EditingMode
|
||||
from packages.domain.recipe import RecipeItem
|
||||
from packages.domain.template import TemplateSegment
|
||||
from packages.domain.template_clip_config import (
|
||||
ClipType,
|
||||
TemplateClipConfig,
|
||||
TransitionEffect,
|
||||
)
|
||||
from packages.domain.template_version import EditTemplateVersion
|
||||
from packages.domain.title_library import TitleLibraryItem
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
|
||||
class TestEditingMode:
|
||||
def test_all_modes_exist(self):
|
||||
assert EditingMode.ONE_TAKE.value == "one_take"
|
||||
assert EditingMode.PIP.value == "pip"
|
||||
assert EditingMode.VOICE_OVER.value == "voice_over"
|
||||
assert EditingMode.VOICE_PIP.value == "voice_pip"
|
||||
|
||||
def test_from_string(self):
|
||||
assert EditingMode("one_take") == EditingMode.ONE_TAKE
|
||||
assert EditingMode("voice_over") == EditingMode.VOICE_OVER
|
||||
|
||||
def test_invalid_mode_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
EditingMode("invalid_mode")
|
||||
|
||||
def test_is_str_enum(self):
|
||||
# StrEnum 的值是字符串,可以直接比较
|
||||
assert EditingMode.ONE_TAKE == "one_take"
|
||||
|
||||
|
||||
class TestTemplateSegment:
|
||||
def test_create_minimal(self):
|
||||
seg = TemplateSegment(
|
||||
id="seg1",
|
||||
template_id="tpl1",
|
||||
segment_order=1,
|
||||
duration_min=5.0,
|
||||
duration_max=10.0,
|
||||
)
|
||||
assert seg.id == "seg1"
|
||||
assert seg.template_id == "tpl1"
|
||||
assert seg.segment_order == 1
|
||||
assert seg.duration_min == 5.0
|
||||
assert seg.duration_max == 10.0
|
||||
assert seg.material_type is None
|
||||
assert isinstance(seg.created_at, datetime)
|
||||
|
||||
def test_create_with_material_type(self):
|
||||
seg = TemplateSegment(
|
||||
id="seg2",
|
||||
template_id="tpl1",
|
||||
segment_order=2,
|
||||
duration_min=3.0,
|
||||
duration_max=8.0,
|
||||
material_type="人物",
|
||||
)
|
||||
assert seg.material_type == "人物"
|
||||
|
||||
|
||||
class TestClipType:
|
||||
def test_basic_types_exist(self):
|
||||
assert hasattr(ClipType, "MAIN")
|
||||
assert hasattr(ClipType, "INTRO")
|
||||
assert hasattr(ClipType, "OUTRO")
|
||||
assert hasattr(ClipType, "TRANSITION")
|
||||
|
||||
def test_values_are_strings(self):
|
||||
for ct in ClipType:
|
||||
assert isinstance(ct.value, str)
|
||||
|
||||
|
||||
class TestTransitionEffect:
|
||||
def test_effects_exist(self):
|
||||
assert TransitionEffect.CUT.value == "cut"
|
||||
assert TransitionEffect.FADE.value == "fade"
|
||||
assert TransitionEffect.DISSOLVE.value == "dissolve"
|
||||
# 至少有 5 种以上转场效果
|
||||
assert len(list(TransitionEffect)) >= 5
|
||||
|
||||
|
||||
class TestTemplateClipConfig:
|
||||
def test_create_minimal(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="tpl1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
min_duration=3.0,
|
||||
max_duration=8.0,
|
||||
)
|
||||
assert config.id is not None
|
||||
assert config.template_id == "tpl1"
|
||||
assert config.clip_type == ClipType.MAIN
|
||||
assert config.order == 1
|
||||
assert config.min_duration == 3.0
|
||||
assert config.max_duration == 8.0
|
||||
|
||||
def test_create_with_string_type(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="tpl1",
|
||||
clip_type="intro",
|
||||
order=0,
|
||||
min_duration=2.0,
|
||||
max_duration=5.0,
|
||||
)
|
||||
assert config.clip_type == ClipType.INTRO
|
||||
|
||||
def test_has_duration_range_true(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
min_duration=3.0,
|
||||
max_duration=8.0,
|
||||
)
|
||||
assert config.has_duration_range is True
|
||||
|
||||
def test_has_duration_range_false_when_both_zero(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
)
|
||||
assert config.has_duration_range is False
|
||||
|
||||
def test_default_duration_midpoint(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
min_duration=4.0,
|
||||
max_duration=6.0,
|
||||
)
|
||||
assert config.default_duration == pytest.approx(5.0)
|
||||
|
||||
def test_default_duration_when_only_max(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
max_duration=5.0,
|
||||
)
|
||||
assert config.default_duration == 5.0
|
||||
|
||||
def test_create_negative_min_duration_raises(self):
|
||||
with pytest.raises(ValueError, match="min_duration"):
|
||||
TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
min_duration=-1.0,
|
||||
)
|
||||
|
||||
def test_create_min_greater_than_max_raises(self):
|
||||
with pytest.raises(ValueError, match="min_duration.*max_duration"):
|
||||
TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
min_duration=10.0,
|
||||
max_duration=5.0,
|
||||
)
|
||||
|
||||
def test_create_empty_template_id_raises(self):
|
||||
with pytest.raises(ValueError, match="template_id"):
|
||||
TemplateClipConfig.create(
|
||||
template_id="",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
)
|
||||
|
||||
def test_default_transition_is_cut(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
)
|
||||
assert config.transition_effect == TransitionEffect.CUT
|
||||
|
||||
def test_custom_transition_effect(self):
|
||||
config = TemplateClipConfig.create(
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
transition_effect="fade",
|
||||
)
|
||||
assert config.transition_effect == TransitionEffect.FADE
|
||||
|
||||
|
||||
class TestEditTemplateVersion:
|
||||
def test_create_minimal(self):
|
||||
version = EditTemplateVersion.create(
|
||||
template_id="tpl1",
|
||||
version=1,
|
||||
)
|
||||
assert version.id is not None
|
||||
assert len(version.id) == 32
|
||||
assert version.template_id == "tpl1"
|
||||
assert version.version == 1
|
||||
assert version.config == {}
|
||||
assert version.clip_configs == []
|
||||
assert version.published_by == ""
|
||||
assert version.change_note == ""
|
||||
assert version.name == ""
|
||||
assert version.editing_mode == "one_take"
|
||||
assert isinstance(version.created_at, datetime)
|
||||
|
||||
def test_create_with_config_and_clip_configs(self):
|
||||
version = EditTemplateVersion.create(
|
||||
template_id="tpl1",
|
||||
version=2,
|
||||
config={"layout": "one_take"},
|
||||
clip_configs=[{"clip_id": "c1", "type": "main"}],
|
||||
published_by="user1",
|
||||
change_note="添加了片头效果",
|
||||
)
|
||||
assert version.config == {"layout": "one_take"}
|
||||
assert len(version.clip_configs) == 1
|
||||
assert version.published_by == "user1"
|
||||
assert version.change_note == "添加了片头效果"
|
||||
|
||||
def test_create_with_name_and_mode(self):
|
||||
version = EditTemplateVersion.create(
|
||||
template_id="t1",
|
||||
version=1,
|
||||
name="v1.0 正式版",
|
||||
editing_mode="voice_over",
|
||||
)
|
||||
assert version.name == "v1.0 正式版"
|
||||
assert version.editing_mode == "voice_over"
|
||||
|
||||
def test_create_unique_ids(self):
|
||||
v1 = EditTemplateVersion.create("t1", 1)
|
||||
v2 = EditTemplateVersion.create("t1", 2)
|
||||
assert v1.id != v2.id
|
||||
|
||||
def test_none_config_defaults_to_empty_dict(self):
|
||||
version = EditTemplateVersion.create("t1", 1, config=None)
|
||||
assert version.config == {}
|
||||
|
||||
def test_none_clip_configs_defaults_to_empty_list(self):
|
||||
version = EditTemplateVersion.create("t1", 1, clip_configs=None)
|
||||
assert version.clip_configs == []
|
||||
|
||||
|
||||
class TestVoiceLibraryItem:
|
||||
def test_create_minimal(self):
|
||||
item = VoiceLibraryItem(
|
||||
id="v1",
|
||||
user_id="u1",
|
||||
name="我的配音",
|
||||
)
|
||||
assert item.id == "v1"
|
||||
assert item.user_id == "u1"
|
||||
assert item.name == "我的配音"
|
||||
assert item.text == ""
|
||||
assert item.voice_provider == ""
|
||||
assert item.duration == 0
|
||||
assert item.status == "completed"
|
||||
assert item.tags == []
|
||||
assert item.project_id is None
|
||||
assert isinstance(item.created_at, datetime)
|
||||
|
||||
def test_create_with_all_fields(self):
|
||||
item = VoiceLibraryItem(
|
||||
id="v2",
|
||||
user_id="u1",
|
||||
name="产品介绍",
|
||||
text="欢迎来到我们的产品",
|
||||
voice_provider="cosyvoice",
|
||||
voice_id="voice_001",
|
||||
voice_name="温柔女声",
|
||||
audio_url="https://cdn/v2.mp3",
|
||||
duration=30.5,
|
||||
file_size=102400,
|
||||
status="processing",
|
||||
project_id="proj1",
|
||||
tags=["产品", "介绍"],
|
||||
)
|
||||
assert item.text == "欢迎来到我们的产品"
|
||||
assert item.voice_provider == "cosyvoice"
|
||||
assert item.voice_id == "voice_001"
|
||||
assert item.audio_url == "https://cdn/v2.mp3"
|
||||
assert item.duration == 30.5
|
||||
assert item.file_size == 102400
|
||||
assert item.status == "processing"
|
||||
assert item.project_id == "proj1"
|
||||
assert item.tags == ["产品", "介绍"]
|
||||
|
||||
|
||||
class TestTitleLibraryItem:
|
||||
def test_create_minimal(self):
|
||||
item = TitleLibraryItem(
|
||||
id="t1",
|
||||
user_id="u1",
|
||||
name="爆款标题1",
|
||||
text="这是一个爆款标题",
|
||||
)
|
||||
assert item.id == "t1"
|
||||
assert item.user_id == "u1"
|
||||
assert item.name == "爆款标题1"
|
||||
assert item.text == "这是一个爆款标题"
|
||||
assert item.category == "default"
|
||||
assert item.description == ""
|
||||
assert item.tags == []
|
||||
assert item.usage_count == 0
|
||||
assert item.is_active is True
|
||||
|
||||
def test_create_with_category(self):
|
||||
item = TitleLibraryItem(
|
||||
id="t2",
|
||||
user_id="u1",
|
||||
name="美食标题",
|
||||
text="太好吃了!",
|
||||
category="美食",
|
||||
)
|
||||
assert item.category == "美食"
|
||||
|
||||
def test_inactive_item(self):
|
||||
item = TitleLibraryItem(
|
||||
id="t3",
|
||||
user_id="u1",
|
||||
name="旧标题",
|
||||
text="旧文案",
|
||||
is_active=False,
|
||||
)
|
||||
assert item.is_active is False
|
||||
|
||||
def test_usage_count_increment(self):
|
||||
item = TitleLibraryItem(
|
||||
id="t4",
|
||||
user_id="u1",
|
||||
name="T",
|
||||
text="T",
|
||||
)
|
||||
item.usage_count += 1
|
||||
assert item.usage_count == 1
|
||||
|
||||
|
||||
class TestRecipeItem:
|
||||
def test_create_minimal(self):
|
||||
item = RecipeItem(
|
||||
id="ri1",
|
||||
recipe_id="r1",
|
||||
item_type="asset",
|
||||
item_id="asset_001",
|
||||
)
|
||||
assert item.id == "ri1"
|
||||
assert item.recipe_id == "r1"
|
||||
assert item.item_type == "asset"
|
||||
assert item.item_id == "asset_001"
|
||||
assert item.position == 0
|
||||
assert item.metadata_ == {}
|
||||
|
||||
def test_create_with_position_and_metadata(self):
|
||||
item = RecipeItem(
|
||||
id="ri2",
|
||||
recipe_id="r1",
|
||||
item_type="title",
|
||||
item_id="title_001",
|
||||
position=2,
|
||||
metadata_={"style": "bold"},
|
||||
)
|
||||
assert item.position == 2
|
||||
assert item.metadata_ == {"style": "bold"}
|
||||
|
||||
def test_item_types_variety(self):
|
||||
asset_item = RecipeItem(id="a", recipe_id="r", item_type="asset", item_id="i1")
|
||||
title_item = RecipeItem(id="t", recipe_id="r", item_type="title", item_id="i2")
|
||||
voice_item = RecipeItem(id="v", recipe_id="r", item_type="voice", item_id="i3")
|
||||
assert asset_item.item_type == "asset"
|
||||
assert title_item.item_type == "title"
|
||||
assert voice_item.item_type == "voice"
|
||||
@@ -15,7 +15,6 @@ from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from worker_app.tasks.generation_plan_builder import (
|
||||
VirtualClip,
|
||||
VirtualPlan,
|
||||
|
||||
+335
-320
@@ -1,4 +1,6 @@
|
||||
"""Job 领域层单元测试 - job.py"""
|
||||
"""Job 领域模型单元测试。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -10,45 +12,45 @@ from packages.domain.job import (
|
||||
)
|
||||
|
||||
|
||||
class TestJobType:
|
||||
"""JobType 枚举测试"""
|
||||
class TestJobTypeEnum:
|
||||
def test_all_types_exist(self):
|
||||
assert JobType.VIDEO_COMPOSE.value == "video_compose"
|
||||
assert JobType.RENDER_EDIT_PLAN.value == "render_edit_plan"
|
||||
assert JobType.ASSET_INGEST.value == "asset_ingest"
|
||||
assert JobType.CLASSIFICATION.value == "classification"
|
||||
assert JobType.VOICE_EXTRACTION.value == "voice_extraction"
|
||||
assert JobType.GENERATION.value == "generation"
|
||||
|
||||
def test_all_types_have_values(self):
|
||||
"""所有枚举成员都有字符串值"""
|
||||
for jt in JobType:
|
||||
assert isinstance(jt.value, str)
|
||||
assert jt.value
|
||||
def test_from_string(self):
|
||||
assert JobType("video_compose") == JobType.VIDEO_COMPOSE
|
||||
assert JobType("generation") == JobType.GENERATION
|
||||
|
||||
def test_str_enum_behavior(self):
|
||||
"""是 str 枚举"""
|
||||
assert JobType.VIDEO_COMPOSE == "video_compose"
|
||||
assert isinstance(JobType.VIDEO_COMPOSE, str)
|
||||
|
||||
def test_known_types_exist(self):
|
||||
"""核心任务类型都存在"""
|
||||
assert JobType.VIDEO_COMPOSE
|
||||
assert JobType.RENDER_EDIT_PLAN
|
||||
assert JobType.ASSET_INGEST
|
||||
assert JobType.CLASSIFICATION
|
||||
assert JobType.GENERATION
|
||||
def test_invalid_type_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
JobType("invalid_type")
|
||||
|
||||
|
||||
class TestJobStatus:
|
||||
"""JobStatus 枚举测试"""
|
||||
class TestJobStatusEnum:
|
||||
def test_all_statuses_exist(self):
|
||||
assert JobStatus.PENDING.value == "pending"
|
||||
assert JobStatus.RUNNING.value == "running"
|
||||
assert JobStatus.SUCCESS.value == "success"
|
||||
assert JobStatus.FAILED.value == "failed"
|
||||
assert JobStatus.CANCELLED.value == "cancelled"
|
||||
|
||||
def test_all_statuses_have_values(self):
|
||||
for js in JobStatus:
|
||||
assert isinstance(js.value, str)
|
||||
assert js.value
|
||||
def test_from_string(self):
|
||||
assert JobStatus("pending") == JobStatus.PENDING
|
||||
assert JobStatus("success") == JobStatus.SUCCESS
|
||||
|
||||
def test_str_enum_behavior(self):
|
||||
assert JobStatus.PENDING == "pending"
|
||||
assert isinstance(JobStatus.PENDING, str)
|
||||
|
||||
def test_terminal_statuses(self):
|
||||
"""终态集合包含成功/失败/取消"""
|
||||
class TestTerminalStatuses:
|
||||
def test_success_is_terminal(self):
|
||||
assert JobStatus.SUCCESS in TERMINAL_STATUSES
|
||||
|
||||
def test_failed_is_terminal(self):
|
||||
assert JobStatus.FAILED in TERMINAL_STATUSES
|
||||
|
||||
def test_cancelled_is_terminal(self):
|
||||
assert JobStatus.CANCELLED in TERMINAL_STATUSES
|
||||
|
||||
def test_pending_not_terminal(self):
|
||||
@@ -59,372 +61,376 @@ class TestJobStatus:
|
||||
|
||||
|
||||
class TestJobCreate:
|
||||
"""Job.create 工厂方法测试"""
|
||||
|
||||
def test_create_basic(self):
|
||||
"""基本创建"""
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
)
|
||||
assert job.id
|
||||
assert len(job.id) == 32 # uuid4 hex
|
||||
assert job.project_id == "proj-1"
|
||||
def test_create_minimal(self):
|
||||
job = Job.create(project_id="proj1", job_type=JobType.VIDEO_COMPOSE)
|
||||
assert job.id is not None
|
||||
assert len(job.id) == 32
|
||||
assert job.project_id == "proj1"
|
||||
assert job.job_type == JobType.VIDEO_COMPOSE
|
||||
assert job.status == JobStatus.PENDING
|
||||
assert job.progress == 0.0
|
||||
assert job.current_stage == ""
|
||||
assert job.payload == {}
|
||||
assert job.result == {}
|
||||
assert job.error_message == ""
|
||||
assert job.retry_count == 0
|
||||
assert job.max_retries == 3
|
||||
assert job.created_at
|
||||
assert job.updated_at
|
||||
assert job.celery_task_id == ""
|
||||
assert job.source_id == ""
|
||||
assert job.created_by_user_id == ""
|
||||
assert job.started_at is None
|
||||
assert job.completed_at is None
|
||||
assert isinstance(job.created_at, datetime)
|
||||
assert isinstance(job.updated_at, datetime)
|
||||
|
||||
def test_create_with_string_job_type(self):
|
||||
"""用字符串创建任务类型"""
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type="video_compose",
|
||||
)
|
||||
def test_create_with_enum_type(self):
|
||||
job = Job.create("p1", JobType.GENERATION)
|
||||
assert job.job_type == JobType.GENERATION
|
||||
|
||||
def test_create_with_string_type(self):
|
||||
job = Job.create("p1", "video_compose")
|
||||
assert job.job_type == JobType.VIDEO_COMPOSE
|
||||
|
||||
def test_create_invalid_string_job_type_raises(self):
|
||||
"""无效的任务类型字符串抛 ValueError"""
|
||||
with pytest.raises(ValueError, match="不支持的任务类型"):
|
||||
Job.create(project_id="proj-1", job_type="invalid_type")
|
||||
|
||||
def test_create_empty_project_id_raises(self):
|
||||
"""空 project_id 抛 ValueError"""
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
Job.create(project_id=" ", job_type=JobType.VIDEO_COMPOSE)
|
||||
|
||||
def test_create_with_payload(self):
|
||||
"""带 payload 创建"""
|
||||
payload = {"video_id": "v1", "quality": "1080p"}
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
payload=payload,
|
||||
)
|
||||
payload = {"edit_plan_id": "plan123", "resolution": "1080p"}
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, payload=payload)
|
||||
assert job.payload == payload
|
||||
|
||||
def test_create_with_source_id(self):
|
||||
"""带 source_id 创建"""
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
source_id="plan-123",
|
||||
)
|
||||
assert job.source_id == "plan-123"
|
||||
|
||||
def test_create_with_created_by(self):
|
||||
"""带创建人"""
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
assert job.created_by_user_id == "user-1"
|
||||
|
||||
def test_create_with_custom_max_retries(self):
|
||||
"""自定义最大重试次数"""
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
max_retries=5,
|
||||
)
|
||||
assert job.max_retries == 5
|
||||
|
||||
def test_create_project_id_stripped(self):
|
||||
"""project_id 会被 strip"""
|
||||
job = Job.create(
|
||||
project_id=" proj-1 ",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
)
|
||||
assert job.project_id == "proj-1"
|
||||
|
||||
def test_create_source_id_stripped(self):
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
source_id=" src-1 ",
|
||||
)
|
||||
assert job.source_id == "src-1"
|
||||
|
||||
def test_create_created_by_stripped(self):
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
created_by_user_id=" user-1 ",
|
||||
)
|
||||
assert job.created_by_user_id == "user-1"
|
||||
|
||||
def test_create_none_payload_defaults_to_empty_dict(self):
|
||||
"""payload=None 时默认为空 dict"""
|
||||
job = Job.create(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
payload=None,
|
||||
)
|
||||
def test_create_with_none_payload(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, payload=None)
|
||||
assert job.payload == {}
|
||||
|
||||
def test_create_with_source_id(self):
|
||||
job = Job.create("p1", JobType.GENERATION, source_id="gen123")
|
||||
assert job.source_id == "gen123"
|
||||
|
||||
class TestJobIsTerminal:
|
||||
"""is_terminal 属性测试"""
|
||||
def test_create_with_user_id(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, created_by_user_id="user1")
|
||||
assert job.created_by_user_id == "user1"
|
||||
|
||||
def test_create_with_custom_max_retries(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=5)
|
||||
assert job.max_retries == 5
|
||||
|
||||
def test_create_strips_project_id(self):
|
||||
job = Job.create(" proj1 ", JobType.VIDEO_COMPOSE)
|
||||
assert job.project_id == "proj1"
|
||||
|
||||
def test_create_strips_source_id(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, source_id=" src1 ")
|
||||
assert job.source_id == "src1"
|
||||
|
||||
def test_create_strips_user_id(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, created_by_user_id=" u1 ")
|
||||
assert job.created_by_user_id == "u1"
|
||||
|
||||
def test_create_empty_project_id(self):
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
Job.create("", JobType.VIDEO_COMPOSE)
|
||||
|
||||
def test_create_whitespace_project_id(self):
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
Job.create(" \t ", JobType.VIDEO_COMPOSE)
|
||||
|
||||
def test_create_invalid_job_type_string(self):
|
||||
with pytest.raises(ValueError, match="不支持的任务类型"):
|
||||
Job.create("p1", "invalid_type")
|
||||
|
||||
def test_create_unique_ids(self):
|
||||
j1 = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
j2 = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
assert j1.id != j2.id
|
||||
|
||||
|
||||
class TestIsTerminal:
|
||||
def test_pending_not_terminal(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
assert job.is_terminal is False
|
||||
|
||||
def test_running_not_terminal(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
assert job.is_terminal is False
|
||||
|
||||
def test_success_is_terminal(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.SUCCESS)
|
||||
assert job.is_terminal is True
|
||||
|
||||
def test_failed_is_terminal(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
assert job.is_terminal is True
|
||||
|
||||
def test_cancelled_is_terminal(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.CANCELLED)
|
||||
assert job.is_terminal is True
|
||||
|
||||
|
||||
class TestJobTransitions:
|
||||
"""状态转换测试"""
|
||||
class TestIsRetryable:
|
||||
def test_pending_not_retryable(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_running_not_retryable(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_success_not_retryable(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_success()
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_failed_within_limit_is_retryable(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
job.mark_running()
|
||||
job.mark_failed("error")
|
||||
assert job.is_retryable is True
|
||||
|
||||
def test_failed_at_limit_not_retryable(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
job.mark_running()
|
||||
job.mark_failed("error")
|
||||
job.retry_count = 3 # 已达到上限
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_failed_over_limit_not_retryable(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
job.retry_count = 5
|
||||
job.status = JobStatus.FAILED
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_zero_max_retries_not_retryable(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=0)
|
||||
job.status = JobStatus.FAILED
|
||||
assert job.is_retryable is False
|
||||
|
||||
|
||||
class TestTransitionTo:
|
||||
def test_pending_to_running(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
assert job.status == JobStatus.RUNNING
|
||||
assert job.started_at is not None
|
||||
|
||||
def test_pending_to_success(self):
|
||||
"""pending 可以直接到 success(快速成功)"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.SUCCESS)
|
||||
assert job.status == JobStatus.SUCCESS
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_pending_to_cancelled(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.CANCELLED)
|
||||
assert job.status == JobStatus.CANCELLED
|
||||
|
||||
def test_running_to_success(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.SUCCESS)
|
||||
assert job.status == JobStatus.SUCCESS
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_running_to_failed(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
assert job.status == JobStatus.FAILED
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_running_to_cancelled(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.CANCELLED)
|
||||
assert job.status == JobStatus.CANCELLED
|
||||
|
||||
def test_failed_to_pending_retry(self):
|
||||
"""失败后可以回到 pending(重试)"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
job.transition_to(JobStatus.PENDING)
|
||||
assert job.status == JobStatus.PENDING
|
||||
|
||||
def test_invalid_transition_raises(self):
|
||||
"""非法状态转换抛 ValueError"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
# pending 不能直接到 failed
|
||||
def test_pending_to_failed_invalid(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
with pytest.raises(ValueError, match="非法状态转换"):
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
|
||||
def test_success_to_pending_raises(self):
|
||||
"""成功后不能回到 pending"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_running_to_success(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.SUCCESS)
|
||||
assert job.status == JobStatus.SUCCESS
|
||||
|
||||
def test_running_to_failed(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
assert job.status == JobStatus.FAILED
|
||||
|
||||
def test_running_to_cancelled(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.CANCELLED)
|
||||
assert job.status == JobStatus.CANCELLED
|
||||
|
||||
def test_running_to_pending_invalid(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
with pytest.raises(ValueError, match="非法状态转换"):
|
||||
job.transition_to(JobStatus.PENDING)
|
||||
|
||||
def test_failed_to_pending(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
# 注意:_VALID_TRANSITIONS 中 FAILED → PENDING 是允许的
|
||||
job.transition_to(JobStatus.PENDING)
|
||||
assert job.status == JobStatus.PENDING
|
||||
|
||||
def test_success_to_anything_invalid(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
job.transition_to(JobStatus.SUCCESS)
|
||||
with pytest.raises(ValueError):
|
||||
job.transition_to(JobStatus.PENDING)
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
with pytest.raises(ValueError):
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
|
||||
def test_transition_with_string_status(self):
|
||||
"""用字符串做状态转换"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to("running")
|
||||
assert job.status == JobStatus.RUNNING
|
||||
|
||||
def test_transition_invalid_string_raises(self):
|
||||
"""无效状态字符串抛 ValueError"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_transition_with_invalid_string(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
with pytest.raises(ValueError, match="无效状态"):
|
||||
job.transition_to("invalid_status")
|
||||
|
||||
def test_transition_updates_updated_at(self):
|
||||
"""状态转换更新 updated_at"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
old_updated = job.updated_at
|
||||
import time
|
||||
|
||||
time.sleep(0.001)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
old_time = job.updated_at
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
assert job.updated_at >= old_updated
|
||||
assert job.updated_at >= old_time
|
||||
|
||||
def test_started_at_only_set_once(self):
|
||||
"""started_at 只在第一次 RUNNING 时设置"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
first_started = job.started_at
|
||||
job.transition_to(JobStatus.SUCCESS)
|
||||
# 回到 pending 再 running(模拟重试场景,但started_at是None时才设置)
|
||||
# 注意:正常重试是通过 prepare_retry 重置的
|
||||
assert first_started is not None
|
||||
first_start = job.started_at
|
||||
# 再次 RUNNING 不合法,但我们测试 started_at 在多次 running→success→retry→running 时的行为
|
||||
# 先失败重试
|
||||
job.transition_to(JobStatus.FAILED)
|
||||
job.transition_to(JobStatus.PENDING)
|
||||
job.started_at = None # 模拟 prepare_retry 的重置
|
||||
job.transition_to(JobStatus.RUNNING)
|
||||
assert job.started_at is not None
|
||||
assert job.started_at != first_start
|
||||
|
||||
|
||||
class TestJobMarkMethods:
|
||||
"""便捷标记方法测试"""
|
||||
|
||||
def test_mark_running(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.mark_running("合成中")
|
||||
assert job.status == JobStatus.RUNNING
|
||||
assert job.current_stage == "合成中"
|
||||
|
||||
def test_mark_running_no_stage(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
class TestMarkRunning:
|
||||
def test_mark_running_basic(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
assert job.status == JobStatus.RUNNING
|
||||
assert job.current_stage == ""
|
||||
assert job.started_at is not None
|
||||
|
||||
def test_mark_success(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_success({"output_url": "http://..."})
|
||||
assert job.status == JobStatus.SUCCESS
|
||||
assert job.progress == 100.0
|
||||
assert job.current_stage == "完成"
|
||||
assert job.result == {"output_url": "http://..."}
|
||||
def test_mark_running_with_stage(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running(stage="下载素材")
|
||||
assert job.status == JobStatus.RUNNING
|
||||
assert job.current_stage == "下载素材"
|
||||
|
||||
def test_mark_success_no_result(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_mark_running_empty_stage_unchanged(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.current_stage = "已有阶段"
|
||||
job.mark_running() # 不传 stage
|
||||
assert job.current_stage == "已有阶段"
|
||||
|
||||
|
||||
class TestMarkSuccess:
|
||||
def test_mark_success_basic(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_success()
|
||||
assert job.status == JobStatus.SUCCESS
|
||||
assert job.result == {}
|
||||
assert job.progress == 100.0
|
||||
assert job.current_stage == "完成"
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_mark_failed(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_mark_success_with_result(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
result = {"video_url": "https://...", "duration": 30}
|
||||
job.mark_success(result=result)
|
||||
assert job.result == result
|
||||
|
||||
def test_mark_success_without_result(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
original_result = job.result.copy()
|
||||
job.mark_success()
|
||||
assert job.result == original_result # 不变
|
||||
|
||||
|
||||
class TestMarkFailed:
|
||||
def test_mark_failed_basic(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_failed("网络超时")
|
||||
assert job.status == JobStatus.FAILED
|
||||
assert job.error_message == "网络超时"
|
||||
assert job.current_stage == "失败"
|
||||
assert job.completed_at is not None
|
||||
|
||||
def test_mark_cancelled(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_mark_failed_empty_message(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_failed("")
|
||||
assert job.error_message == ""
|
||||
|
||||
|
||||
class TestMarkCancelled:
|
||||
def test_mark_cancelled_from_pending(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_cancelled()
|
||||
assert job.status == JobStatus.CANCELLED
|
||||
assert job.current_stage == "已取消"
|
||||
|
||||
def test_mark_cancelled_from_running(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_cancelled()
|
||||
assert job.status == JobStatus.CANCELLED
|
||||
|
||||
class TestJobProgress:
|
||||
"""进度更新测试"""
|
||||
|
||||
def test_update_progress(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.update_progress(50.0, "渲染中")
|
||||
class TestUpdateProgress:
|
||||
def test_update_progress_valid(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.update_progress(50.0)
|
||||
assert job.progress == 50.0
|
||||
assert job.current_stage == "渲染中"
|
||||
|
||||
def test_update_progress_zero(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.update_progress(0.0)
|
||||
assert job.progress == 0.0
|
||||
|
||||
def test_update_progress_100(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_update_progress_hundred(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.update_progress(100.0)
|
||||
assert job.progress == 100.0
|
||||
|
||||
def test_update_progress_negative_raises(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_update_progress_negative(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
with pytest.raises(ValueError, match="进度必须在 0~100 之间"):
|
||||
job.update_progress(-1.0)
|
||||
|
||||
def test_update_progress_over_100_raises(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
def test_update_progress_over_100(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
with pytest.raises(ValueError, match="进度必须在 0~100 之间"):
|
||||
job.update_progress(101.0)
|
||||
|
||||
def test_update_progress_without_stage(self):
|
||||
"""不传 stage 时不修改 current_stage"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.current_stage = "初始阶段"
|
||||
job.update_progress(30.0)
|
||||
def test_update_progress_with_stage(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.update_progress(30.0, stage="渲染中")
|
||||
assert job.progress == 30.0
|
||||
assert job.current_stage == "初始阶段"
|
||||
assert job.current_stage == "渲染中"
|
||||
|
||||
def test_update_progress_updates_updated_at(self):
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
old_updated = job.updated_at
|
||||
import time
|
||||
|
||||
time.sleep(0.001)
|
||||
def test_update_progress_without_stage_unchanged(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.current_stage = "原阶段"
|
||||
job.update_progress(50.0)
|
||||
assert job.updated_at >= old_updated
|
||||
assert job.current_stage == "原阶段"
|
||||
|
||||
def test_update_progress_updates_timestamp(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
old_time = job.updated_at
|
||||
job.update_progress(25.0)
|
||||
assert job.updated_at >= old_time
|
||||
|
||||
|
||||
class TestJobRetry:
|
||||
"""重试逻辑测试"""
|
||||
|
||||
def test_is_retryable_failed_within_limit(self):
|
||||
"""失败且未超过重试上限时可重试"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
job.mark_running()
|
||||
job.mark_failed("错误")
|
||||
assert job.is_retryable is True
|
||||
|
||||
def test_is_retryable_failed_at_limit(self):
|
||||
"""达到重试上限时不可重试"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE, max_retries=1)
|
||||
job.mark_running()
|
||||
job.mark_failed("错误")
|
||||
job.retry_count = 1
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_is_retryable_pending_false(self):
|
||||
"""pending 状态不可重试"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_is_retryable_success_false(self):
|
||||
"""成功状态不可重试"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_success()
|
||||
assert job.is_retryable is False
|
||||
|
||||
def test_prepare_retry(self):
|
||||
"""准备重试"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
class TestPrepareRetry:
|
||||
def test_prepare_retry_success(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
job.mark_running()
|
||||
job.mark_failed("网络错误")
|
||||
job.celery_task_id = "task-123"
|
||||
|
||||
job.prepare_retry()
|
||||
|
||||
@@ -437,38 +443,41 @@ class TestJobRetry:
|
||||
assert job.completed_at is None
|
||||
assert job.celery_task_id == ""
|
||||
|
||||
def test_prepare_retry_not_retryable_raises(self):
|
||||
"""不可重试时抛 ValueError"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE, max_retries=0)
|
||||
def test_prepare_retry_increments_count(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=5)
|
||||
job.mark_running()
|
||||
job.mark_failed("错误")
|
||||
with pytest.raises(ValueError, match="任务不可重试"):
|
||||
job.prepare_retry()
|
||||
job.mark_failed("err")
|
||||
|
||||
def test_prepare_retry_increments_correctly(self):
|
||||
"""多次重试计数正确"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
job.mark_running()
|
||||
job.mark_failed("错误1")
|
||||
job.prepare_retry()
|
||||
assert job.retry_count == 1
|
||||
|
||||
# 再次失败重试
|
||||
job.mark_running()
|
||||
job.mark_failed("错误2")
|
||||
job.mark_failed("err2")
|
||||
job.prepare_retry()
|
||||
assert job.retry_count == 2
|
||||
|
||||
def test_prepare_retry_not_retryable_raises(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=0)
|
||||
job.mark_running()
|
||||
job.mark_failed("err")
|
||||
with pytest.raises(ValueError, match="任务不可重试"):
|
||||
job.prepare_retry()
|
||||
|
||||
class TestJobToDict:
|
||||
"""to_dict 序列化测试"""
|
||||
def test_prepare_retry_wrong_status_raises(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
with pytest.raises(ValueError, match="任务不可重试"):
|
||||
job.prepare_retry()
|
||||
|
||||
def test_to_dict_contains_all_fields(self):
|
||||
|
||||
class TestToDict:
|
||||
def test_to_dict_structure(self):
|
||||
job = Job.create(
|
||||
project_id="p1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
payload={"key": "value"},
|
||||
source_id="src-1",
|
||||
created_by_user_id="user-1",
|
||||
"p1",
|
||||
JobType.VIDEO_COMPOSE,
|
||||
payload={"key": "val"},
|
||||
source_id="src1",
|
||||
created_by_user_id="u1",
|
||||
)
|
||||
d = job.to_dict()
|
||||
assert d["id"] == job.id
|
||||
@@ -476,33 +485,39 @@ class TestJobToDict:
|
||||
assert d["job_type"] == "video_compose"
|
||||
assert d["status"] == "pending"
|
||||
assert d["progress"] == 0.0
|
||||
assert d["payload"] == {"key": "value"}
|
||||
assert d["source_id"] == "src-1"
|
||||
assert d["created_by_user_id"] == "user-1"
|
||||
assert d["current_stage"] == ""
|
||||
assert d["payload"] == {"key": "val"}
|
||||
assert d["result"] == {}
|
||||
assert d["error_message"] == ""
|
||||
assert d["retry_count"] == 0
|
||||
assert d["max_retries"] == 3
|
||||
assert d["celery_task_id"] == ""
|
||||
assert d["source_id"] == "src1"
|
||||
assert d["created_by_user_id"] == "u1"
|
||||
assert d["is_retryable"] is False
|
||||
|
||||
def test_to_dict_datetime_fields_are_strings(self):
|
||||
"""时间字段序列化为 ISO 字符串"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
d = job.to_dict()
|
||||
assert isinstance(d["created_at"], str)
|
||||
assert isinstance(d["updated_at"], str)
|
||||
|
||||
def test_to_dict_none_datetime_fields(self):
|
||||
"""未设置的时间字段为 None"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
d = job.to_dict()
|
||||
assert d["started_at"] is None
|
||||
assert d["completed_at"] is None
|
||||
assert d["created_at"] is not None
|
||||
assert d["updated_at"] is not None
|
||||
|
||||
def test_to_dict_after_success(self):
|
||||
"""成功后 to_dict 状态正确"""
|
||||
job = Job.create(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
job.mark_running()
|
||||
job.mark_success({"url": "http://..."})
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE)
|
||||
job.mark_running("渲染")
|
||||
job.mark_success({"url": "https://..."})
|
||||
d = job.to_dict()
|
||||
assert d["status"] == "success"
|
||||
assert d["progress"] == 100.0
|
||||
assert d["result"] == {"url": "http://..."}
|
||||
assert d["is_retryable"] is False
|
||||
assert d["started_at"] is not None
|
||||
assert d["completed_at"] is not None
|
||||
assert isinstance(d["started_at"], str)
|
||||
assert isinstance(d["completed_at"], str)
|
||||
|
||||
def test_to_dict_after_failed(self):
|
||||
job = Job.create("p1", JobType.VIDEO_COMPOSE, max_retries=3)
|
||||
job.mark_running()
|
||||
job.mark_failed("timeout")
|
||||
d = job.to_dict()
|
||||
assert d["status"] == "failed"
|
||||
assert d["error_message"] == "timeout"
|
||||
assert d["is_retryable"] is True
|
||||
|
||||
Executable
+299
@@ -0,0 +1,299 @@
|
||||
"""media_validation 领域模块单元测试。"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.media_validation import (
|
||||
MIN_AUDIO_FILE_SIZE,
|
||||
MIN_IMAGE_FILE_SIZE,
|
||||
MIN_VIDEO_FILE_SIZE,
|
||||
SUPPORTED_VIDEO_CODECS,
|
||||
is_valid_media,
|
||||
safe_parse_fps,
|
||||
)
|
||||
|
||||
|
||||
class TestSafeParseFpsBasic:
|
||||
def test_integer_fps(self):
|
||||
assert safe_parse_fps("30") == 30.0
|
||||
|
||||
def test_decimal_fps(self):
|
||||
assert safe_parse_fps("29.97") == pytest.approx(29.97)
|
||||
|
||||
def test_fraction_simple(self):
|
||||
assert safe_parse_fps("30/1") == 30.0
|
||||
|
||||
def test_fraction_ntsc(self):
|
||||
assert safe_parse_fps("30000/1001") == pytest.approx(29.97002997)
|
||||
|
||||
def test_fraction_pal(self):
|
||||
assert safe_parse_fps("25/1") == 25.0
|
||||
|
||||
def test_fraction_24fps_cine(self):
|
||||
assert safe_parse_fps("24000/1001") == pytest.approx(23.976023976)
|
||||
|
||||
def test_zero_fps(self):
|
||||
assert safe_parse_fps("0") == 0.0
|
||||
|
||||
def test_zero_fraction(self):
|
||||
assert safe_parse_fps("0/1") == 0.0
|
||||
|
||||
|
||||
class TestSafeParseFpsEdgeCases:
|
||||
def test_zero_denominator(self):
|
||||
assert safe_parse_fps("30/0") == 0.0
|
||||
|
||||
def test_empty_string(self):
|
||||
assert safe_parse_fps("") == 0.0
|
||||
|
||||
def test_garbage_string(self):
|
||||
assert safe_parse_fps("not_a_number") == 0.0
|
||||
|
||||
def test_multiple_slashes(self):
|
||||
# split("/", 1) 只切第一个,后面的作为 den 的一部分会解析失败
|
||||
assert safe_parse_fps("30/1/2") == 0.0
|
||||
|
||||
def test_negative_fps(self):
|
||||
assert safe_parse_fps("-30") == -30.0
|
||||
|
||||
def test_negative_fraction(self):
|
||||
assert safe_parse_fps("-30/1") == -30.0
|
||||
|
||||
def test_very_high_fps(self):
|
||||
assert safe_parse_fps("240/1") == 240.0
|
||||
|
||||
def test_fraction_float_num(self):
|
||||
assert safe_parse_fps("29.97/1") == pytest.approx(29.97)
|
||||
|
||||
def test_fraction_float_den(self):
|
||||
assert safe_parse_fps("30/1.001") == pytest.approx(29.97002997)
|
||||
|
||||
def test_whitespace_in_string(self):
|
||||
# float(" 30 ") 能解析,所以应该返回 30.0
|
||||
assert safe_parse_fps(" 30 ") == 30.0
|
||||
|
||||
|
||||
class TestMinFileSizeConstants:
|
||||
def test_min_video_size_is_1kb(self):
|
||||
assert MIN_VIDEO_FILE_SIZE == 1024
|
||||
|
||||
def test_min_audio_size(self):
|
||||
assert MIN_AUDIO_FILE_SIZE == 100
|
||||
|
||||
def test_min_image_size(self):
|
||||
assert MIN_IMAGE_FILE_SIZE == 100
|
||||
|
||||
|
||||
class TestSupportedVideoCodecs:
|
||||
def test_h264_family_present(self):
|
||||
assert "h264" in SUPPORTED_VIDEO_CODECS
|
||||
assert "avc1" in SUPPORTED_VIDEO_CODECS
|
||||
assert "avc" in SUPPORTED_VIDEO_CODECS
|
||||
|
||||
def test_h265_family_present(self):
|
||||
assert "hevc" in SUPPORTED_VIDEO_CODECS
|
||||
assert "h265" in SUPPORTED_VIDEO_CODECS
|
||||
assert "hev1" in SUPPORTED_VIDEO_CODECS
|
||||
assert "hvc1" in SUPPORTED_VIDEO_CODECS
|
||||
|
||||
def test_vp9_av1_present(self):
|
||||
assert "vp9" in SUPPORTED_VIDEO_CODECS
|
||||
assert "vp09" in SUPPORTED_VIDEO_CODECS
|
||||
assert "av1" in SUPPORTED_VIDEO_CODECS
|
||||
assert "av01" in SUPPORTED_VIDEO_CODECS
|
||||
|
||||
def test_vp8_present(self):
|
||||
assert "vp8" in SUPPORTED_VIDEO_CODECS
|
||||
assert "vp08" in SUPPORTED_VIDEO_CODECS
|
||||
|
||||
def test_mpeg_family_present(self):
|
||||
assert "mpeg4" in SUPPORTED_VIDEO_CODECS
|
||||
assert "mp4v" in SUPPORTED_VIDEO_CODECS
|
||||
assert "mpeg2video" in SUPPORTED_VIDEO_CODECS
|
||||
|
||||
def test_prores_family_present(self):
|
||||
assert "prores" in SUPPORTED_VIDEO_CODECS
|
||||
assert "apcn" in SUPPORTED_VIDEO_CODECS
|
||||
assert "apch" in SUPPORTED_VIDEO_CODECS
|
||||
|
||||
def test_unknown_codec_not_present(self):
|
||||
assert "unknown_codec_xyz" not in SUPPORTED_VIDEO_CODECS
|
||||
|
||||
def test_codecs_count_reasonable(self):
|
||||
# 白名单应该有足够多的编码格式
|
||||
assert len(SUPPORTED_VIDEO_CODECS) >= 30
|
||||
|
||||
|
||||
class TestIsValidMediaVideo:
|
||||
def test_valid_video(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_too_small(self):
|
||||
metadata = {"size_bytes": 500, "duration": 10.0, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is False
|
||||
|
||||
def test_video_exact_min_size(self):
|
||||
metadata = {"size_bytes": 1024, "duration": 10.0, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_zero_duration(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 0, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is False
|
||||
|
||||
def test_video_negative_duration(self):
|
||||
metadata = {"size_bytes": 5000, "duration": -1.0, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is False
|
||||
|
||||
def test_video_missing_size_default_zero(self):
|
||||
metadata = {"duration": 10.0, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is False
|
||||
|
||||
def test_video_missing_duration_default_zero(self):
|
||||
metadata = {"size_bytes": 5000, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is False
|
||||
|
||||
def test_video_unknown_codec_still_valid(self):
|
||||
# 非白名单编码仍允许通过(渲染层统一转码)
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": "some_unknown_codec"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_missing_codec_still_valid(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_codec_case_insensitive(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": "H264"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_empty_codec(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": ""}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_hevc_codec(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": "hevc"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_vp9_codec(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": "vp9"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_av1_codec(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": "av1"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_prores_codec(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0, "codec": "prores"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_video_empty_metadata(self):
|
||||
assert is_valid_media({}, "video") is False
|
||||
|
||||
|
||||
class TestIsValidMediaAudio:
|
||||
def test_valid_audio(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 30.0}
|
||||
assert is_valid_media(metadata, "audio") is True
|
||||
|
||||
def test_audio_too_small(self):
|
||||
metadata = {"size_bytes": 50, "duration": 30.0}
|
||||
assert is_valid_media(metadata, "audio") is False
|
||||
|
||||
def test_audio_exact_min_size(self):
|
||||
metadata = {"size_bytes": 100, "duration": 10.0}
|
||||
assert is_valid_media(metadata, "audio") is True
|
||||
|
||||
def test_audio_zero_duration(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 0}
|
||||
assert is_valid_media(metadata, "audio") is False
|
||||
|
||||
def test_audio_negative_duration(self):
|
||||
metadata = {"size_bytes": 5000, "duration": -1.0}
|
||||
assert is_valid_media(metadata, "audio") is False
|
||||
|
||||
def test_audio_missing_size(self):
|
||||
metadata = {"duration": 10.0}
|
||||
assert is_valid_media(metadata, "audio") is False
|
||||
|
||||
def test_audio_missing_duration(self):
|
||||
metadata = {"size_bytes": 5000}
|
||||
assert is_valid_media(metadata, "audio") is False
|
||||
|
||||
def test_audio_with_codec_info(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 30.0, "codec": "aac"}
|
||||
assert is_valid_media(metadata, "audio") is True
|
||||
|
||||
def test_audio_empty_metadata(self):
|
||||
assert is_valid_media({}, "audio") is False
|
||||
|
||||
|
||||
class TestIsValidMediaImage:
|
||||
def test_valid_image(self):
|
||||
metadata = {"size_bytes": 5000, "width": 1920, "height": 1080}
|
||||
assert is_valid_media(metadata, "image") is True
|
||||
|
||||
def test_image_too_small(self):
|
||||
metadata = {"size_bytes": 50, "width": 1920, "height": 1080}
|
||||
assert is_valid_media(metadata, "image") is False
|
||||
|
||||
def test_image_exact_min_size(self):
|
||||
metadata = {"size_bytes": 100, "width": 100, "height": 100}
|
||||
assert is_valid_media(metadata, "image") is True
|
||||
|
||||
def test_image_zero_width(self):
|
||||
metadata = {"size_bytes": 5000, "width": 0, "height": 1080}
|
||||
assert is_valid_media(metadata, "image") is False
|
||||
|
||||
def test_image_zero_height(self):
|
||||
metadata = {"size_bytes": 5000, "width": 1920, "height": 0}
|
||||
assert is_valid_media(metadata, "image") is False
|
||||
|
||||
def test_image_negative_dimensions(self):
|
||||
metadata = {"size_bytes": 5000, "width": -1, "height": 1080}
|
||||
assert is_valid_media(metadata, "image") is False
|
||||
|
||||
def test_image_missing_width(self):
|
||||
metadata = {"size_bytes": 5000, "height": 1080}
|
||||
assert is_valid_media(metadata, "image") is False
|
||||
|
||||
def test_image_missing_height(self):
|
||||
metadata = {"size_bytes": 5000, "width": 1920}
|
||||
assert is_valid_media(metadata, "image") is False
|
||||
|
||||
def test_image_missing_size(self):
|
||||
metadata = {"width": 1920, "height": 1080}
|
||||
assert is_valid_media(metadata, "image") is False
|
||||
|
||||
def test_image_small_but_valid(self):
|
||||
metadata = {"size_bytes": 100, "width": 1, "height": 1}
|
||||
assert is_valid_media(metadata, "image") is True
|
||||
|
||||
def test_image_empty_metadata(self):
|
||||
assert is_valid_media({}, "image") is False
|
||||
|
||||
|
||||
class TestIsValidMediaUnknownType:
|
||||
def test_unknown_type_returns_false(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0}
|
||||
assert is_valid_media(metadata, "unknown") is False
|
||||
|
||||
def test_empty_type_returns_false(self):
|
||||
metadata = {"size_bytes": 5000, "duration": 10.0}
|
||||
assert is_valid_media(metadata, "") is False
|
||||
|
||||
def test_text_type_returns_false(self):
|
||||
metadata = {"size_bytes": 5000}
|
||||
assert is_valid_media(metadata, "text") is False
|
||||
|
||||
|
||||
class TestIsValidMediaSizeTypes:
|
||||
def test_size_as_string(self):
|
||||
# int("5000") 能解析
|
||||
metadata = {"size_bytes": "5000", "duration": 10.0, "codec": "h264"}
|
||||
assert is_valid_media(metadata, "video") is True
|
||||
|
||||
def test_size_as_none(self):
|
||||
# int(None) 会 TypeError,但 metadata.get 返回 0 默认值
|
||||
metadata = {"size_bytes": None, "duration": 10.0, "codec": "h264"}
|
||||
# int(None) 会抛 TypeError
|
||||
with pytest.raises(TypeError):
|
||||
is_valid_media(metadata, "video")
|
||||
@@ -17,7 +17,6 @@ from packages.domain.plan_generator_utils import (
|
||||
)
|
||||
from packages.domain.template_clip_config import ClipType, TemplateClipConfig
|
||||
|
||||
|
||||
# ── 辅助函数 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -184,11 +183,7 @@ class TestDistributeVoicePip:
|
||||
|
||||
def test_three_assets_full_distribution(self):
|
||||
"""3个素材:background + corner_voice + b_roll 各一个."""
|
||||
clips = (
|
||||
_make_clips(1, "background")
|
||||
+ _make_clips(1, "corner_voice")
|
||||
+ _make_clips(1, "b_roll")
|
||||
)
|
||||
clips = _make_clips(1, "background") + _make_clips(1, "corner_voice") + _make_clips(1, "b_roll")
|
||||
distribute_assets(clips, ["a1", "a2", "a3"], EditingMode.VOICE_PIP.value)
|
||||
bgs = [c for c in clips if c.clip_type == "background"]
|
||||
voices = [c for c in clips if c.clip_type == "corner_voice"]
|
||||
@@ -199,11 +194,7 @@ class TestDistributeVoicePip:
|
||||
|
||||
def test_single_asset_only_background(self):
|
||||
"""1个素材:只分配给 background."""
|
||||
clips = (
|
||||
_make_clips(1, "background")
|
||||
+ _make_clips(1, "corner_voice")
|
||||
+ _make_clips(2, "b_roll")
|
||||
)
|
||||
clips = _make_clips(1, "background") + _make_clips(1, "corner_voice") + _make_clips(2, "b_roll")
|
||||
distribute_assets(clips, ["a1"], EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == ""
|
||||
@@ -212,11 +203,7 @@ class TestDistributeVoicePip:
|
||||
|
||||
def test_two_assets_bg_and_voice(self):
|
||||
"""2个素材:background + corner_voice."""
|
||||
clips = (
|
||||
_make_clips(1, "background")
|
||||
+ _make_clips(1, "corner_voice")
|
||||
+ _make_clips(2, "b_roll")
|
||||
)
|
||||
clips = _make_clips(1, "background") + _make_clips(1, "corner_voice") + _make_clips(2, "b_roll")
|
||||
distribute_assets(clips, ["a1", "a2"], EditingMode.VOICE_PIP.value)
|
||||
bgs = [c for c in clips if c.clip_type == "background"]
|
||||
voices = [c for c in clips if c.clip_type == "corner_voice"]
|
||||
@@ -227,11 +214,7 @@ class TestDistributeVoicePip:
|
||||
|
||||
def test_many_broll_clips(self):
|
||||
"""多个 b_roll clip:按顺序分配剩余素材."""
|
||||
clips = (
|
||||
_make_clips(1, "background")
|
||||
+ _make_clips(1, "corner_voice")
|
||||
+ _make_clips(5, "b_roll")
|
||||
)
|
||||
clips = _make_clips(1, "background") + _make_clips(1, "corner_voice") + _make_clips(5, "b_roll")
|
||||
distribute_assets(
|
||||
clips,
|
||||
["a1", "a2", "a3", "a4", "a5"],
|
||||
@@ -337,11 +320,7 @@ class TestMapClipTypesForMode:
|
||||
|
||||
def test_non_main_clips_unchanged(self):
|
||||
"""非 MAIN 类型 clip 不受影响."""
|
||||
clips = (
|
||||
_make_clips(1, "intro")
|
||||
+ _make_clips(3) # main
|
||||
+ _make_clips(1, "outro")
|
||||
)
|
||||
clips = _make_clips(1, "intro") + _make_clips(3) + _make_clips(1, "outro") # main
|
||||
map_clip_types_for_mode(clips, EditingMode.PIP.value)
|
||||
assert clips[0].clip_type == "intro"
|
||||
assert clips[1].clip_type == "main" # 第1个 main
|
||||
|
||||
Executable
+201
@@ -0,0 +1,201 @@
|
||||
"""Preset BGM 预设背景音乐单元测试。"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.preset_bgm import (
|
||||
BGM_STYLES,
|
||||
PRESET_BGM_LIBRARY,
|
||||
PresetBGM,
|
||||
get_preset_bgm,
|
||||
list_preset_bgm_by_style,
|
||||
search_preset_bgm,
|
||||
)
|
||||
|
||||
|
||||
class TestPresetBGMDataclass:
|
||||
def test_creation_required_fields(self):
|
||||
bgm = PresetBGM(id="test_001", name="Test BGM", style="upbeat", duration=120.0)
|
||||
assert bgm.id == "test_001"
|
||||
assert bgm.name == "Test BGM"
|
||||
assert bgm.style == "upbeat"
|
||||
assert bgm.duration == 120.0
|
||||
assert bgm.artist == ""
|
||||
assert bgm.description == ""
|
||||
assert bgm.tags == []
|
||||
assert bgm.audio_url == ""
|
||||
|
||||
def test_creation_all_fields(self):
|
||||
bgm = PresetBGM(
|
||||
id="test_002",
|
||||
name="Full BGM",
|
||||
style="relax",
|
||||
duration=180.5,
|
||||
artist="Artist Name",
|
||||
description="A test description",
|
||||
tags=["tag1", "tag2"],
|
||||
audio_url="https://cdn/test.mp3",
|
||||
)
|
||||
assert bgm.artist == "Artist Name"
|
||||
assert bgm.description == "A test description"
|
||||
assert bgm.tags == ["tag1", "tag2"]
|
||||
assert bgm.audio_url == "https://cdn/test.mp3"
|
||||
|
||||
def test_frozen_immutable(self):
|
||||
bgm = PresetBGM(id="t1", name="T", style="upbeat", duration=60.0)
|
||||
with pytest.raises(Exception): # FrozenInstanceError
|
||||
bgm.name = "new name"
|
||||
|
||||
def test_equality(self):
|
||||
bgm1 = PresetBGM(id="same", name="N", style="upbeat", duration=60.0)
|
||||
bgm2 = PresetBGM(id="same", name="N", style="upbeat", duration=60.0)
|
||||
assert bgm1 == bgm2
|
||||
|
||||
def test_inequality(self):
|
||||
bgm1 = PresetBGM(id="a", name="A", style="upbeat", duration=60.0)
|
||||
bgm2 = PresetBGM(id="b", name="B", style="upbeat", duration=60.0)
|
||||
assert bgm1 != bgm2
|
||||
|
||||
def test_frozen_with_list_field_not_hashable(self):
|
||||
# 包含 list 字段的 frozen dataclass 仍然不可哈希(list 不可哈希)
|
||||
bgm = PresetBGM(id="h1", name="H", style="upbeat", duration=60.0, tags=["a"])
|
||||
with pytest.raises(TypeError, match="unhashable"):
|
||||
hash(bgm)
|
||||
|
||||
|
||||
class TestPresetBGMLibrary:
|
||||
def test_library_not_empty(self):
|
||||
assert len(PRESET_BGM_LIBRARY) > 0
|
||||
|
||||
def test_library_has_entries(self):
|
||||
assert len(PRESET_BGM_LIBRARY) >= 10
|
||||
|
||||
def test_all_have_unique_ids(self):
|
||||
ids = [bgm.id for bgm in PRESET_BGM_LIBRARY]
|
||||
assert len(ids) == len(set(ids))
|
||||
|
||||
def test_all_have_valid_styles(self):
|
||||
for bgm in PRESET_BGM_LIBRARY:
|
||||
assert bgm.style in BGM_STYLES
|
||||
|
||||
def test_all_have_positive_duration(self):
|
||||
for bgm in PRESET_BGM_LIBRARY:
|
||||
assert bgm.duration > 0
|
||||
|
||||
def test_all_have_non_empty_name(self):
|
||||
for bgm in PRESET_BGM_LIBRARY:
|
||||
assert bgm.name.strip() != ""
|
||||
|
||||
|
||||
class TestBGMStyles:
|
||||
def test_styles_dict_keys(self):
|
||||
assert "upbeat" in BGM_STYLES
|
||||
assert "relax" in BGM_STYLES
|
||||
assert "tech" in BGM_STYLES
|
||||
assert "commerce" in BGM_STYLES
|
||||
assert "emotional" in BGM_STYLES
|
||||
assert "cinematic" in BGM_STYLES
|
||||
|
||||
def test_styles_have_chinese_names(self):
|
||||
for key, value in BGM_STYLES.items():
|
||||
assert isinstance(value, str)
|
||||
assert len(value) > 0
|
||||
|
||||
|
||||
class TestGetPresetBGM:
|
||||
def test_get_existing(self):
|
||||
bgm = get_preset_bgm("bgm_upbeat_001")
|
||||
assert bgm is not None
|
||||
assert bgm.id == "bgm_upbeat_001"
|
||||
assert bgm.name == "阳光清晨"
|
||||
assert bgm.style == "upbeat"
|
||||
|
||||
def test_get_nonexistent(self):
|
||||
assert get_preset_bgm("nonexistent_id") is None
|
||||
|
||||
def test_get_empty_string(self):
|
||||
assert get_preset_bgm("") is None
|
||||
|
||||
def test_get_returns_same_object(self):
|
||||
bgm1 = get_preset_bgm("bgm_relax_001")
|
||||
bgm2 = get_preset_bgm("bgm_relax_001")
|
||||
assert bgm1 is bgm2 # 同一实例(引用同一列表中的对象)
|
||||
|
||||
|
||||
class TestListPresetBGMByStyle:
|
||||
def test_list_upbeat(self):
|
||||
results = list_preset_bgm_by_style("upbeat")
|
||||
assert len(results) >= 3
|
||||
for bgm in results:
|
||||
assert bgm.style == "upbeat"
|
||||
|
||||
def test_list_relax(self):
|
||||
results = list_preset_bgm_by_style("relax")
|
||||
assert len(results) >= 3
|
||||
for bgm in results:
|
||||
assert bgm.style == "relax"
|
||||
|
||||
def test_list_tech(self):
|
||||
results = list_preset_bgm_by_style("tech")
|
||||
assert len(results) >= 2
|
||||
for bgm in results:
|
||||
assert bgm.style == "tech"
|
||||
|
||||
def test_list_commerce(self):
|
||||
results = list_preset_bgm_by_style("commerce")
|
||||
assert len(results) >= 2
|
||||
for bgm in results:
|
||||
assert bgm.style == "commerce"
|
||||
|
||||
def test_list_empty_style(self):
|
||||
results = list_preset_bgm_by_style("nonexistent_style")
|
||||
assert results == []
|
||||
|
||||
def test_list_preserves_order(self):
|
||||
results = list_preset_bgm_by_style("upbeat")
|
||||
ids = [b.id for b in results]
|
||||
# 应该按照在列表中的出现顺序排列
|
||||
assert ids == sorted(ids, key=lambda x: PRESET_BGM_LIBRARY.index(get_preset_bgm(x)))
|
||||
|
||||
|
||||
class TestSearchPresetBGM:
|
||||
def test_search_by_name(self):
|
||||
results = search_preset_bgm("阳光")
|
||||
assert len(results) >= 1
|
||||
assert any("阳光" in b.name for b in results)
|
||||
|
||||
def test_search_by_description(self):
|
||||
results = search_preset_bgm("钢琴")
|
||||
assert len(results) >= 1
|
||||
# 钢琴出现在名称或描述或标签中
|
||||
found = False
|
||||
for b in results:
|
||||
if "钢琴" in b.description or "钢琴" in b.name or "钢琴" in b.tags:
|
||||
found = True
|
||||
break
|
||||
assert found
|
||||
|
||||
def test_search_by_tag(self):
|
||||
results = search_preset_bgm("科技")
|
||||
assert len(results) >= 1
|
||||
found_tech = any(b.style == "tech" for b in results)
|
||||
assert found_tech
|
||||
|
||||
def test_search_case_insensitive(self):
|
||||
results1 = search_preset_bgm("Tech")
|
||||
results2 = search_preset_bgm("tech")
|
||||
assert len(results1) == len(results2)
|
||||
|
||||
def test_search_no_match(self):
|
||||
results = search_preset_bgm("zzzzzzzzzzz_nonexistent_keyword")
|
||||
assert results == []
|
||||
|
||||
def test_search_empty_keyword(self):
|
||||
# 空字符串应该匹配所有(因为空字符串 in 任何字符串都是 True)
|
||||
results = search_preset_bgm("")
|
||||
assert len(results) == len(PRESET_BGM_LIBRARY)
|
||||
|
||||
def test_search_no_duplicates(self):
|
||||
# 确保同一个 BGM 不会出现多次
|
||||
results = search_preset_bgm("电子")
|
||||
ids = [b.id for b in results]
|
||||
assert len(ids) == len(set(ids))
|
||||
+222
-284
@@ -1,6 +1,4 @@
|
||||
"""Quota 领域层单元测试 - quota.py"""
|
||||
|
||||
import math
|
||||
"""Quota 配额系统单元测试。"""
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -19,103 +17,114 @@ from packages.domain.quota import (
|
||||
|
||||
|
||||
class TestQuotaDimension:
|
||||
"""QuotaDimension 枚举测试"""
|
||||
def test_core_dimensions_exist(self):
|
||||
assert QuotaDimension.STORAGE_GB.value == "storage_gb"
|
||||
assert QuotaDimension.VIDEOS_PER_MONTH.value == "videos_per_month"
|
||||
assert QuotaDimension.MAX_CONCURRENT.value == "max_concurrent"
|
||||
assert QuotaDimension.MAX_TEMPLATES.value == "max_templates"
|
||||
assert QuotaDimension.MAX_TITLES.value == "max_titles"
|
||||
assert QuotaDimension.MAX_VOICEOVERS.value == "max_voiceovers"
|
||||
assert QuotaDimension.AI_VOICE_ENABLED.value == "ai_voice_enabled"
|
||||
|
||||
def test_all_dimensions_have_values(self):
|
||||
"""所有枚举成员都有字符串值"""
|
||||
def test_extended_dimensions_exist(self):
|
||||
assert QuotaDimension.AI_VOICE_CREDITS.value == "ai_voice_credits"
|
||||
assert QuotaDimension.BATCH_EXPORT_ENABLED.value == "batch_export_enabled"
|
||||
assert QuotaDimension.MULTI_PLATFORM_ENABLED.value == "multi_platform_enabled"
|
||||
assert QuotaDimension.DEDUP_REPORT_ENABLED.value == "dedup_report_enabled"
|
||||
|
||||
def test_all_dimensions_are_strings(self):
|
||||
for dim in QuotaDimension:
|
||||
assert isinstance(dim.value, str)
|
||||
assert dim.value
|
||||
|
||||
def test_dimension_count(self):
|
||||
"""配额维度数量 >= 内置维度"""
|
||||
# 至少有 storage_gb, videos_per_month, max_concurrent, max_templates 等
|
||||
assert len(QuotaDimension) >= 7
|
||||
|
||||
def test_str_enum_behavior(self):
|
||||
"""是 str 枚举,可直接当字符串用"""
|
||||
assert QuotaDimension.STORAGE_GB == "storage_gb"
|
||||
assert isinstance(QuotaDimension.STORAGE_GB, str)
|
||||
|
||||
|
||||
class TestQuotaTier:
|
||||
"""QuotaTier 测试"""
|
||||
|
||||
def test_get_limit_defined(self):
|
||||
"""已定义的维度返回正确值"""
|
||||
tier = QuotaTier(name="test", limits={"storage": 10, "videos": 5})
|
||||
assert tier.get_limit("storage") == 10
|
||||
tier = QuotaTier(name="test", limits={"storage_gb": 10, "videos": 5})
|
||||
assert tier.get_limit("storage_gb") == 10
|
||||
assert tier.get_limit("videos") == 5
|
||||
|
||||
def test_get_limit_undefined_returns_zero(self):
|
||||
"""未定义的维度返回 0"""
|
||||
tier = QuotaTier(name="test", limits={"storage": 10})
|
||||
assert tier.get_limit("unknown") == 0
|
||||
tier = QuotaTier(name="test", limits={"storage_gb": 10})
|
||||
assert tier.get_limit("unknown_dim") == 0
|
||||
|
||||
def test_is_unlimited_true(self):
|
||||
"""不限量判断 - inf"""
|
||||
def test_is_unlimited_false_for_finite(self):
|
||||
tier = QuotaTier(name="test", limits={"storage_gb": 10})
|
||||
assert tier.is_unlimited("storage_gb") is False
|
||||
|
||||
def test_is_unlimited_true_for_inf(self):
|
||||
tier = QuotaTier(name="test", limits={"templates": float("inf")})
|
||||
assert tier.is_unlimited("templates") is True
|
||||
|
||||
def test_is_unlimited_false(self):
|
||||
"""限量判断"""
|
||||
tier = QuotaTier(name="test", limits={"storage": 10})
|
||||
assert tier.is_unlimited("storage") is False
|
||||
|
||||
def test_is_unlimited_undefined_returns_true(self):
|
||||
"""未定义的维度默认 inf,is_unlimited 返回 True"""
|
||||
def test_is_unlimited_undefined(self):
|
||||
tier = QuotaTier(name="test", limits={})
|
||||
# get_limit 用 dict.get 默认 0,但 is_unlimited 用 dict.get 默认 inf
|
||||
# 未定义的维度,limits.get 返回默认 inf,所以 is_unlimited 返回 True
|
||||
assert tier.is_unlimited("unknown") is True
|
||||
|
||||
def test_empty_limits(self):
|
||||
tier = QuotaTier(name="empty")
|
||||
assert tier.limits == {}
|
||||
assert tier.name == "empty"
|
||||
|
||||
|
||||
class TestQuotaTiers:
|
||||
"""内置套餐配额测试"""
|
||||
|
||||
def test_three_tiers_exist(self):
|
||||
"""三个套餐等级都存在"""
|
||||
assert "free" in QUOTA_TIERS
|
||||
assert "basic" in QUOTA_TIERS
|
||||
assert "premium" in QUOTA_TIERS
|
||||
|
||||
def test_free_tier_storage(self):
|
||||
"""free 套餐 2GB 存储"""
|
||||
assert QUOTA_TIERS["free"].get_limit(QuotaDimension.STORAGE_GB) == 2
|
||||
def test_free_tier_limits(self):
|
||||
free = QUOTA_TIERS["free"]
|
||||
assert free.get_limit("storage_gb") == 2
|
||||
assert free.get_limit("videos_per_month") == 5
|
||||
assert free.get_limit("max_concurrent") == 3
|
||||
assert free.get_limit("max_templates") == 3
|
||||
assert free.get_limit("max_titles") == 50
|
||||
assert free.get_limit("max_voiceovers") == 10
|
||||
assert free.get_limit("ai_voice_enabled") == 0
|
||||
|
||||
def test_basic_tier_storage(self):
|
||||
"""basic 套餐 20GB 存储"""
|
||||
assert QUOTA_TIERS["basic"].get_limit(QuotaDimension.STORAGE_GB) == 20
|
||||
def test_basic_tier_limits(self):
|
||||
basic = QUOTA_TIERS["basic"]
|
||||
assert basic.get_limit("storage_gb") == 20
|
||||
assert basic.get_limit("videos_per_month") == 30
|
||||
assert basic.get_limit("max_concurrent") == 10
|
||||
assert basic.get_limit("max_templates") == 15
|
||||
assert basic.get_limit("max_titles") == 500
|
||||
assert basic.get_limit("max_voiceovers") == 100
|
||||
assert basic.get_limit("ai_voice_enabled") == 1
|
||||
assert basic.get_limit("ai_voice_credits") == 100
|
||||
assert basic.get_limit("batch_export_enabled") == 1
|
||||
|
||||
def test_premium_tier_storage(self):
|
||||
"""premium 套餐 100GB 存储"""
|
||||
assert QUOTA_TIERS["premium"].get_limit(QuotaDimension.STORAGE_GB) == 100
|
||||
def test_premium_tier_limits(self):
|
||||
premium = QUOTA_TIERS["premium"]
|
||||
assert premium.get_limit("storage_gb") == 100
|
||||
assert premium.get_limit("videos_per_month") == 100
|
||||
assert premium.get_limit("max_concurrent") == 20
|
||||
assert premium.is_unlimited("max_templates") is True
|
||||
assert premium.get_limit("ai_voice_enabled") == 1
|
||||
assert premium.get_limit("ai_voice_credits") == 500
|
||||
assert premium.get_limit("batch_export_enabled") == 1
|
||||
assert premium.get_limit("multi_platform_enabled") == 1
|
||||
assert premium.get_limit("dedup_report_enabled") == 1
|
||||
|
||||
def test_free_no_ai_voice(self):
|
||||
"""free 套餐没有 AI 配音"""
|
||||
assert QUOTA_TIERS["free"].get_limit(QuotaDimension.AI_VOICE_ENABLED) == 0
|
||||
|
||||
def test_basic_has_ai_voice(self):
|
||||
"""basic 套餐有 AI 配音"""
|
||||
assert QUOTA_TIERS["basic"].get_limit(QuotaDimension.AI_VOICE_ENABLED) == 1
|
||||
|
||||
def test_premium_templates_unlimited(self):
|
||||
"""premium 套餐模板不限量"""
|
||||
assert QUOTA_TIERS["premium"].is_unlimited(QuotaDimension.MAX_TEMPLATES) is True
|
||||
|
||||
def test_free_videos_per_month(self):
|
||||
"""free 每月 5 个视频"""
|
||||
assert QUOTA_TIERS["free"].get_limit(QuotaDimension.VIDEOS_PER_MONTH) == 5
|
||||
|
||||
def test_premium_multi_platform_enabled(self):
|
||||
"""premium 支持多平台发布"""
|
||||
assert QUOTA_TIERS["premium"].get_limit(QuotaDimension.MULTI_PLATFORM_ENABLED) == 1
|
||||
def test_tier_increase_monotonic(self):
|
||||
free = QUOTA_TIERS["free"]
|
||||
basic = QUOTA_TIERS["basic"]
|
||||
premium = QUOTA_TIERS["premium"]
|
||||
# 高级套餐应该 >= 低级套餐的所有限制
|
||||
for dim in [
|
||||
"storage_gb",
|
||||
"videos_per_month",
|
||||
"max_concurrent",
|
||||
"max_titles",
|
||||
"max_voiceovers",
|
||||
"ai_voice_credits",
|
||||
]:
|
||||
assert basic.get_limit(dim) >= free.get_limit(dim)
|
||||
assert premium.get_limit(dim) >= basic.get_limit(dim)
|
||||
|
||||
|
||||
class TestQuotaWarningLevel:
|
||||
"""告警级别常量测试"""
|
||||
|
||||
def test_level_values(self):
|
||||
"""四个告警级别都有定义"""
|
||||
def test_levels_exist(self):
|
||||
assert QuotaWarningLevel.NORMAL == "normal"
|
||||
assert QuotaWarningLevel.WARNING == "warning"
|
||||
assert QuotaWarningLevel.CRITICAL == "critical"
|
||||
@@ -123,302 +132,231 @@ class TestQuotaWarningLevel:
|
||||
|
||||
|
||||
class TestQuotaCheckResult:
|
||||
"""QuotaCheckResult 测试"""
|
||||
|
||||
def test_usage_percent_normal(self):
|
||||
"""正常使用百分比计算"""
|
||||
result = QuotaCheckResult(
|
||||
allowed=True,
|
||||
dimension="storage",
|
||||
dimension="storage_gb",
|
||||
limit=100,
|
||||
used=30,
|
||||
remaining=70,
|
||||
warning_level=QuotaWarningLevel.NORMAL,
|
||||
used=50,
|
||||
remaining=50,
|
||||
warning_level="normal",
|
||||
)
|
||||
assert result.usage_percent == 30.0
|
||||
assert result.usage_percent == 50.0
|
||||
|
||||
def test_usage_percent_capped_at_100(self):
|
||||
"""超过 100% 时截断为 100%"""
|
||||
def test_usage_percent_exceeded(self):
|
||||
result = QuotaCheckResult(
|
||||
allowed=False,
|
||||
dimension="storage",
|
||||
limit=100,
|
||||
used=150,
|
||||
remaining=0,
|
||||
warning_level=QuotaWarningLevel.EXCEEDED,
|
||||
allowed=False, dimension="d", limit=100, used=150, remaining=0, warning_level="exceeded"
|
||||
)
|
||||
assert result.usage_percent == 100.0
|
||||
assert result.usage_percent == 100.0 # min(100, 150%)
|
||||
|
||||
def test_usage_percent_zero_used(self):
|
||||
result = QuotaCheckResult(allowed=True, dimension="d", limit=100, used=0, remaining=100, warning_level="normal")
|
||||
assert result.usage_percent == 0.0
|
||||
|
||||
def test_usage_percent_zero_limit_with_usage(self):
|
||||
"""limit=0 但有使用量,返回 100%"""
|
||||
result = QuotaCheckResult(
|
||||
allowed=False,
|
||||
dimension="storage",
|
||||
limit=0,
|
||||
used=5,
|
||||
remaining=0,
|
||||
warning_level=QuotaWarningLevel.EXCEEDED,
|
||||
)
|
||||
result = QuotaCheckResult(allowed=False, dimension="d", limit=0, used=10, remaining=0, warning_level="exceeded")
|
||||
assert result.usage_percent == 100.0
|
||||
|
||||
def test_usage_percent_zero_limit_no_usage(self):
|
||||
"""limit=0 且无使用量,返回 0%"""
|
||||
result = QuotaCheckResult(
|
||||
allowed=True,
|
||||
dimension="storage",
|
||||
limit=0,
|
||||
used=0,
|
||||
remaining=0,
|
||||
warning_level=QuotaWarningLevel.NORMAL,
|
||||
)
|
||||
result = QuotaCheckResult(allowed=True, dimension="d", limit=0, used=0, remaining=0, warning_level="normal")
|
||||
assert result.usage_percent == 0.0
|
||||
|
||||
def test_usage_percent_unlimited(self):
|
||||
"""不限量时使用百分比为 0"""
|
||||
result = QuotaCheckResult(
|
||||
allowed=True,
|
||||
dimension="templates",
|
||||
dimension="d",
|
||||
limit=float("inf"),
|
||||
used=50,
|
||||
used=1000,
|
||||
remaining=float("inf"),
|
||||
warning_level=QuotaWarningLevel.NORMAL,
|
||||
warning_level="normal",
|
||||
)
|
||||
assert result.usage_percent == 0.0
|
||||
|
||||
|
||||
class TestQuotaRegistry:
|
||||
"""QuotaRegistry 测试"""
|
||||
|
||||
def test_initial_dimensions(self):
|
||||
"""初始化时内置维度已注册"""
|
||||
registry = QuotaRegistry()
|
||||
dims = registry.list_dimensions()
|
||||
assert QuotaDimension.STORAGE_GB in dims
|
||||
assert QuotaDimension.VIDEOS_PER_MONTH in dims
|
||||
reg = QuotaRegistry()
|
||||
dims = reg.list_dimensions()
|
||||
assert "storage_gb" in dims
|
||||
assert "videos_per_month" in dims
|
||||
assert len(dims) == len(QuotaDimension)
|
||||
|
||||
def test_initial_tiers(self):
|
||||
"""初始化时三个套餐已注册"""
|
||||
registry = QuotaRegistry()
|
||||
tiers = registry.list_tiers()
|
||||
def test_list_tiers(self):
|
||||
reg = QuotaRegistry()
|
||||
tiers = reg.list_tiers()
|
||||
assert "free" in tiers
|
||||
assert "basic" in tiers
|
||||
assert "premium" in tiers
|
||||
|
||||
def test_register_new_dimension(self):
|
||||
"""注册新的配额维度"""
|
||||
registry = QuotaRegistry()
|
||||
registry.register_dimension("custom_dim", "自定义维度")
|
||||
dims = registry.list_dimensions()
|
||||
assert "custom_dim" in dims
|
||||
assert dims["custom_dim"] == "自定义维度"
|
||||
|
||||
def test_register_dimension_idempotent(self):
|
||||
"""重复注册是幂等的"""
|
||||
registry = QuotaRegistry()
|
||||
registry.register_dimension("custom", "描述1")
|
||||
registry.register_dimension("custom", "描述2")
|
||||
# 保留第一次注册的描述
|
||||
assert registry.list_dimensions()["custom"] == "描述1"
|
||||
|
||||
def test_register_with_default_limits(self):
|
||||
"""注册时指定各套餐的默认限制"""
|
||||
registry = QuotaRegistry()
|
||||
registry.register_dimension(
|
||||
"custom",
|
||||
"自定义",
|
||||
default_limits={"free": 1, "basic": 10, "premium": 100},
|
||||
)
|
||||
assert registry.get_limit("free", "custom") == 1
|
||||
assert registry.get_limit("basic", "custom") == 10
|
||||
assert registry.get_limit("premium", "custom") == 100
|
||||
|
||||
def test_register_without_default_limits_defaults_to_zero(self):
|
||||
"""不指定默认限制时各套餐该维度为 0"""
|
||||
registry = QuotaRegistry()
|
||||
registry.register_dimension("custom_no_limit", "自定义")
|
||||
assert registry.get_limit("free", "custom_no_limit") == 0
|
||||
assert registry.get_limit("basic", "custom_no_limit") == 0
|
||||
|
||||
def test_register_default_limits_ignores_unknown_plan(self):
|
||||
"""默认限制中未知的套餐名被忽略"""
|
||||
registry = QuotaRegistry()
|
||||
registry.register_dimension(
|
||||
"custom",
|
||||
"自定义",
|
||||
default_limits={"nonexistent": 999},
|
||||
)
|
||||
# 不报错,但也不会创建新套餐
|
||||
assert registry.get_tier("nonexistent") is None
|
||||
assert len(tiers) == 3
|
||||
|
||||
def test_get_tier_existing(self):
|
||||
"""获取存在的套餐"""
|
||||
registry = QuotaRegistry()
|
||||
tier = registry.get_tier("free")
|
||||
reg = QuotaRegistry()
|
||||
tier = reg.get_tier("free")
|
||||
assert tier is not None
|
||||
assert tier.name == "free"
|
||||
|
||||
def test_get_tier_nonexistent(self):
|
||||
"""获取不存在的套餐返回 None"""
|
||||
registry = QuotaRegistry()
|
||||
assert registry.get_tier("enterprise") is None
|
||||
def test_get_tier_unknown(self):
|
||||
reg = QuotaRegistry()
|
||||
assert reg.get_tier("unknown_plan") is None
|
||||
|
||||
def test_get_limit_existing(self):
|
||||
"""获取存在的套餐和维度的限制"""
|
||||
registry = QuotaRegistry()
|
||||
assert registry.get_limit("free", QuotaDimension.STORAGE_GB) == 2
|
||||
def test_get_limit_known(self):
|
||||
reg = QuotaRegistry()
|
||||
assert reg.get_limit("free", "storage_gb") == 2
|
||||
assert reg.get_limit("premium", "storage_gb") == 100
|
||||
|
||||
def test_get_limit_nonexistent_plan(self):
|
||||
"""不存在的套餐返回 0"""
|
||||
registry = QuotaRegistry()
|
||||
assert registry.get_limit("unknown", QuotaDimension.STORAGE_GB) == 0
|
||||
def test_get_limit_unknown_plan(self):
|
||||
reg = QuotaRegistry()
|
||||
assert reg.get_limit("unknown", "storage_gb") == 0
|
||||
|
||||
def test_list_dimensions_returns_copy(self):
|
||||
"""list_dimensions 返回副本,修改不影响内部"""
|
||||
registry = QuotaRegistry()
|
||||
dims = registry.list_dimensions()
|
||||
dims["fake"] = "fake"
|
||||
assert "fake" not in registry.list_dimensions()
|
||||
def test_register_new_dimension(self):
|
||||
reg = QuotaRegistry()
|
||||
reg.register_dimension("new_feature", "新功能", default_limits={"free": 0, "basic": 1, "premium": 5})
|
||||
dims = reg.list_dimensions()
|
||||
assert "new_feature" in dims
|
||||
assert dims["new_feature"] == "新功能"
|
||||
assert reg.get_limit("free", "new_feature") == 0
|
||||
assert reg.get_limit("basic", "new_feature") == 1
|
||||
assert reg.get_limit("premium", "new_feature") == 5
|
||||
|
||||
def test_list_tiers_returns_all_three(self):
|
||||
"""列出所有套餐"""
|
||||
registry = QuotaRegistry()
|
||||
tiers = registry.list_tiers()
|
||||
assert len(tiers) == 3
|
||||
assert set(tiers) == {"free", "basic", "premium"}
|
||||
def test_register_dimension_idempotent(self):
|
||||
reg = QuotaRegistry()
|
||||
reg.register_dimension("storage_gb", "should not change", default_limits={"free": 999})
|
||||
# 已经存在的不覆盖
|
||||
assert reg.get_limit("free", "storage_gb") == 2
|
||||
|
||||
def test_register_without_defaults(self):
|
||||
reg = QuotaRegistry()
|
||||
reg.register_dimension("new_dim", "描述")
|
||||
assert reg.get_limit("free", "new_dim") == 0
|
||||
assert reg.get_limit("basic", "new_dim") == 0
|
||||
assert reg.get_limit("premium", "new_dim") == 0
|
||||
|
||||
def test_register_partial_limits(self):
|
||||
reg = QuotaRegistry()
|
||||
reg.register_dimension("partial", "partial", default_limits={"premium": 42})
|
||||
assert reg.get_limit("free", "partial") == 0 # 未设置的保持 0
|
||||
assert reg.get_limit("premium", "partial") == 42
|
||||
|
||||
|
||||
class TestQuotaChecker:
|
||||
"""QuotaChecker 测试"""
|
||||
|
||||
def test_check_under_limit_allowed(self):
|
||||
"""使用量低于限制,允许"""
|
||||
def test_check_within_limit(self):
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("free", QuotaDimension.STORAGE_GB, 1.0)
|
||||
result = checker.check("free", "storage_gb", 1)
|
||||
assert result.allowed is True
|
||||
assert result.remaining == 1.0
|
||||
assert result.warning_level == QuotaWarningLevel.NORMAL
|
||||
assert result.limit == 2
|
||||
assert result.used == 1
|
||||
assert result.remaining == 1
|
||||
assert result.dimension == "storage_gb"
|
||||
|
||||
def test_check_at_limit_not_allowed(self):
|
||||
"""使用量等于限制,不允许(used < limit 判定)"""
|
||||
def test_check_exceeded(self):
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("free", QuotaDimension.STORAGE_GB, 2.0)
|
||||
result = checker.check("free", "storage_gb", 3)
|
||||
assert result.allowed is False
|
||||
assert result.remaining == 0
|
||||
assert result.warning_level == QuotaWarningLevel.EXCEEDED
|
||||
assert result.warning_level == "exceeded"
|
||||
|
||||
def test_check_over_limit(self):
|
||||
"""使用量超过限制"""
|
||||
def test_check_exact_limit_not_allowed(self):
|
||||
# used < limit 才 allowed,等于不算
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("free", QuotaDimension.STORAGE_GB, 3.0)
|
||||
result = checker.check("free", "storage_gb", 2)
|
||||
assert result.allowed is False
|
||||
assert result.remaining == 0
|
||||
assert result.warning_level == QuotaWarningLevel.EXCEEDED
|
||||
|
||||
def test_check_warning_level_80_percent(self):
|
||||
"""80% 触发 WARNING"""
|
||||
def test_check_unlimited(self):
|
||||
checker = QuotaChecker()
|
||||
# 100GB 的 80% = 80GB
|
||||
result = checker.check("premium", QuotaDimension.STORAGE_GB, 80.0)
|
||||
assert result.warning_level == QuotaWarningLevel.WARNING
|
||||
result = checker.check("premium", "max_templates", 999999)
|
||||
assert result.allowed is True
|
||||
assert result.remaining == float("inf")
|
||||
assert result.warning_level == "normal"
|
||||
|
||||
def test_check_warning_level_95_percent(self):
|
||||
"""95% 触发 CRITICAL"""
|
||||
def test_check_warning_level_normal(self):
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("premium", QuotaDimension.STORAGE_GB, 95.0)
|
||||
assert result.warning_level == QuotaWarningLevel.CRITICAL
|
||||
result = checker.check("free", "storage_gb", 1) # 50%
|
||||
assert result.warning_level == "normal"
|
||||
|
||||
def test_check_warning_level_warning(self):
|
||||
checker = QuotaChecker()
|
||||
# 80% <= used < 95%
|
||||
result = checker.check("free", "max_templates", 2.5) # 2.5/3 = 83%
|
||||
assert result.warning_level == "warning"
|
||||
|
||||
def test_check_warning_level_critical(self):
|
||||
checker = QuotaChecker()
|
||||
# 95% <= used < 100%
|
||||
result = checker.check("free", "max_templates", 2.9) # 2.9/3 = 97%
|
||||
assert result.warning_level == "critical"
|
||||
|
||||
def test_check_warning_level_exceeded(self):
|
||||
"""100% 及以上触发 EXCEEDED"""
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("premium", QuotaDimension.STORAGE_GB, 100.0)
|
||||
assert result.warning_level == QuotaWarningLevel.EXCEEDED
|
||||
|
||||
def test_check_unlimited_always_allowed(self):
|
||||
"""不限量的维度始终允许"""
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("premium", QuotaDimension.MAX_TEMPLATES, 9999)
|
||||
assert result.allowed is True
|
||||
assert math.isinf(result.remaining)
|
||||
assert result.warning_level == QuotaWarningLevel.NORMAL
|
||||
|
||||
def test_check_unknown_plan_zero_limit(self):
|
||||
"""未知套餐限制为 0,used=0 时不允许(0 < 0 为 False)"""
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("unknown", QuotaDimension.STORAGE_GB, 0)
|
||||
assert result.limit == 0
|
||||
assert result.allowed is False
|
||||
result = checker.check("free", "storage_gb", 5) # 250%
|
||||
assert result.warning_level == "exceeded"
|
||||
|
||||
def test_check_multiple(self):
|
||||
"""批量检查多个维度"""
|
||||
checker = QuotaChecker()
|
||||
results = checker.check_multiple(
|
||||
"free",
|
||||
{
|
||||
QuotaDimension.STORAGE_GB: 1.0,
|
||||
QuotaDimension.VIDEOS_PER_MONTH: 3,
|
||||
},
|
||||
{"storage_gb": 1, "max_templates": 2, "max_titles": 10},
|
||||
)
|
||||
assert len(results) == 2
|
||||
assert len(results) == 3
|
||||
assert results[0].dimension == "storage_gb"
|
||||
assert results[1].dimension == "max_templates"
|
||||
assert results[2].dimension == "max_titles"
|
||||
assert all(r.allowed for r in results)
|
||||
dims = {r.dimension for r in results}
|
||||
assert QuotaDimension.STORAGE_GB in dims
|
||||
assert QuotaDimension.VIDEOS_PER_MONTH in dims
|
||||
|
||||
def test_check_with_custom_registry(self):
|
||||
"""使用自定义注册表"""
|
||||
registry = QuotaRegistry()
|
||||
registry.register_dimension("custom", "自定义", default_limits={"free": 5})
|
||||
checker = QuotaChecker(registry)
|
||||
result = checker.check("free", "custom", 3)
|
||||
assert result.allowed is True
|
||||
assert result.limit == 5
|
||||
def test_check_zero_limit(self):
|
||||
checker = QuotaChecker()
|
||||
result = checker.check("free", "ai_voice_enabled", 0)
|
||||
# limit=0, used=0: used < limit 为 False → allowed=False
|
||||
assert result.allowed is False
|
||||
assert result.remaining == 0
|
||||
assert result.warning_level == "normal"
|
||||
|
||||
def test_compute_warning_level_zero_limit_no_usage(self):
|
||||
"""limit=0, used=0 → NORMAL"""
|
||||
level = QuotaChecker._compute_warning_level(0, 0)
|
||||
assert level == QuotaWarningLevel.NORMAL
|
||||
def test_compute_warning_level_normal(self):
|
||||
assert QuotaChecker._compute_warning_level(50, 100) == "normal"
|
||||
assert QuotaChecker._compute_warning_level(79, 100) == "normal"
|
||||
|
||||
def test_compute_warning_level_warning_boundary(self):
|
||||
assert QuotaChecker._compute_warning_level(80, 100) == "warning"
|
||||
assert QuotaChecker._compute_warning_level(94, 100) == "warning"
|
||||
|
||||
def test_compute_warning_level_critical_boundary(self):
|
||||
assert QuotaChecker._compute_warning_level(95, 100) == "critical"
|
||||
assert QuotaChecker._compute_warning_level(99, 100) == "critical"
|
||||
|
||||
def test_compute_warning_level_exceeded(self):
|
||||
assert QuotaChecker._compute_warning_level(100, 100) == "exceeded"
|
||||
assert QuotaChecker._compute_warning_level(150, 100) == "exceeded"
|
||||
|
||||
def test_compute_warning_level_unlimited(self):
|
||||
assert QuotaChecker._compute_warning_level(9999, float("inf")) == "normal"
|
||||
|
||||
def test_compute_warning_level_zero_limit_with_usage(self):
|
||||
"""limit=0, used>0 → EXCEEDED"""
|
||||
level = QuotaChecker._compute_warning_level(1, 0)
|
||||
assert level == QuotaWarningLevel.EXCEEDED
|
||||
assert QuotaChecker._compute_warning_level(1, 0) == "exceeded"
|
||||
|
||||
def test_compute_warning_level_zero_limit_no_usage(self):
|
||||
assert QuotaChecker._compute_warning_level(0, 0) == "normal"
|
||||
|
||||
def test_compute_warning_level_negative_limit(self):
|
||||
"""limit<0 视同 0 处理"""
|
||||
level = QuotaChecker._compute_warning_level(1, -1)
|
||||
assert level == QuotaWarningLevel.EXCEEDED
|
||||
# limit <= 0 且 used=0 → NORMAL
|
||||
assert QuotaChecker._compute_warning_level(0, -1) == "normal"
|
||||
|
||||
|
||||
class TestGetWarningLevel:
|
||||
"""get_warning_level 便捷函数测试"""
|
||||
|
||||
def test_normal(self):
|
||||
assert get_warning_level(50, 100) == QuotaWarningLevel.NORMAL
|
||||
|
||||
def test_warning(self):
|
||||
assert get_warning_level(85, 100) == QuotaWarningLevel.WARNING
|
||||
|
||||
def test_critical(self):
|
||||
assert get_warning_level(97, 100) == QuotaWarningLevel.CRITICAL
|
||||
|
||||
def test_exceeded(self):
|
||||
assert get_warning_level(100, 100) == QuotaWarningLevel.EXCEEDED
|
||||
|
||||
def test_unlimited(self):
|
||||
assert get_warning_level(9999, float("inf")) == QuotaWarningLevel.NORMAL
|
||||
def test_convenience_function(self):
|
||||
assert get_warning_level(50, 100) == "normal"
|
||||
assert get_warning_level(99, 100) == "critical"
|
||||
assert get_warning_level(100, 100) == "exceeded"
|
||||
assert get_warning_level(0, 0) == "normal"
|
||||
assert get_warning_level(1, 0) == "exceeded"
|
||||
|
||||
|
||||
class TestGlobalSingletons:
|
||||
"""全局单例测试"""
|
||||
|
||||
def test_quota_registry_is_instance(self):
|
||||
assert isinstance(quota_registry, QuotaRegistry)
|
||||
|
||||
def test_quota_checker_is_instance(self):
|
||||
assert isinstance(quota_checker, QuotaChecker)
|
||||
|
||||
def test_global_checker_uses_global_registry(self):
|
||||
"""全局 checker 使用全局 registry"""
|
||||
# 验证能正常工作
|
||||
result = quota_checker.check("free", QuotaDimension.STORAGE_GB, 1.0)
|
||||
def test_global_checker_works(self):
|
||||
result = quota_checker.check("free", "storage_gb", 1)
|
||||
assert result.allowed is True
|
||||
|
||||
Executable
+361
@@ -0,0 +1,361 @@
|
||||
"""render_layer_utils 模块单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.render_layer_utils import (
|
||||
LAYER_Z_INDEX,
|
||||
MAIN_LAYER_ROLES,
|
||||
PIP_DEFAULT_SCALE,
|
||||
can_pass_through,
|
||||
clip_adjusted_duration,
|
||||
clip_effective_duration,
|
||||
clip_playback_speed,
|
||||
estimate_total_duration,
|
||||
get_layer_z_index,
|
||||
resolve_layer_role,
|
||||
)
|
||||
|
||||
# ── 辅助数据类 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeClip:
|
||||
duration: float = 0.0
|
||||
actual_duration: float = 0.0
|
||||
playback_speed: Any = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeLayer:
|
||||
role: str = "main"
|
||||
clips: list[FakeClip] = field(default_factory=list)
|
||||
|
||||
|
||||
# ── 常量验证 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConstants:
|
||||
def test_layer_z_index_has_expected_keys(self):
|
||||
assert set(LAYER_Z_INDEX.keys()) == {
|
||||
"background",
|
||||
"broll",
|
||||
"main",
|
||||
"overlay",
|
||||
"corner_voice",
|
||||
"audio",
|
||||
}
|
||||
|
||||
def test_layer_z_index_ordering(self):
|
||||
assert LAYER_Z_INDEX["background"] < LAYER_Z_INDEX["main"]
|
||||
assert LAYER_Z_INDEX["main"] == LAYER_Z_INDEX["broll"]
|
||||
assert LAYER_Z_INDEX["overlay"] > LAYER_Z_INDEX["main"]
|
||||
assert LAYER_Z_INDEX["corner_voice"] > LAYER_Z_INDEX["main"]
|
||||
assert LAYER_Z_INDEX["audio"] > LAYER_Z_INDEX["overlay"]
|
||||
|
||||
def test_pip_default_scale_positive(self):
|
||||
assert 0 < PIP_DEFAULT_SCALE < 1
|
||||
|
||||
def test_main_layer_roles(self):
|
||||
assert "main" in MAIN_LAYER_ROLES
|
||||
assert "broll" in MAIN_LAYER_ROLES
|
||||
assert "background" in MAIN_LAYER_ROLES
|
||||
assert "overlay" not in MAIN_LAYER_ROLES
|
||||
|
||||
|
||||
# ── resolve_layer_role ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResolveLayerRole:
|
||||
def test_intro_maps_to_main(self):
|
||||
assert resolve_layer_role("intro") == "main"
|
||||
|
||||
def test_outro_maps_to_main(self):
|
||||
assert resolve_layer_role("outro") == "main"
|
||||
|
||||
def test_overlay_maps_to_overlay(self):
|
||||
assert resolve_layer_role("overlay") == "overlay"
|
||||
|
||||
def test_corner_voice_maps_to_corner_voice(self):
|
||||
assert resolve_layer_role("corner_voice") == "corner_voice"
|
||||
|
||||
def test_background_maps_to_background(self):
|
||||
assert resolve_layer_role("background") == "background"
|
||||
|
||||
def test_b_roll_maps_to_broll(self):
|
||||
assert resolve_layer_role("b_roll") == "broll"
|
||||
|
||||
def test_main_defaults_to_main(self):
|
||||
assert resolve_layer_role("main") == "main"
|
||||
|
||||
def test_main_with_b_roll_role(self):
|
||||
assert resolve_layer_role("main", {"role": "b_roll"}) == "broll"
|
||||
|
||||
def test_main_with_audio_role(self):
|
||||
assert resolve_layer_role("main", {"role": "audio"}) == "audio"
|
||||
|
||||
def test_main_with_other_role_stays_main(self):
|
||||
assert resolve_layer_role("main", {"role": "overlay"}) == "main"
|
||||
|
||||
def test_none_config(self):
|
||||
assert resolve_layer_role("main", None) == "main"
|
||||
|
||||
def test_empty_config(self):
|
||||
assert resolve_layer_role("main", {}) == "main"
|
||||
|
||||
def test_unknown_type_defaults_to_main(self):
|
||||
assert resolve_layer_role("unknown_type") == "main"
|
||||
|
||||
|
||||
# ── get_layer_z_index ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGetLayerZIndex:
|
||||
def test_known_roles(self):
|
||||
for role, expected in LAYER_Z_INDEX.items():
|
||||
assert get_layer_z_index(role) == expected
|
||||
|
||||
def test_unknown_role_returns_zero(self):
|
||||
assert get_layer_z_index("nonexistent") == 0
|
||||
|
||||
def test_empty_string_returns_zero(self):
|
||||
assert get_layer_z_index("") == 0
|
||||
|
||||
|
||||
# ── clip_effective_duration ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClipEffectiveDuration:
|
||||
def test_explicit_duration_no_actual(self):
|
||||
assert clip_effective_duration(5.0) == 5.0
|
||||
|
||||
def test_explicit_duration_with_shorter_actual(self):
|
||||
assert clip_effective_duration(5.0, 3.0) == 3.0
|
||||
|
||||
def test_explicit_duration_with_longer_actual(self):
|
||||
assert clip_effective_duration(5.0, 10.0) == 5.0
|
||||
|
||||
def test_zero_duration_uses_actual(self):
|
||||
assert clip_effective_duration(0, 8.0) == 8.0
|
||||
|
||||
def test_negative_duration_uses_actual(self):
|
||||
assert clip_effective_duration(-1.0, 8.0) == 8.0
|
||||
|
||||
def test_zero_duration_zero_actual(self):
|
||||
assert clip_effective_duration(0, 0) == 0.0
|
||||
|
||||
def test_no_args_returns_zero(self):
|
||||
assert clip_effective_duration(0) == 0.0
|
||||
|
||||
def test_equal_duration_and_actual(self):
|
||||
assert clip_effective_duration(5.0, 5.0) == 5.0
|
||||
|
||||
|
||||
# ── clip_playback_speed ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClipPlaybackSpeed:
|
||||
def test_normal_speed(self):
|
||||
assert clip_playback_speed(1.0) == 1.0
|
||||
|
||||
def test_fast_speed(self):
|
||||
assert clip_playback_speed(2.0) == 2.0
|
||||
|
||||
def test_slow_speed(self):
|
||||
assert clip_playback_speed(0.5) == 0.5
|
||||
|
||||
def test_zero_speed_defaults_to_one(self):
|
||||
assert clip_playback_speed(0) == 1.0
|
||||
|
||||
def test_negative_speed_defaults_to_one(self):
|
||||
assert clip_playback_speed(-1.0) == 1.0
|
||||
|
||||
def test_none_defaults_to_one(self):
|
||||
assert clip_playback_speed(None) == 1.0
|
||||
|
||||
def test_string_defaults_to_one(self):
|
||||
assert clip_playback_speed("fast") == 1.0
|
||||
|
||||
def test_int_speed(self):
|
||||
assert clip_playback_speed(2) == 2.0
|
||||
|
||||
|
||||
# ── clip_adjusted_duration ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClipAdjustedDuration:
|
||||
def test_normal_speed_same_as_effective(self):
|
||||
assert clip_adjusted_duration(5.0, 10.0, 1.0) == 5.0
|
||||
|
||||
def test_double_speed_half_duration(self):
|
||||
assert clip_adjusted_duration(10.0, 10.0, 2.0) == pytest.approx(5.0)
|
||||
|
||||
def test_half_speed_double_duration(self):
|
||||
assert clip_adjusted_duration(5.0, 10.0, 0.5) == pytest.approx(10.0)
|
||||
|
||||
def test_invalid_speed_uses_default(self):
|
||||
assert clip_adjusted_duration(5.0, 10.0, 0) == 5.0
|
||||
|
||||
def test_zero_duration(self):
|
||||
assert clip_adjusted_duration(0, 0, 1.0) == 0.0
|
||||
|
||||
def test_actual_duration_only(self):
|
||||
assert clip_adjusted_duration(0, 8.0, 1.0) == 8.0
|
||||
|
||||
def test_actual_duration_only_with_speed(self):
|
||||
assert clip_adjusted_duration(0, 8.0, 2.0) == pytest.approx(4.0)
|
||||
|
||||
def test_very_close_to_normal_speed(self):
|
||||
# 1.0000001 应该被认为接近 1.0,不做除法
|
||||
result = clip_adjusted_duration(5.0, 10.0, 1.0 + 1e-10)
|
||||
assert result == 5.0
|
||||
|
||||
|
||||
# ── estimate_total_duration ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestEstimateTotalDuration:
|
||||
def test_empty_layers(self):
|
||||
assert estimate_total_duration([]) == 0.0
|
||||
|
||||
def test_no_main_layer(self):
|
||||
layers = [FakeLayer(role="overlay", clips=[FakeClip(duration=5.0)])]
|
||||
assert estimate_total_duration(layers) == 0.0
|
||||
|
||||
def test_single_clip_main_layer(self):
|
||||
layers = [FakeLayer(role="main", clips=[FakeClip(duration=5.0)])]
|
||||
assert estimate_total_duration(layers) == pytest.approx(5.0)
|
||||
|
||||
def test_multiple_clips_no_transition(self):
|
||||
layers = [
|
||||
FakeLayer(
|
||||
role="main",
|
||||
clips=[
|
||||
FakeClip(duration=3.0),
|
||||
FakeClip(duration=2.0),
|
||||
FakeClip(duration=5.0),
|
||||
],
|
||||
)
|
||||
]
|
||||
assert estimate_total_duration(layers) == pytest.approx(10.0)
|
||||
|
||||
def test_multiple_clips_with_transition(self):
|
||||
layers = [
|
||||
FakeLayer(
|
||||
role="main",
|
||||
clips=[
|
||||
FakeClip(duration=3.0),
|
||||
FakeClip(duration=2.0),
|
||||
FakeClip(duration=5.0),
|
||||
],
|
||||
)
|
||||
]
|
||||
# 3 + 2 + 5 - 2 * 0.5 = 9.0
|
||||
assert estimate_total_duration(layers, transition_duration=0.5) == pytest.approx(9.0)
|
||||
|
||||
def test_prefers_main_over_broll(self):
|
||||
layers = [
|
||||
FakeLayer(role="broll", clips=[FakeClip(duration=10.0)]),
|
||||
FakeLayer(role="main", clips=[FakeClip(duration=5.0)]),
|
||||
]
|
||||
assert estimate_total_duration(layers) == pytest.approx(5.0)
|
||||
|
||||
def test_prefers_broll_over_background(self):
|
||||
layers = [
|
||||
FakeLayer(role="background", clips=[FakeClip(duration=10.0)]),
|
||||
FakeLayer(role="broll", clips=[FakeClip(duration=5.0)]),
|
||||
]
|
||||
assert estimate_total_duration(layers) == pytest.approx(5.0)
|
||||
|
||||
def test_main_layer_empty_clips(self):
|
||||
layers = [FakeLayer(role="main", clips=[])]
|
||||
assert estimate_total_duration(layers) == 0.0
|
||||
|
||||
def test_minimum_duration(self):
|
||||
layers = [
|
||||
FakeLayer(
|
||||
role="main",
|
||||
clips=[
|
||||
FakeClip(duration=0.01),
|
||||
FakeClip(duration=0.01),
|
||||
],
|
||||
)
|
||||
]
|
||||
result = estimate_total_duration(layers, transition_duration=0.5)
|
||||
assert result >= 0.1
|
||||
|
||||
def test_with_playback_speed(self):
|
||||
layers = [
|
||||
FakeLayer(
|
||||
role="main",
|
||||
clips=[
|
||||
FakeClip(duration=10.0, playback_speed=2.0),
|
||||
FakeClip(duration=10.0, playback_speed=0.5),
|
||||
],
|
||||
)
|
||||
]
|
||||
# 5 + 20 = 25
|
||||
assert estimate_total_duration(layers) == pytest.approx(25.0)
|
||||
|
||||
|
||||
# ── can_pass_through ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCanPassThrough:
|
||||
def test_single_main_clip_no_effects(self):
|
||||
layers = [FakeLayer(role="main", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers) is True
|
||||
|
||||
def test_single_broll_clip(self):
|
||||
layers = [FakeLayer(role="broll", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers) is True
|
||||
|
||||
def test_single_background_clip(self):
|
||||
layers = [FakeLayer(role="background", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers) is True
|
||||
|
||||
def test_multiple_layers(self):
|
||||
layers = [
|
||||
FakeLayer(role="main", clips=[FakeClip(duration=5.0)]),
|
||||
FakeLayer(role="overlay", clips=[FakeClip(duration=3.0)]),
|
||||
]
|
||||
assert can_pass_through(layers) is False
|
||||
|
||||
def test_overlay_layer(self):
|
||||
layers = [FakeLayer(role="overlay", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers) is False
|
||||
|
||||
def test_multiple_clips_in_layer(self):
|
||||
layers = [
|
||||
FakeLayer(
|
||||
role="main",
|
||||
clips=[
|
||||
FakeClip(duration=3.0),
|
||||
FakeClip(duration=2.0),
|
||||
],
|
||||
)
|
||||
]
|
||||
assert can_pass_through(layers) is False
|
||||
|
||||
def test_with_stickers(self):
|
||||
layers = [FakeLayer(role="main", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers, has_stickers=True) is False
|
||||
|
||||
def test_with_watermark(self):
|
||||
layers = [FakeLayer(role="main", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers, has_watermark=True) is False
|
||||
|
||||
def test_with_stickers_and_watermark(self):
|
||||
layers = [FakeLayer(role="main", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers, has_stickers=True, has_watermark=True) is False
|
||||
|
||||
def test_empty_layer_list(self):
|
||||
assert can_pass_through([]) is False
|
||||
|
||||
def test_audio_layer_only(self):
|
||||
layers = [FakeLayer(role="audio", clips=[FakeClip(duration=5.0)])]
|
||||
assert can_pass_through(layers) is False
|
||||
Executable
+401
@@ -0,0 +1,401 @@
|
||||
"""subtitle_style 领域模型单测 — 纯逻辑."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.subtitle_style import (
|
||||
ALLOWED_SUBTITLE_EXTENSIONS,
|
||||
DEFAULT_COLOR,
|
||||
DEFAULT_FONT,
|
||||
DEFAULT_FONT_SIZE,
|
||||
DEFAULT_MAX_CHARS_PER_LINE,
|
||||
DEFAULT_POSITION,
|
||||
DEFAULT_STROKE_COLOR,
|
||||
DEFAULT_STROKE_WIDTH,
|
||||
POSITION_ALIASES,
|
||||
POSITION_ALIGNMENT,
|
||||
SubtitleSegment,
|
||||
SubtitleStyle,
|
||||
escape_ass_text,
|
||||
format_ass_time,
|
||||
hex_to_ass_bgr,
|
||||
hex_to_ass_color,
|
||||
opacity_to_ass_alpha,
|
||||
wrap_text,
|
||||
)
|
||||
|
||||
# ── 常量测试 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConstants:
|
||||
def test_position_alignment_has_9_positions(self):
|
||||
assert len(POSITION_ALIGNMENT) == 9
|
||||
assert POSITION_ALIGNMENT["bottom_center"] == 2
|
||||
assert POSITION_ALIGNMENT["top_center"] == 8
|
||||
assert POSITION_ALIGNMENT["center"] == 5
|
||||
|
||||
def test_position_aliases(self):
|
||||
assert POSITION_ALIASES["top"] == "top_center"
|
||||
assert POSITION_ALIASES["bottom"] == "bottom_center"
|
||||
assert POSITION_ALIASES["middle"] == "center"
|
||||
|
||||
def test_default_values(self):
|
||||
assert DEFAULT_FONT == "思源黑体"
|
||||
assert DEFAULT_FONT_SIZE == 24
|
||||
assert DEFAULT_COLOR == "#FFFFFF"
|
||||
assert DEFAULT_POSITION == "bottom_center"
|
||||
|
||||
def test_allowed_extensions(self):
|
||||
assert ".srt" in ALLOWED_SUBTITLE_EXTENSIONS
|
||||
assert ".ass" in ALLOWED_SUBTITLE_EXTENSIONS
|
||||
assert ".vtt" in ALLOWED_SUBTITLE_EXTENSIONS
|
||||
|
||||
|
||||
# ── 工具函数测试 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestHexToAssColor:
|
||||
def test_white(self):
|
||||
assert hex_to_ass_color("#FFFFFF") == "&H00FFFFFF"
|
||||
|
||||
def test_black(self):
|
||||
assert hex_to_ass_color("#000000") == "&H00000000"
|
||||
|
||||
def test_red(self):
|
||||
assert hex_to_ass_color("#FF0000") == "&H000000FF"
|
||||
|
||||
def test_blue(self):
|
||||
assert hex_to_ass_color("#0000FF") == "&H00FF0000"
|
||||
|
||||
def test_green(self):
|
||||
assert hex_to_ass_color("#00FF00") == "&H0000FF00"
|
||||
|
||||
def test_without_hash(self):
|
||||
assert hex_to_ass_color("FF0000") == "&H000000FF"
|
||||
|
||||
def test_lowercase(self):
|
||||
assert hex_to_ass_color("#ff0000") == "&H000000FF"
|
||||
|
||||
def test_invalid_length_returns_white(self):
|
||||
assert hex_to_ass_color("#FFF") == "&H00FFFFFF"
|
||||
assert hex_to_ass_color("#FF000000") == "&H00FFFFFF"
|
||||
|
||||
def test_empty_string(self):
|
||||
assert hex_to_ass_color("") == "&H00FFFFFF"
|
||||
|
||||
|
||||
class TestHexToAssBgr:
|
||||
def test_white(self):
|
||||
assert hex_to_ass_bgr("#FFFFFF") == "FFFFFF"
|
||||
|
||||
def test_red(self):
|
||||
assert hex_to_ass_bgr("#FF0000") == "0000FF"
|
||||
|
||||
def test_blue(self):
|
||||
assert hex_to_ass_bgr("#0000FF") == "FF0000"
|
||||
|
||||
def test_invalid_length(self):
|
||||
assert hex_to_ass_bgr("#FFF") == "FFFFFF"
|
||||
|
||||
|
||||
class TestOpacityToAssAlpha:
|
||||
def test_fully_opaque(self):
|
||||
assert opacity_to_ass_alpha(1.0) == "00"
|
||||
|
||||
def test_fully_transparent(self):
|
||||
assert opacity_to_ass_alpha(0.0) == "FF"
|
||||
|
||||
def test_half(self):
|
||||
assert opacity_to_ass_alpha(0.5) == "80"
|
||||
|
||||
def test_above_1_clamped(self):
|
||||
assert opacity_to_ass_alpha(1.5) == "00"
|
||||
|
||||
def test_below_0_clamped(self):
|
||||
assert opacity_to_ass_alpha(-0.5) == "FF"
|
||||
|
||||
|
||||
class TestEscapeAssText:
|
||||
def test_newline_unix(self):
|
||||
assert escape_ass_text("hello\nworld") == "hello\\Nworld"
|
||||
|
||||
def test_newline_windows(self):
|
||||
assert escape_ass_text("hello\r\nworld") == "hello\\Nworld"
|
||||
|
||||
def test_newline_mac(self):
|
||||
assert escape_ass_text("hello\rworld") == "hello\\Nworld"
|
||||
|
||||
def test_curly_braces(self):
|
||||
assert escape_ass_text("{text}") == "(text)"
|
||||
|
||||
def test_mixed(self):
|
||||
assert escape_ass_text("hello\n{world}\r\nend") == "hello\\N(world)\\Nend"
|
||||
|
||||
def test_empty(self):
|
||||
assert escape_ass_text("") == ""
|
||||
|
||||
|
||||
class TestFormatAssTime:
|
||||
def test_zero(self):
|
||||
assert format_ass_time(0.0) == "0:00:00.00"
|
||||
|
||||
def test_seconds_only(self):
|
||||
assert format_ass_time(5.5) == "0:00:05.50"
|
||||
|
||||
def test_minutes(self):
|
||||
assert format_ass_time(65.25) == "0:01:05.25"
|
||||
|
||||
def test_hours(self):
|
||||
assert format_ass_time(3661.5) == "1:01:01.50"
|
||||
|
||||
def test_negative_returns_zero(self):
|
||||
assert format_ass_time(-1.0) == "0:00:00.00"
|
||||
|
||||
def test_centiseconds_precision(self):
|
||||
assert format_ass_time(1.234) == "0:00:01.23"
|
||||
|
||||
|
||||
class TestWrapText:
|
||||
def test_short_text_no_wrap(self):
|
||||
assert wrap_text("你好", 10) == ["你好"]
|
||||
|
||||
def test_exact_length_no_wrap(self):
|
||||
text = "你" * 10
|
||||
result = wrap_text(text, 10)
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) == 10
|
||||
|
||||
def test_long_text_breaks_at_max(self):
|
||||
text = "你" * 25
|
||||
result = wrap_text(text, 10)
|
||||
assert len(result) == 3
|
||||
assert len(result[0]) == 10
|
||||
assert len(result[1]) == 10
|
||||
assert len(result[2]) == 5
|
||||
|
||||
def test_breaks_at_punctuation(self):
|
||||
# "一二三四五六。七八九十"共10字,"。"在索引6
|
||||
# max_chars=8 时,从 8 往回找到 4,会命中索引6的"。"
|
||||
text = "一二三四五六。七八九十"
|
||||
result = wrap_text(text, 8)
|
||||
assert len(result) == 2
|
||||
assert result[0] == "一二三四五六。"
|
||||
assert result[1] == "七八九十"
|
||||
|
||||
def test_no_punctuation_breaks_at_max(self):
|
||||
text = "一二三四五六七八九十一二三四五六七八九十"
|
||||
result = wrap_text(text, 10)
|
||||
assert len(result[0]) == 10
|
||||
|
||||
def test_empty_text(self):
|
||||
assert wrap_text("", 10) == [""]
|
||||
|
||||
def test_zero_max_chars(self):
|
||||
assert wrap_text("hello", 0) == ["hello"]
|
||||
|
||||
def test_negative_max_chars(self):
|
||||
result = wrap_text("hello", -5)
|
||||
assert isinstance(result, list)
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
# ── SubtitleStyle 测试 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSubtitleStyleDefaults:
|
||||
def test_default_values(self):
|
||||
style = SubtitleStyle()
|
||||
assert style.font_name == DEFAULT_FONT
|
||||
assert style.font_size == DEFAULT_FONT_SIZE
|
||||
assert style.font_color == DEFAULT_COLOR
|
||||
assert style.bold is False
|
||||
assert style.italic is False
|
||||
assert style.stroke_enabled is True
|
||||
assert style.stroke_color == DEFAULT_STROKE_COLOR
|
||||
assert style.stroke_width == DEFAULT_STROKE_WIDTH
|
||||
assert style.position == DEFAULT_POSITION
|
||||
assert style.max_chars_per_line == DEFAULT_MAX_CHARS_PER_LINE
|
||||
|
||||
|
||||
class TestSubtitleStyleFromDict:
|
||||
def test_none_returns_default(self):
|
||||
style = SubtitleStyle.from_dict(None)
|
||||
assert style.font_name == DEFAULT_FONT
|
||||
|
||||
def test_empty_dict_returns_default(self):
|
||||
style = SubtitleStyle.from_dict({})
|
||||
assert style.font_size == DEFAULT_FONT_SIZE
|
||||
|
||||
def test_custom_font(self):
|
||||
style = SubtitleStyle.from_dict({"font": "微软雅黑", "size": 32})
|
||||
assert style.font_name == "微软雅黑"
|
||||
assert style.font_size == 32
|
||||
|
||||
def test_color(self):
|
||||
style = SubtitleStyle.from_dict({"color": "#FF0000"})
|
||||
assert style.font_color == "#FF0000"
|
||||
|
||||
def test_bold_italic(self):
|
||||
style = SubtitleStyle.from_dict({"bold": True, "italic": True})
|
||||
assert style.bold is True
|
||||
assert style.italic is True
|
||||
|
||||
def test_stroke_config(self):
|
||||
style = SubtitleStyle.from_dict(
|
||||
{
|
||||
"stroke_enabled": False,
|
||||
"stroke_color": "#00FF00",
|
||||
"stroke_width": 2.0,
|
||||
}
|
||||
)
|
||||
assert style.stroke_enabled is False
|
||||
assert style.stroke_color == "#00FF00"
|
||||
assert style.stroke_width == 2.0
|
||||
|
||||
def test_shadow_config(self):
|
||||
style = SubtitleStyle.from_dict(
|
||||
{
|
||||
"shadow_enabled": True,
|
||||
"shadow_color": "#111111",
|
||||
"shadow_offset_x": 4,
|
||||
"shadow_offset_y": 4,
|
||||
"shadow_blur": 1.5,
|
||||
}
|
||||
)
|
||||
assert style.shadow_enabled is True
|
||||
assert style.shadow_color == "#111111"
|
||||
assert style.shadow_offset_x == 4
|
||||
assert style.shadow_offset_y == 4
|
||||
assert style.shadow_blur == 1.5
|
||||
|
||||
def test_background_config(self):
|
||||
style = SubtitleStyle.from_dict(
|
||||
{
|
||||
"background_enabled": True,
|
||||
"background_color": "#000000",
|
||||
"background_opacity": 0.7,
|
||||
"background_padding": 10,
|
||||
"background_radius": 6,
|
||||
}
|
||||
)
|
||||
assert style.background_enabled is True
|
||||
assert style.background_opacity == 0.7
|
||||
assert style.background_padding == 10
|
||||
|
||||
def test_background_opacity_clamped_0_to_1(self):
|
||||
style = SubtitleStyle.from_dict({"background_opacity": -0.5})
|
||||
assert style.background_opacity == 0.0
|
||||
style2 = SubtitleStyle.from_dict({"background_opacity": 1.5})
|
||||
assert style2.background_opacity == 1.0
|
||||
|
||||
def test_position_valid(self):
|
||||
style = SubtitleStyle.from_dict({"position": "top_center"})
|
||||
assert style.position == "top_center"
|
||||
|
||||
def test_position_alias(self):
|
||||
style = SubtitleStyle.from_dict({"position": "top"})
|
||||
assert style.position == "top_center"
|
||||
|
||||
def test_position_invalid_falls_back(self):
|
||||
style = SubtitleStyle.from_dict({"position": "invalid_pos"})
|
||||
assert style.position == DEFAULT_POSITION
|
||||
|
||||
def test_margins(self):
|
||||
style = SubtitleStyle.from_dict({"margin_v": 80, "margin_l": 50, "margin_r": 50})
|
||||
assert style.margin_v == 80
|
||||
assert style.margin_l == 50
|
||||
assert style.margin_r == 50
|
||||
|
||||
def test_max_chars_per_line(self):
|
||||
style = SubtitleStyle.from_dict({"max_chars_per_line": 15})
|
||||
assert style.max_chars_per_line == 15
|
||||
|
||||
def test_line_spacing(self):
|
||||
style = SubtitleStyle.from_dict({"line_spacing": 4})
|
||||
assert style.line_spacing == 4
|
||||
|
||||
def test_fade_in_out(self):
|
||||
style = SubtitleStyle.from_dict({"fade_in": 0.5, "fade_out": 1.0})
|
||||
assert style.fade_in == 0.5
|
||||
assert style.fade_out == 1.0
|
||||
|
||||
def test_fade_negative_clamped(self):
|
||||
style = SubtitleStyle.from_dict({"fade_in": -1, "fade_out": -2})
|
||||
assert style.fade_in == 0.0
|
||||
assert style.fade_out == 0.0
|
||||
|
||||
def test_animation_type(self):
|
||||
style = SubtitleStyle.from_dict({"animation_type": "fade"})
|
||||
assert style.animation_type == "fade"
|
||||
|
||||
def test_invalid_int_falls_back(self):
|
||||
style = SubtitleStyle.from_dict({"size": "not_a_number"})
|
||||
assert style.font_size == DEFAULT_FONT_SIZE
|
||||
|
||||
def test_invalid_float_falls_back(self):
|
||||
style = SubtitleStyle.from_dict({"stroke_width": "abc"})
|
||||
assert style.stroke_width == DEFAULT_STROKE_WIDTH
|
||||
|
||||
|
||||
class TestSubtitleStyleProperties:
|
||||
def test_alignment_bottom_center(self):
|
||||
style = SubtitleStyle(position="bottom_center")
|
||||
assert style.alignment == 2
|
||||
|
||||
def test_alignment_top_center(self):
|
||||
style = SubtitleStyle(position="top_center")
|
||||
assert style.alignment == 8
|
||||
|
||||
def test_ass_font_color(self):
|
||||
style = SubtitleStyle(font_color="#FF0000")
|
||||
assert style.ass_font_color == "&H000000FF"
|
||||
|
||||
def test_ass_stroke_color(self):
|
||||
style = SubtitleStyle(stroke_color="#00FF00")
|
||||
assert style.ass_stroke_color == "&H0000FF00"
|
||||
|
||||
def test_ass_shadow_color(self):
|
||||
style = SubtitleStyle(shadow_color="#0000FF")
|
||||
assert style.ass_shadow_color == "&H00FF0000"
|
||||
|
||||
def test_ass_background_color(self):
|
||||
style = SubtitleStyle(background_color="#FF0000", background_opacity=0.5)
|
||||
# alpha = 255 - 127 = 128 = 0x80, bgr of red = 0000FF
|
||||
assert style.ass_background_color == "&H800000FF"
|
||||
|
||||
|
||||
# ── SubtitleSegment 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSubtitleSegment:
|
||||
def test_basic(self):
|
||||
seg = SubtitleSegment(start=1.0, end=3.0, text="你好")
|
||||
assert seg.start == 1.0
|
||||
assert seg.end == 3.0
|
||||
assert seg.text == "你好"
|
||||
assert seg.style_name == "Default"
|
||||
|
||||
def test_custom_style(self):
|
||||
seg = SubtitleSegment(start=0, end=2, text="hi", style_name="Title")
|
||||
assert seg.style_name == "Title"
|
||||
|
||||
def test_duration(self):
|
||||
seg = SubtitleSegment(start=1.5, end=4.0, text="test")
|
||||
assert seg.duration == 2.5
|
||||
|
||||
def test_duration_zero_when_end_before_start(self):
|
||||
seg = SubtitleSegment(start=5.0, end=3.0, text="test")
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_is_valid_true(self):
|
||||
seg = SubtitleSegment(start=0, end=2, text="hello")
|
||||
assert seg.is_valid is True
|
||||
|
||||
def test_is_valid_empty_text(self):
|
||||
seg = SubtitleSegment(start=0, end=2, text="")
|
||||
assert seg.is_valid is False
|
||||
|
||||
def test_is_valid_zero_duration(self):
|
||||
seg = SubtitleSegment(start=1, end=1, text="hello")
|
||||
assert seg.is_valid is False
|
||||
Executable
+480
@@ -0,0 +1,480 @@
|
||||
"""template_clip_converter 模块单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.template_clip_config import ClipType, TransitionEffect
|
||||
from packages.domain.template_clip_converter import (
|
||||
clip_config_to_snapshot,
|
||||
clip_configs_to_snapshots,
|
||||
clip_to_template_clip_config,
|
||||
clips_to_template_clip_configs,
|
||||
filter_clip_config,
|
||||
filter_plan_config_to_template,
|
||||
safe_parse_clip_type,
|
||||
safe_parse_transition_effect,
|
||||
snapshot_to_template_clip_config,
|
||||
snapshots_to_template_clip_configs,
|
||||
validate_template_name,
|
||||
)
|
||||
|
||||
# ── 辅助数据类 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeClip:
|
||||
"""模拟剪辑计划片段对象."""
|
||||
|
||||
clip_type: Any = "main"
|
||||
order: int = 0
|
||||
duration: float = 5.0
|
||||
text_content: str = ""
|
||||
transition_effect: Any = "cut"
|
||||
playback_speed: float | None = None
|
||||
config: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeClipConfig:
|
||||
"""模拟模板片段配置对象."""
|
||||
|
||||
clip_type: Any = "main"
|
||||
order: int = 0
|
||||
min_duration: float = 0.0
|
||||
max_duration: float = 0.0
|
||||
text_template: str = ""
|
||||
transition_effect: Any = "cut"
|
||||
config: dict[str, Any] | None = None
|
||||
|
||||
|
||||
# ── safe_parse_transition_effect ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSafeParseTransitionEffect:
|
||||
def test_enum_value_passthrough(self):
|
||||
assert safe_parse_transition_effect(TransitionEffect.DISSOLVE) == TransitionEffect.DISSOLVE
|
||||
|
||||
def test_valid_string(self):
|
||||
assert safe_parse_transition_effect("dissolve") == TransitionEffect.DISSOLVE
|
||||
|
||||
def test_invalid_string_defaults_to_cut(self):
|
||||
assert safe_parse_transition_effect("invalid_effect") == TransitionEffect.CUT
|
||||
|
||||
def test_none_defaults_to_cut(self):
|
||||
assert safe_parse_transition_effect(None) == TransitionEffect.CUT
|
||||
|
||||
def test_custom_default(self):
|
||||
assert safe_parse_transition_effect("bad", default=TransitionEffect.FADE) == TransitionEffect.FADE
|
||||
|
||||
def test_int_value(self):
|
||||
assert safe_parse_transition_effect(123) == TransitionEffect.CUT
|
||||
|
||||
|
||||
# ── safe_parse_clip_type ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSafeParseClipType:
|
||||
def test_enum_value_passthrough(self):
|
||||
assert safe_parse_clip_type(ClipType.INTRO) == ClipType.INTRO
|
||||
|
||||
def test_valid_string(self):
|
||||
assert safe_parse_clip_type("intro") == ClipType.INTRO
|
||||
|
||||
def test_invalid_string_defaults_to_main(self):
|
||||
assert safe_parse_clip_type("invalid_type") == ClipType.MAIN
|
||||
|
||||
def test_none_defaults_to_main(self):
|
||||
assert safe_parse_clip_type(None) == ClipType.MAIN
|
||||
|
||||
def test_custom_default(self):
|
||||
assert safe_parse_clip_type("bad", default=ClipType.OUTRO) == ClipType.OUTRO
|
||||
|
||||
def test_int_value(self):
|
||||
assert safe_parse_clip_type(42) == ClipType.MAIN
|
||||
|
||||
|
||||
# ── filter_clip_config ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFilterClipConfig:
|
||||
def test_none_config_no_speed(self):
|
||||
result = filter_clip_config(None)
|
||||
assert result == {}
|
||||
|
||||
def test_empty_config_no_speed(self):
|
||||
result = filter_clip_config({})
|
||||
assert result == {}
|
||||
|
||||
def test_playback_speed_added_when_not_default(self):
|
||||
result = filter_clip_config(None, playback_speed=1.5)
|
||||
assert result == {"playback_speed": 1.5}
|
||||
|
||||
def test_playback_speed_skipped_when_default(self):
|
||||
result = filter_clip_config(None, playback_speed=1.0)
|
||||
assert result == {}
|
||||
|
||||
def test_playback_speed_none_skipped(self):
|
||||
result = filter_clip_config(None, playback_speed=None)
|
||||
assert result == {}
|
||||
|
||||
def test_config_merged(self):
|
||||
result = filter_clip_config({"filter": "vintage", "intensity": 0.5})
|
||||
assert result == {"filter": "vintage", "intensity": 0.5}
|
||||
|
||||
def test_asset_info_removed(self):
|
||||
result = filter_clip_config({"asset_info": {"name": "test.mp4"}, "filter": "vintage"})
|
||||
assert "asset_info" not in result
|
||||
assert result["filter"] == "vintage"
|
||||
|
||||
def test_source_asset_id_removed(self):
|
||||
result = filter_clip_config({"source_asset_id": "abc123", "filter": "vintage"})
|
||||
assert "source_asset_id" not in result
|
||||
assert result["filter"] == "vintage"
|
||||
|
||||
def test_speed_overrides_config_playback_speed(self):
|
||||
result = filter_clip_config({"playback_speed": 2.0}, playback_speed=0.5)
|
||||
assert result["playback_speed"] == 2.0 # config 优先级更高
|
||||
|
||||
def test_custom_skip_keys(self):
|
||||
skip = frozenset({"custom_field"})
|
||||
result = filter_clip_config(
|
||||
{"custom_field": "x", "asset_info": "keep_it"},
|
||||
skip_keys=skip,
|
||||
)
|
||||
assert "custom_field" not in result
|
||||
assert "asset_info" in result # 自定义 skip 覆盖默认
|
||||
|
||||
|
||||
# ── filter_plan_config_to_template ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFilterPlanConfigToTemplate:
|
||||
def test_none_config(self):
|
||||
assert filter_plan_config_to_template(None) == {}
|
||||
|
||||
def test_empty_config(self):
|
||||
assert filter_plan_config_to_template({}) == {}
|
||||
|
||||
def test_draft_flag_removed(self):
|
||||
result = filter_plan_config_to_template({"is_template_draft": True, "editing_mode": "one_take"})
|
||||
assert "is_template_draft" not in result
|
||||
assert result["editing_mode"] == "one_take"
|
||||
|
||||
def test_asset_ids_removed(self):
|
||||
result = filter_plan_config_to_template({"asset_ids": ["a", "b"], "resolution": "1080p"})
|
||||
assert "asset_ids" not in result
|
||||
assert result["resolution"] == "1080p"
|
||||
|
||||
def test_source_edit_plan_id_removed(self):
|
||||
result = filter_plan_config_to_template({"source_edit_plan_id": "plan123", "theme": "dark"})
|
||||
assert "source_edit_plan_id" not in result
|
||||
assert result["theme"] == "dark"
|
||||
|
||||
def test_generation_task_id_removed(self):
|
||||
result = filter_plan_config_to_template({"generation_task_id": "task123", "bgm": "on"})
|
||||
assert "generation_task_id" not in result
|
||||
assert result["bgm"] == "on"
|
||||
|
||||
def test_normal_fields_preserved(self):
|
||||
config = {
|
||||
"editing_mode": "pip",
|
||||
"resolution": "720p",
|
||||
"duration": 30,
|
||||
"style": "cinematic",
|
||||
}
|
||||
result = filter_plan_config_to_template(config)
|
||||
assert result == config
|
||||
|
||||
def test_custom_skip_keys(self):
|
||||
skip = frozenset({"secret_field"})
|
||||
result = filter_plan_config_to_template(
|
||||
{"secret_field": "x", "is_template_draft": "keep"},
|
||||
skip_keys=skip,
|
||||
)
|
||||
assert "secret_field" not in result
|
||||
assert "is_template_draft" in result
|
||||
|
||||
|
||||
# ── clip_to_template_clip_config ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClipToTemplateClipConfig:
|
||||
def test_basic_conversion(self):
|
||||
clip = FakeClip(
|
||||
clip_type="main",
|
||||
order=2,
|
||||
duration=3.5,
|
||||
text_content="Hello world",
|
||||
transition_effect="dissolve",
|
||||
)
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.template_id == "tmpl_001"
|
||||
assert result.clip_type == ClipType.MAIN
|
||||
assert result.order == 2
|
||||
assert result.min_duration == 3.5
|
||||
assert result.max_duration == 3.5
|
||||
assert result.text_template == "Hello world"
|
||||
assert result.transition_effect == TransitionEffect.DISSOLVE
|
||||
|
||||
def test_playback_speed_in_config(self):
|
||||
clip = FakeClip(playback_speed=2.0)
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.config["playback_speed"] == 2.0
|
||||
|
||||
def test_default_speed_not_in_config(self):
|
||||
clip = FakeClip(playback_speed=1.0)
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert "playback_speed" not in result.config
|
||||
|
||||
def test_config_preserved_and_filtered(self):
|
||||
clip = FakeClip(config={"filter": "vintage", "asset_info": {"id": "x"}})
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.config["filter"] == "vintage"
|
||||
assert "asset_info" not in result.config
|
||||
|
||||
def test_text_content_none_becomes_empty(self):
|
||||
clip = FakeClip(text_content=None)
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.text_template == ""
|
||||
|
||||
def test_invalid_type_falls_back(self):
|
||||
clip = FakeClip(clip_type="nonexistent")
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.clip_type == ClipType.MAIN
|
||||
|
||||
def test_invalid_transition_falls_back(self):
|
||||
clip = FakeClip(transition_effect="nonexistent")
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.transition_effect == TransitionEffect.CUT
|
||||
|
||||
def test_enum_type_input(self):
|
||||
clip = FakeClip(clip_type=ClipType.INTRO, transition_effect=TransitionEffect.FADE)
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.clip_type == ClipType.INTRO
|
||||
assert result.transition_effect == TransitionEffect.FADE
|
||||
|
||||
def test_duration_none_defaults_zero(self):
|
||||
clip = FakeClip(duration=None)
|
||||
result = clip_to_template_clip_config("tmpl_001", clip)
|
||||
assert result.min_duration == 0.0
|
||||
assert result.max_duration == 0.0
|
||||
|
||||
|
||||
class TestClipsToTemplateClipConfigs:
|
||||
def test_empty_list(self):
|
||||
result = clips_to_template_clip_configs("tmpl_001", [])
|
||||
assert result == []
|
||||
|
||||
def test_multiple_clips(self):
|
||||
clips = [
|
||||
FakeClip(clip_type="intro", order=0, duration=2.0),
|
||||
FakeClip(clip_type="main", order=1, duration=5.0),
|
||||
FakeClip(clip_type="outro", order=2, duration=3.0),
|
||||
]
|
||||
result = clips_to_template_clip_configs("tmpl_001", clips)
|
||||
assert len(result) == 3
|
||||
assert result[0].clip_type == ClipType.INTRO
|
||||
assert result[1].clip_type == ClipType.MAIN
|
||||
assert result[2].clip_type == ClipType.OUTRO
|
||||
assert all(r.template_id == "tmpl_001" for r in result)
|
||||
|
||||
|
||||
# ── clip_config_to_snapshot ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClipConfigToSnapshot:
|
||||
def test_basic_snapshot(self):
|
||||
cfg = FakeClipConfig(
|
||||
clip_type="main",
|
||||
order=1,
|
||||
min_duration=2.0,
|
||||
max_duration=5.0,
|
||||
text_template="hello",
|
||||
transition_effect="dissolve",
|
||||
config={"filter": "vintage"},
|
||||
)
|
||||
snap = clip_config_to_snapshot(cfg)
|
||||
assert snap["clip_type"] == "main"
|
||||
assert snap["order"] == 1
|
||||
assert snap["min_duration"] == 2.0
|
||||
assert snap["max_duration"] == 5.0
|
||||
assert snap["text_template"] == "hello"
|
||||
assert snap["transition_effect"] == "dissolve"
|
||||
assert snap["config"] == {"filter": "vintage"}
|
||||
|
||||
def test_enum_values_converted_to_strings(self):
|
||||
cfg = FakeClipConfig(
|
||||
clip_type=ClipType.INTRO,
|
||||
transition_effect=TransitionEffect.FADE,
|
||||
)
|
||||
snap = clip_config_to_snapshot(cfg)
|
||||
assert snap["clip_type"] == "intro"
|
||||
assert snap["transition_effect"] == "fade"
|
||||
|
||||
def test_none_text_template_becomes_empty(self):
|
||||
cfg = FakeClipConfig(text_template=None)
|
||||
snap = clip_config_to_snapshot(cfg)
|
||||
assert snap["text_template"] == ""
|
||||
|
||||
def test_none_config_becomes_empty_dict(self):
|
||||
cfg = FakeClipConfig(config=None)
|
||||
snap = clip_config_to_snapshot(cfg)
|
||||
assert snap["config"] == {}
|
||||
|
||||
def test_config_is_copy_not_reference(self):
|
||||
original = {"key": "value"}
|
||||
cfg = FakeClipConfig(config=original)
|
||||
snap = clip_config_to_snapshot(cfg)
|
||||
snap["config"]["key"] = "modified"
|
||||
assert original["key"] == "value"
|
||||
|
||||
|
||||
class TestClipConfigsToSnapshots:
|
||||
def test_empty_list(self):
|
||||
assert clip_configs_to_snapshots([]) == []
|
||||
|
||||
def test_multiple_configs(self):
|
||||
configs = [
|
||||
FakeClipConfig(clip_type="intro", order=0),
|
||||
FakeClipConfig(clip_type="main", order=1),
|
||||
]
|
||||
result = clip_configs_to_snapshots(configs)
|
||||
assert len(result) == 2
|
||||
assert result[0]["clip_type"] == "intro"
|
||||
assert result[1]["order"] == 1
|
||||
|
||||
|
||||
# ── snapshot_to_template_clip_config ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSnapshotToTemplateClipConfig:
|
||||
def test_basic_conversion(self):
|
||||
snap = {
|
||||
"clip_type": "intro",
|
||||
"order": 2,
|
||||
"min_duration": 1.0,
|
||||
"max_duration": 3.0,
|
||||
"text_template": "hi",
|
||||
"transition_effect": "dissolve",
|
||||
"config": {"filter": "bw"},
|
||||
}
|
||||
result = snapshot_to_template_clip_config("tmpl_001", snap)
|
||||
assert result.template_id == "tmpl_001"
|
||||
assert result.clip_type == ClipType.INTRO
|
||||
assert result.order == 2
|
||||
assert result.min_duration == 1.0
|
||||
assert result.max_duration == 3.0
|
||||
assert result.text_template == "hi"
|
||||
assert result.transition_effect == TransitionEffect.DISSOLVE
|
||||
assert result.config == {"filter": "bw"}
|
||||
|
||||
def test_missing_fields_get_defaults(self):
|
||||
result = snapshot_to_template_clip_config("tmpl_001", {})
|
||||
assert result.clip_type == ClipType.MAIN
|
||||
assert result.order == 0
|
||||
assert result.min_duration == 0.0
|
||||
assert result.max_duration == 0.0
|
||||
assert result.text_template == ""
|
||||
assert result.transition_effect == TransitionEffect.CUT
|
||||
assert result.config == {}
|
||||
|
||||
def test_invalid_type_falls_back(self):
|
||||
snap = {"clip_type": "invalid"}
|
||||
result = snapshot_to_template_clip_config("tmpl_001", snap)
|
||||
assert result.clip_type == ClipType.MAIN
|
||||
|
||||
def test_invalid_transition_falls_back(self):
|
||||
snap = {"transition_effect": "invalid"}
|
||||
result = snapshot_to_template_clip_config("tmpl_001", snap)
|
||||
assert result.transition_effect == TransitionEffect.CUT
|
||||
|
||||
def test_none_config_becomes_empty_dict(self):
|
||||
snap = {"config": None}
|
||||
result = snapshot_to_template_clip_config("tmpl_001", snap)
|
||||
assert result.config == {}
|
||||
|
||||
|
||||
class TestSnapshotsToTemplateClipConfigs:
|
||||
def test_empty_list(self):
|
||||
assert snapshots_to_template_clip_configs("tmpl_001", []) == []
|
||||
|
||||
def test_multiple_snapshots(self):
|
||||
snaps = [
|
||||
{"clip_type": "intro", "order": 0},
|
||||
{"clip_type": "outro", "order": 2},
|
||||
]
|
||||
result = snapshots_to_template_clip_configs("tmpl_001", snaps)
|
||||
assert len(result) == 2
|
||||
assert result[0].clip_type == ClipType.INTRO
|
||||
assert result[1].clip_type == ClipType.OUTRO
|
||||
assert all(r.template_id == "tmpl_001" for r in result)
|
||||
|
||||
|
||||
# ── 往返一致性测试 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRoundTrip:
|
||||
def test_snapshot_clip_config_roundtrip(self):
|
||||
"""snapshot → TemplateClipConfig → snapshot 应保持一致."""
|
||||
original = {
|
||||
"clip_type": "intro",
|
||||
"order": 3,
|
||||
"min_duration": 1.5,
|
||||
"max_duration": 4.0,
|
||||
"text_template": "test text",
|
||||
"transition_effect": "dissolve",
|
||||
"config": {"key": "value", "nested": {"a": 1}},
|
||||
}
|
||||
cfg = snapshot_to_template_clip_config("tmpl_test", original)
|
||||
result = clip_config_to_snapshot(cfg)
|
||||
assert result == original
|
||||
|
||||
def test_clip_to_config_to_snapshot(self):
|
||||
"""clip → TemplateClipConfig → snapshot 的预期结果."""
|
||||
clip = FakeClip(
|
||||
clip_type="main",
|
||||
order=1,
|
||||
duration=5.0,
|
||||
text_content="hello",
|
||||
transition_effect="fade",
|
||||
playback_speed=1.5,
|
||||
config={"filter": "vintage", "asset_info": "should_remove"},
|
||||
)
|
||||
cfg = clip_to_template_clip_config("tmpl_001", clip)
|
||||
snap = clip_config_to_snapshot(cfg)
|
||||
assert snap["clip_type"] == "main"
|
||||
assert snap["order"] == 1
|
||||
assert snap["min_duration"] == 5.0
|
||||
assert snap["max_duration"] == 5.0
|
||||
assert snap["text_template"] == "hello"
|
||||
assert snap["transition_effect"] == "fade"
|
||||
assert snap["config"]["playback_speed"] == 1.5
|
||||
assert snap["config"]["filter"] == "vintage"
|
||||
assert "asset_info" not in snap["config"]
|
||||
|
||||
|
||||
# ── validate_template_name ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateTemplateName:
|
||||
def test_valid_name(self):
|
||||
assert validate_template_name("My Template") == "My Template"
|
||||
|
||||
def test_strips_whitespace(self):
|
||||
assert validate_template_name(" Hello ") == "Hello"
|
||||
|
||||
def test_empty_string_raises(self):
|
||||
with pytest.raises(ValueError, match="名称不能为空"):
|
||||
validate_template_name("")
|
||||
|
||||
def test_whitespace_only_raises(self):
|
||||
with pytest.raises(ValueError, match="名称不能为空"):
|
||||
validate_template_name(" ")
|
||||
|
||||
def test_none_raises(self):
|
||||
with pytest.raises(ValueError, match="名称不能为空"):
|
||||
validate_template_name(None)
|
||||
Executable
+143
@@ -0,0 +1,143 @@
|
||||
"""mark_title_used_for_generation 纯逻辑单测.
|
||||
|
||||
验证 title usage 计数 + updated_at 更新逻辑。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import TitleLibraryModel
|
||||
|
||||
|
||||
def test_module_importable():
|
||||
"""确认模块可以正常导入."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation # noqa: F401
|
||||
|
||||
|
||||
class TestMarkTitleUsedForGeneration:
|
||||
"""mark_title_used_for_generation 测试."""
|
||||
|
||||
def _make_task(self, strategy_id: str = "title-1") -> MagicMock:
|
||||
task = MagicMock()
|
||||
task.strategy_id = strategy_id
|
||||
return task
|
||||
|
||||
def _make_title(self, usage_count: int = 0) -> MagicMock:
|
||||
title = MagicMock(spec=TitleLibraryModel)
|
||||
title.id = "title-1"
|
||||
title.usage_count = usage_count
|
||||
title.updated_at = None
|
||||
return title
|
||||
|
||||
def test_no_strategy_id_returns_early(self):
|
||||
"""无 strategy_id 时直接返回,不查 DB."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
task = self._make_task(strategy_id="")
|
||||
|
||||
mark_title_used_for_generation(db, task)
|
||||
|
||||
db.query.assert_not_called()
|
||||
|
||||
def test_none_strategy_id_returns_early(self):
|
||||
"""strategy_id 为 None 时直接返回."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
task = self._make_task(strategy_id=None)
|
||||
|
||||
mark_title_used_for_generation(db, task)
|
||||
|
||||
db.query.assert_not_called()
|
||||
|
||||
def test_title_not_found_returns_early(self):
|
||||
"""title 不存在时不报错,静默返回."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
db.query.return_value.filter.return_value.first.return_value = None
|
||||
task = self._make_task(strategy_id="missing-id")
|
||||
|
||||
mark_title_used_for_generation(db, task)
|
||||
|
||||
db.add.assert_not_called()
|
||||
db.commit.assert_not_called()
|
||||
|
||||
def test_increments_usage_count_from_zero(self):
|
||||
"""usage_count 从 0 递增到 1."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
title = self._make_title(usage_count=0)
|
||||
db.query.return_value.filter.return_value.first.return_value = title
|
||||
task = self._make_task()
|
||||
|
||||
mark_title_used_for_generation(db, task)
|
||||
|
||||
assert title.usage_count == 1
|
||||
db.add.assert_called_once_with(title)
|
||||
db.commit.assert_called_once()
|
||||
|
||||
def test_increments_usage_count_from_existing(self):
|
||||
"""已有 usage_count 时递增."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
title = self._make_title(usage_count=5)
|
||||
db.query.return_value.filter.return_value.first.return_value = title
|
||||
task = self._make_task()
|
||||
|
||||
mark_title_used_for_generation(db, task)
|
||||
|
||||
assert title.usage_count == 6
|
||||
|
||||
def test_none_usage_count_defaults_to_zero_then_increments(self):
|
||||
"""usage_count 为 None 时按 0 处理,递增到 1."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
title = self._make_title(usage_count=None)
|
||||
db.query.return_value.filter.return_value.first.return_value = title
|
||||
task = self._make_task()
|
||||
|
||||
mark_title_used_for_generation(db, task)
|
||||
|
||||
assert title.usage_count == 1
|
||||
|
||||
def test_updates_updated_at_to_utc_now(self):
|
||||
"""updated_at 更新为当前 UTC 时间."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
title = self._make_title(usage_count=3)
|
||||
db.query.return_value.filter.return_value.first.return_value = title
|
||||
task = self._make_task()
|
||||
|
||||
before = datetime.now(timezone.utc)
|
||||
mark_title_used_for_generation(db, task)
|
||||
after = datetime.now(timezone.utc)
|
||||
|
||||
assert before <= title.updated_at <= after
|
||||
assert title.updated_at.tzinfo is not None # 带时区
|
||||
|
||||
def test_correct_query_filter(self):
|
||||
"""查询时使用正确的 id 过滤."""
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
db = MagicMock()
|
||||
title = self._make_title()
|
||||
db.query.return_value.filter.return_value.first.return_value = title
|
||||
task = self._make_task(strategy_id="title-abc")
|
||||
|
||||
mark_title_used_for_generation(db, task)
|
||||
|
||||
# 验证 query 模型正确
|
||||
db.query.assert_called_once_with(TitleLibraryModel)
|
||||
# 验证 filter 条件
|
||||
filter_call = db.query.return_value.filter
|
||||
assert filter_call.called
|
||||
# first 被调用
|
||||
filter_call.return_value.first.assert_called_once()
|
||||
Executable
+441
@@ -0,0 +1,441 @@
|
||||
"""TTSStreamingService 纯逻辑单测 — 分段策略 + 分块推送 + 错误处理.
|
||||
|
||||
mock 掉 WebSocket 和 CosyVoiceService,验证核心逻辑:
|
||||
- 文本长度路由(短文本/长文本)
|
||||
- 空文本/超长文本校验
|
||||
- 音频分块推送算法
|
||||
- 错误处理路径
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.cosyvoice_service import CosyVoiceError
|
||||
from packages.application.tts_job.streaming_service import (
|
||||
_AUDIO_CHUNK_SIZE,
|
||||
TTSStreamingError,
|
||||
TTSStreamingService,
|
||||
)
|
||||
|
||||
# ── Fixtures ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_cosyvoice():
|
||||
"""mock CosyVoiceService."""
|
||||
svc = MagicMock()
|
||||
svc.submit_synthesize_task.return_value = {
|
||||
"audio_url": "https://example.com/audio.mp3",
|
||||
"duration": 3.5,
|
||||
}
|
||||
return svc
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def streaming_service(mock_cosyvoice):
|
||||
"""TTSStreamingService 实例."""
|
||||
return TTSStreamingService(mock_cosyvoice)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_ws():
|
||||
"""mock WebSocket."""
|
||||
ws = AsyncMock()
|
||||
ws.send_bytes = AsyncMock()
|
||||
ws.send_json = AsyncMock()
|
||||
return ws
|
||||
|
||||
|
||||
class FakeAudioBytes:
|
||||
"""生成指定大小的假音频数据."""
|
||||
|
||||
@staticmethod
|
||||
def make(size: int) -> bytes:
|
||||
return b"\x00" * size
|
||||
|
||||
|
||||
# ── 合成路由测试 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSynthesizeRouting:
|
||||
"""synthesize_and_stream 路由逻辑测试."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_text_returns_error(self, streaming_service, mock_ws):
|
||||
"""空文本返回错误,不调用合成."""
|
||||
await streaming_service.synthesize_and_stream(mock_ws, {"text": ""})
|
||||
|
||||
# 应发送 error 消息
|
||||
mock_ws.send_json.assert_called()
|
||||
last_call = mock_ws.send_json.call_args
|
||||
assert last_call[0][0]["type"] == "error"
|
||||
assert "不能为空" in last_call[0][0]["message"]
|
||||
# 不应调用合成
|
||||
streaming_service._cosyvoice.submit_synthesize_task.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_text_key_returns_error(self, streaming_service, mock_ws):
|
||||
"""缺少 text 字段返回错误."""
|
||||
await streaming_service.synthesize_and_stream(mock_ws, {})
|
||||
|
||||
mock_ws.send_json.assert_called()
|
||||
last_call = mock_ws.send_json.call_args
|
||||
assert last_call[0][0]["type"] == "error"
|
||||
streaming_service._cosyvoice.submit_synthesize_task.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_too_long_text_returns_error(self, streaming_service, mock_ws):
|
||||
"""超长文本返回错误."""
|
||||
long_text = "你" * 10001
|
||||
await streaming_service.synthesize_and_stream(mock_ws, {"text": long_text})
|
||||
|
||||
mock_ws.send_json.assert_called()
|
||||
last_call = mock_ws.send_json.call_args
|
||||
assert last_call[0][0]["type"] == "error"
|
||||
assert "最大" in last_call[0][0]["message"]
|
||||
streaming_service._cosyvoice.submit_synthesize_task.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_text_uses_short_path(self, streaming_service, mock_ws):
|
||||
"""短文本走 _stream_short_text 路径."""
|
||||
with patch.object(streaming_service, "_stream_short_text", new_callable=AsyncMock) as mock_short:
|
||||
await streaming_service.synthesize_and_stream(mock_ws, {"text": "hello"})
|
||||
mock_short.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_long_text_uses_long_path(self, streaming_service, mock_ws):
|
||||
"""长文本走 _stream_long_text 路径."""
|
||||
long_text = "你" * 501
|
||||
with patch.object(streaming_service, "_stream_long_text", new_callable=AsyncMock) as mock_long:
|
||||
await streaming_service.synthesize_and_stream(mock_ws, {"text": long_text})
|
||||
mock_long.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exactly_threshold_uses_short_path(self, streaming_service, mock_ws):
|
||||
"""恰好等于阈值走短文本路径."""
|
||||
text = "你" * 500
|
||||
with (
|
||||
patch.object(streaming_service, "_stream_short_text", new_callable=AsyncMock) as mock_short,
|
||||
patch.object(streaming_service, "_stream_long_text", new_callable=AsyncMock) as mock_long,
|
||||
):
|
||||
await streaming_service.synthesize_and_stream(mock_ws, {"text": text})
|
||||
mock_short.assert_called_once()
|
||||
mock_long.assert_not_called()
|
||||
|
||||
|
||||
# ── 短文本流测试 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStreamShortText:
|
||||
"""短文本流式合成测试."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_happy_path_sends_started_then_done(self, streaming_service, mock_ws):
|
||||
"""短文本正常流程:started → 音频块 → done."""
|
||||
audio_data = FakeAudioBytes.make(5000)
|
||||
with patch.object(streaming_service, "_download_audio", return_value=audio_data):
|
||||
await streaming_service._stream_short_text(
|
||||
mock_ws,
|
||||
{"text": "hello", "voice_id": "v1", "format": "mp3", "speed": 1.0},
|
||||
)
|
||||
|
||||
# 检查 started 消息
|
||||
calls = mock_ws.send_json.call_args_list
|
||||
assert calls[0][0][0]["type"] == "started"
|
||||
assert calls[0][0][0]["segment_count"] == 1
|
||||
|
||||
# 检查 done 消息
|
||||
last_msg = calls[-1][0][0]
|
||||
assert last_msg["type"] == "done"
|
||||
assert last_msg["file_size"] == 5000
|
||||
assert last_msg["format"] == "mp3"
|
||||
assert last_msg["duration"] == 3.5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_calls_cosyvoice_with_correct_params(self, streaming_service, mock_ws):
|
||||
"""正确传递参数给 CosyVoice."""
|
||||
audio_data = FakeAudioBytes.make(1000)
|
||||
with patch.object(streaming_service, "_download_audio", return_value=audio_data):
|
||||
await streaming_service._stream_short_text(
|
||||
mock_ws,
|
||||
{
|
||||
"text": "test text",
|
||||
"voice_id": "voice-123",
|
||||
"sample_rate": 22050,
|
||||
"format": "wav",
|
||||
"speed": 1.5,
|
||||
},
|
||||
)
|
||||
|
||||
streaming_service._cosyvoice.submit_synthesize_task.assert_called_once_with(
|
||||
text="test text",
|
||||
voice_id="voice-123",
|
||||
sample_rate=22050,
|
||||
format="wav",
|
||||
speed=1.5,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cosyvoice_error_returns_error(self, streaming_service, mock_ws):
|
||||
"""CosyVoice 错误返回 error 消息."""
|
||||
streaming_service._cosyvoice.submit_synthesize_task.side_effect = CosyVoiceError("API quota exceeded")
|
||||
|
||||
await streaming_service._stream_short_text(mock_ws, {"text": "hello"})
|
||||
|
||||
# 最后一条应该是 error
|
||||
last_msg = mock_ws.send_json.call_args_list[-1][0][0]
|
||||
assert last_msg["type"] == "error"
|
||||
assert "API quota exceeded" in last_msg["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generic_exception_returns_error(self, streaming_service, mock_ws):
|
||||
"""普通异常返回 error 消息."""
|
||||
streaming_service._cosyvoice.submit_synthesize_task.side_effect = RuntimeError("boom")
|
||||
|
||||
await streaming_service._stream_short_text(mock_ws, {"text": "hello"})
|
||||
|
||||
last_msg = mock_ws.send_json.call_args_list[-1][0][0]
|
||||
assert last_msg["type"] == "error"
|
||||
assert "合成失败" in last_msg["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_audio_url_returns_error(self, streaming_service, mock_ws):
|
||||
"""合成结果无 audio_url 返回错误."""
|
||||
streaming_service._cosyvoice.submit_synthesize_task.return_value = {"duration": 3.0}
|
||||
|
||||
await streaming_service._stream_short_text(mock_ws, {"text": "hello"})
|
||||
|
||||
last_msg = mock_ws.send_json.call_args_list[-1][0][0]
|
||||
assert last_msg["type"] == "error"
|
||||
assert "音频 URL" in last_msg["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_failure_returns_error(self, streaming_service, mock_ws):
|
||||
"""音频下载失败返回 error."""
|
||||
with patch.object(
|
||||
streaming_service,
|
||||
"_download_audio",
|
||||
side_effect=Exception("download failed"),
|
||||
):
|
||||
await streaming_service._stream_short_text(mock_ws, {"text": "hello"})
|
||||
|
||||
last_msg = mock_ws.send_json.call_args_list[-1][0][0]
|
||||
assert last_msg["type"] == "error"
|
||||
assert "音频推送失败" in last_msg["message"]
|
||||
|
||||
|
||||
# ── 音频分块测试 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStreamAudioChunks:
|
||||
"""_stream_audio_chunks 分块推送测试."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exact_one_chunk(self, streaming_service, mock_ws):
|
||||
"""恰好一个 chunk 大小的数据."""
|
||||
data = FakeAudioBytes.make(_AUDIO_CHUNK_SIZE)
|
||||
total = await streaming_service._stream_audio_chunks(mock_ws, data)
|
||||
|
||||
assert total == _AUDIO_CHUNK_SIZE
|
||||
assert mock_ws.send_bytes.call_count == 1
|
||||
assert len(mock_ws.send_bytes.call_args[0][0]) == _AUDIO_CHUNK_SIZE
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_smaller_than_one_chunk(self, streaming_service, mock_ws):
|
||||
"""小于一个 chunk 的数据."""
|
||||
data = FakeAudioBytes.make(1000)
|
||||
total = await streaming_service._stream_audio_chunks(mock_ws, data)
|
||||
|
||||
assert total == 1000
|
||||
assert mock_ws.send_bytes.call_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_full_chunks(self, streaming_service, mock_ws):
|
||||
"""多个完整 chunk."""
|
||||
num_chunks = 5
|
||||
data = FakeAudioBytes.make(_AUDIO_CHUNK_SIZE * num_chunks)
|
||||
total = await streaming_service._stream_audio_chunks(mock_ws, data)
|
||||
|
||||
assert total == _AUDIO_CHUNK_SIZE * num_chunks
|
||||
assert mock_ws.send_bytes.call_count == num_chunks
|
||||
for c in mock_ws.send_bytes.call_args_list:
|
||||
assert len(c[0][0]) == _AUDIO_CHUNK_SIZE
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_partial_last_chunk(self, streaming_service, mock_ws):
|
||||
"""最后一个 chunk 不完整."""
|
||||
data = FakeAudioBytes.make(_AUDIO_CHUNK_SIZE * 2 + 1234)
|
||||
total = await streaming_service._stream_audio_chunks(mock_ws, data)
|
||||
|
||||
assert total == _AUDIO_CHUNK_SIZE * 2 + 1234
|
||||
assert mock_ws.send_bytes.call_count == 3
|
||||
# 最后一块是 1234 字节
|
||||
last_chunk = mock_ws.send_bytes.call_args_list[-1][0][0]
|
||||
assert len(last_chunk) == 1234
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_audio_sends_zero_chunks(self, streaming_service, mock_ws):
|
||||
"""空音频不发送任何 chunk."""
|
||||
total = await streaming_service._stream_audio_chunks(mock_ws, b"")
|
||||
assert total == 0
|
||||
mock_ws.send_bytes.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunks_are_consecutive(self, streaming_service, mock_ws):
|
||||
"""所有 chunk 拼接起来等于原始数据."""
|
||||
data = bytes(range(256)) * 50 # 12800 bytes
|
||||
total = await streaming_service._stream_audio_chunks(mock_ws, data)
|
||||
|
||||
assert total == len(data)
|
||||
# 收集所有 chunk
|
||||
all_bytes = b"".join(c[0][0] for c in mock_ws.send_bytes.call_args_list)
|
||||
assert all_bytes == data
|
||||
|
||||
|
||||
# ── 长文本分段流测试 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStreamLongText:
|
||||
"""长文本分段流式合成测试."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_happy_path_all_segments_ok(self, streaming_service, mock_ws):
|
||||
"""长文本正常流程:多个分段全部成功."""
|
||||
audio_data = FakeAudioBytes.make(2000)
|
||||
with patch.object(streaming_service, "_download_audio", return_value=audio_data):
|
||||
text = "你" * 1200 # 应该分成3段
|
||||
await streaming_service._stream_long_text(
|
||||
mock_ws,
|
||||
{"text": text, "voice_id": "v1", "format": "mp3", "speed": 1.0},
|
||||
)
|
||||
|
||||
# 检查 started 消息
|
||||
calls = mock_ws.send_json.call_args_list
|
||||
assert calls[0][0][0]["type"] == "started"
|
||||
segment_count = calls[0][0][0]["segment_count"]
|
||||
assert segment_count >= 2 # 1200 字至少分 2 段
|
||||
|
||||
# 检查有 segment_done 消息
|
||||
segment_dones = [c for c in calls if c[0][0].get("type") == "segment_done"]
|
||||
assert len(segment_dones) == segment_count
|
||||
|
||||
# 检查最后是 done 消息
|
||||
last_msg = calls[-1][0][0]
|
||||
assert last_msg["type"] == "done"
|
||||
assert last_msg["file_size"] == 2000 * segment_count
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_first_segment_fails_returns_error(self, streaming_service, mock_ws):
|
||||
"""第一个分段失败,立即返回错误."""
|
||||
streaming_service._cosyvoice.submit_synthesize_task.side_effect = CosyVoiceError("segment 0 failed")
|
||||
|
||||
text = "你" * 1200
|
||||
await streaming_service._stream_long_text(mock_ws, {"text": text, "voice_id": "v1"})
|
||||
|
||||
calls = mock_ws.send_json.call_args_list
|
||||
last_msg = calls[-1][0][0]
|
||||
assert last_msg["type"] == "error"
|
||||
assert "分段" in last_msg["message"]
|
||||
assert "1" in last_msg["message"] # 第1段失败
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_segment_without_audio_url_fails(self, streaming_service, mock_ws):
|
||||
"""分段结果无 audio_url 视为失败."""
|
||||
# 第一段正常,第二段返回空 audio_url
|
||||
call_results = [
|
||||
{"audio_url": "https://a.com/1.mp3", "duration": 2.0},
|
||||
{"audio_url": "", "duration": 0},
|
||||
{"audio_url": "https://a.com/3.mp3", "duration": 3.0},
|
||||
]
|
||||
streaming_service._cosyvoice.submit_synthesize_task.side_effect = call_results
|
||||
|
||||
with patch.object(
|
||||
streaming_service,
|
||||
"_download_audio",
|
||||
return_value=FakeAudioBytes.make(1000),
|
||||
):
|
||||
text = "你" * 1500
|
||||
await streaming_service._stream_long_text(mock_ws, {"text": text, "voice_id": "v1"})
|
||||
|
||||
calls = mock_ws.send_json.call_args_list
|
||||
# 应该有错误
|
||||
error_msgs = [c for c in calls if c[0][0].get("type") == "error"]
|
||||
assert len(error_msgs) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_each_segment_gets_correct_text(self, streaming_service, mock_ws):
|
||||
"""每个分段都调用了合成,且 text 参数不同."""
|
||||
with patch.object(
|
||||
streaming_service,
|
||||
"_download_audio",
|
||||
return_value=FakeAudioBytes.make(500),
|
||||
):
|
||||
text = "你" * 1200
|
||||
await streaming_service._stream_long_text(mock_ws, {"text": text, "voice_id": "v1"})
|
||||
|
||||
# 分段数应大于1
|
||||
assert streaming_service._cosyvoice.submit_synthesize_task.call_count >= 2
|
||||
|
||||
# 收集所有传进去的 text
|
||||
texts_called = [
|
||||
c.kwargs.get("text") or c.args[0]
|
||||
for c in streaming_service._cosyvoice.submit_synthesize_task.call_args_list
|
||||
]
|
||||
# 每段文本都应该是原文的一部分(不全部相同)
|
||||
assert len(set(texts_called)) >= 2
|
||||
# 所有文本拼接起来应该约等于原文长度
|
||||
total_len = sum(len(t) for t in texts_called)
|
||||
assert total_len >= len(text) * 0.95 # 允许标点切分的小误差
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_each_segment_has_unique_index(self, streaming_service, mock_ws):
|
||||
"""segment_done 消息的序号不重复且正确."""
|
||||
with patch.object(
|
||||
streaming_service,
|
||||
"_download_audio",
|
||||
return_value=FakeAudioBytes.make(500),
|
||||
):
|
||||
text = "你" * 1200
|
||||
await streaming_service._stream_long_text(mock_ws, {"text": text, "voice_id": "v1"})
|
||||
|
||||
calls = mock_ws.send_json.call_args_list
|
||||
segment_dones = [c[0][0] for c in calls if c[0][0].get("type") == "segment_done"]
|
||||
indices = [s["segment"] for s in segment_dones]
|
||||
total = segment_dones[0]["total"]
|
||||
# 序号从 1 到 total,不重复
|
||||
assert sorted(indices) == list(range(1, total + 1))
|
||||
|
||||
|
||||
# ── WebSocket 发送失败容错 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSendJsonErrorHandling:
|
||||
"""_send_json 容错测试."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_json_failure_logs_warning(self, streaming_service, mock_ws):
|
||||
"""WebSocket send_json 失败不抛异常."""
|
||||
mock_ws.send_json.side_effect = Exception("connection closed")
|
||||
|
||||
# 不应抛出异常
|
||||
await streaming_service._send_json(mock_ws, {"type": "done"})
|
||||
mock_ws.send_json.assert_called_once()
|
||||
|
||||
|
||||
# ── TTSStreamingError 异常类 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTTSStreamingError:
|
||||
"""TTSStreamingError 异常类测试."""
|
||||
|
||||
def test_is_exception(self):
|
||||
"""是 Exception 子类."""
|
||||
assert issubclass(TTSStreamingError, Exception)
|
||||
|
||||
def test_carry_message(self):
|
||||
"""携带错误消息."""
|
||||
err = TTSStreamingError("stream failed")
|
||||
assert str(err) == "stream failed"
|
||||
Executable
+364
@@ -0,0 +1,364 @@
|
||||
"""video_concat 领域模型单测 — 纯逻辑,48个测试用例."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.video_concat import (
|
||||
ALLOWED_VIDEO_EXTENSIONS,
|
||||
CONCAT_DEMUXER_REQUIRED_PARAMS,
|
||||
MAX_CONCAT_SEGMENTS,
|
||||
ConcatConfig,
|
||||
ConcatSegment,
|
||||
)
|
||||
|
||||
# ── ConcatSegment 测试 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConcatSegmentBasics:
|
||||
def test_default_values(self):
|
||||
seg = ConcatSegment(video_path="test.mp4")
|
||||
assert seg.video_path == "test.mp4"
|
||||
assert seg.start_time == 0.0
|
||||
assert seg.duration == 0.0
|
||||
assert seg.has_audio is True
|
||||
|
||||
def test_full_params(self):
|
||||
seg = ConcatSegment(
|
||||
video_path="video.mp4",
|
||||
start_time=5.5,
|
||||
duration=10.0,
|
||||
has_audio=False,
|
||||
)
|
||||
assert seg.video_path == "video.mp4"
|
||||
assert seg.start_time == 5.5
|
||||
assert seg.duration == 10.0
|
||||
assert seg.has_audio is False
|
||||
|
||||
|
||||
class TestConcatSegmentFromDict:
|
||||
def test_normal_dict(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "test.mp4",
|
||||
"start_time": 2.0,
|
||||
"duration": 5.0,
|
||||
"has_audio": False,
|
||||
}
|
||||
)
|
||||
assert seg.video_path == "test.mp4"
|
||||
assert seg.start_time == 2.0
|
||||
assert seg.duration == 5.0
|
||||
assert seg.has_audio is False
|
||||
|
||||
def test_empty_dict(self):
|
||||
seg = ConcatSegment.from_dict({})
|
||||
assert seg.video_path == ""
|
||||
assert seg.start_time == 0.0
|
||||
assert seg.duration == 0.0
|
||||
assert seg.has_audio is True
|
||||
|
||||
def test_none_input(self):
|
||||
seg = ConcatSegment.from_dict(None)
|
||||
assert seg.video_path == ""
|
||||
assert seg.is_valid is False
|
||||
|
||||
def test_non_dict_input(self):
|
||||
seg = ConcatSegment.from_dict("not a dict")
|
||||
assert seg.video_path == ""
|
||||
|
||||
def test_start_time_negative_clamped(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "start_time": -5})
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_duration_negative_clamped(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "duration": -10})
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_start_time_invalid_string(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "start_time": "abc"})
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_duration_invalid_string(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "duration": "xyz"})
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_start_time_int_casted(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "start_time": 3})
|
||||
assert seg.start_time == 3.0
|
||||
|
||||
def test_duration_int_casted(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "duration": 7})
|
||||
assert seg.duration == 7.0
|
||||
|
||||
def test_video_path_casted_to_string(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": 12345})
|
||||
assert seg.video_path == "12345"
|
||||
|
||||
def test_has_audio_false(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "has_audio": False})
|
||||
assert seg.has_audio is False
|
||||
|
||||
def test_has_audio_truthy_value(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "a.mp4", "has_audio": 1})
|
||||
assert seg.has_audio is True
|
||||
|
||||
|
||||
class TestConcatSegmentProperties:
|
||||
def test_is_valid_with_path(self):
|
||||
seg = ConcatSegment(video_path="test.mp4")
|
||||
assert seg.is_valid is True
|
||||
|
||||
def test_is_valid_empty_path(self):
|
||||
seg = ConcatSegment(video_path="")
|
||||
assert seg.is_valid is False
|
||||
|
||||
def test_effective_duration_positive(self):
|
||||
seg = ConcatSegment(video_path="a.mp4", duration=10.5)
|
||||
assert seg.effective_duration == 10.5
|
||||
|
||||
def test_effective_duration_zero(self):
|
||||
seg = ConcatSegment(video_path="a.mp4", duration=0.0)
|
||||
assert seg.effective_duration == 0.0
|
||||
|
||||
def test_effective_duration_negative(self):
|
||||
seg = ConcatSegment(video_path="a.mp4", duration=-5.0)
|
||||
assert seg.effective_duration == 0.0
|
||||
|
||||
|
||||
# ── ConcatConfig 测试 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConcatConfigBasics:
|
||||
def test_default_values(self):
|
||||
cfg = ConcatConfig()
|
||||
assert cfg.segments == []
|
||||
assert cfg.output_width == 0
|
||||
assert cfg.output_height == 0
|
||||
assert cfg.output_fps == 0.0
|
||||
assert cfg.force_reencode is False
|
||||
assert cfg.transition == "none"
|
||||
assert cfg.transition_duration == 0.3
|
||||
|
||||
def test_with_segments(self):
|
||||
segs = [ConcatSegment(video_path="a.mp4")]
|
||||
cfg = ConcatConfig(segments=segs)
|
||||
assert len(cfg.segments) == 1
|
||||
assert cfg.segments[0].video_path == "a.mp4"
|
||||
|
||||
|
||||
class TestConcatConfigFromDict:
|
||||
def test_none_config(self):
|
||||
cfg = ConcatConfig.from_config_dict(None)
|
||||
assert cfg.segments == []
|
||||
assert cfg.output_width == 0
|
||||
|
||||
def test_empty_dict(self):
|
||||
cfg = ConcatConfig.from_config_dict({})
|
||||
assert cfg.segments == []
|
||||
|
||||
def test_non_dict_input(self):
|
||||
cfg = ConcatConfig.from_config_dict("config")
|
||||
assert cfg.segments == []
|
||||
|
||||
def test_with_valid_segments(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "a.mp4", "duration": 10},
|
||||
{"video_path": "b.mp4", "duration": 20},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.segments) == 2
|
||||
assert cfg.segments[0].video_path == "a.mp4"
|
||||
assert cfg.segments[1].video_path == "b.mp4"
|
||||
|
||||
def test_skips_empty_video_path(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "a.mp4"},
|
||||
{"video_path": ""},
|
||||
{"video_path": "b.mp4"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.segments) == 2
|
||||
|
||||
def test_skips_invalid_segment_dict(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "a.mp4"},
|
||||
"not a dict",
|
||||
{"video_path": "b.mp4"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(cfg.segments) == 2
|
||||
|
||||
def test_segments_not_a_list(self):
|
||||
cfg = ConcatConfig.from_config_dict({"segments": "not a list"})
|
||||
assert cfg.segments == []
|
||||
|
||||
def test_output_params(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"output_width": 1920,
|
||||
"output_height": 1080,
|
||||
"output_fps": 30.0,
|
||||
"force_reencode": True,
|
||||
}
|
||||
)
|
||||
assert cfg.output_width == 1920
|
||||
assert cfg.output_height == 1080
|
||||
assert cfg.output_fps == 30.0
|
||||
assert cfg.force_reencode is True
|
||||
|
||||
def test_output_width_negative_clamped(self):
|
||||
cfg = ConcatConfig.from_config_dict({"output_width": -100})
|
||||
assert cfg.output_width == 0
|
||||
|
||||
def test_output_height_invalid_string(self):
|
||||
cfg = ConcatConfig.from_config_dict({"output_height": "abc"})
|
||||
assert cfg.output_height == 0
|
||||
|
||||
def test_output_fps_invalid_string(self):
|
||||
cfg = ConcatConfig.from_config_dict({"output_fps": "xyz"})
|
||||
assert cfg.output_fps == 0.0
|
||||
|
||||
def test_transition_params(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"transition": "crossfade",
|
||||
"transition_duration": 1.0,
|
||||
}
|
||||
)
|
||||
assert cfg.transition == "crossfade"
|
||||
assert cfg.transition_duration == 1.0
|
||||
|
||||
def test_transition_duration_minimum(self):
|
||||
cfg = ConcatConfig.from_config_dict({"transition_duration": 0.01})
|
||||
assert cfg.transition_duration == 0.1
|
||||
|
||||
def test_transition_duration_negative(self):
|
||||
cfg = ConcatConfig.from_config_dict({"transition_duration": -1})
|
||||
assert cfg.transition_duration == 0.1
|
||||
|
||||
def test_force_reencode_false_by_default(self):
|
||||
cfg = ConcatConfig.from_config_dict({})
|
||||
assert cfg.force_reencode is False
|
||||
|
||||
|
||||
class TestConcatConfigProperties:
|
||||
def test_has_effect_two_segments(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="a.mp4"),
|
||||
ConcatSegment(video_path="b.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is True
|
||||
|
||||
def test_has_effect_one_segment(self):
|
||||
cfg = ConcatConfig(segments=[ConcatSegment(video_path="a.mp4")])
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_has_effect_empty(self):
|
||||
cfg = ConcatConfig()
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_has_effect_skips_invalid(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="a.mp4"),
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path="b.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.has_effect is True
|
||||
|
||||
def test_valid_segment_count(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="a.mp4"),
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path="b.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.valid_segment_count == 2
|
||||
|
||||
def test_first_valid_segment(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path="first.mp4"),
|
||||
ConcatSegment(video_path="second.mp4"),
|
||||
]
|
||||
)
|
||||
assert cfg.first_valid_segment is not None
|
||||
assert cfg.first_valid_segment.video_path == "first.mp4"
|
||||
|
||||
def test_first_valid_segment_none_when_all_empty(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path=""),
|
||||
]
|
||||
)
|
||||
assert cfg.first_valid_segment is None
|
||||
|
||||
def test_first_valid_segment_empty_list(self):
|
||||
cfg = ConcatConfig()
|
||||
assert cfg.first_valid_segment is None
|
||||
|
||||
def test_estimated_total_duration(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="a.mp4", duration=10.0),
|
||||
ConcatSegment(video_path="b.mp4", duration=20.0),
|
||||
ConcatSegment(video_path="c.mp4", duration=0.0),
|
||||
]
|
||||
)
|
||||
assert cfg.estimated_total_duration == 30.0
|
||||
|
||||
def test_estimated_total_duration_skips_invalid(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="", duration=10.0),
|
||||
ConcatSegment(video_path="a.mp4", duration=5.0),
|
||||
]
|
||||
)
|
||||
assert cfg.estimated_total_duration == 5.0
|
||||
|
||||
|
||||
class TestConcatConfigClampSegments:
|
||||
def test_clamp_when_over_max(self):
|
||||
segs = [ConcatSegment(video_path=f"s{i}.mp4") for i in range(100)]
|
||||
cfg = ConcatConfig(segments=segs)
|
||||
cfg.clamp_segments(50)
|
||||
assert len(cfg.segments) == 50
|
||||
assert cfg.segments[0].video_path == "s0.mp4"
|
||||
assert cfg.segments[-1].video_path == "s49.mp4"
|
||||
|
||||
def test_no_clamp_when_under_max(self):
|
||||
segs = [ConcatSegment(video_path=f"s{i}.mp4") for i in range(10)]
|
||||
cfg = ConcatConfig(segments=segs)
|
||||
cfg.clamp_segments(50)
|
||||
assert len(cfg.segments) == 10
|
||||
|
||||
def test_default_max_constant(self):
|
||||
assert MAX_CONCAT_SEGMENTS == 50
|
||||
|
||||
|
||||
class TestConstants:
|
||||
def test_allowed_extensions(self):
|
||||
assert ".mp4" in ALLOWED_VIDEO_EXTENSIONS
|
||||
assert ".mov" in ALLOWED_VIDEO_EXTENSIONS
|
||||
assert ".webm" in ALLOWED_VIDEO_EXTENSIONS
|
||||
|
||||
def test_demuxer_params(self):
|
||||
assert "codec_name" in CONCAT_DEMUXER_REQUIRED_PARAMS
|
||||
assert "width" in CONCAT_DEMUXER_REQUIRED_PARAMS
|
||||
assert "r_frame_rate" in CONCAT_DEMUXER_REQUIRED_PARAMS
|
||||
Executable
+855
@@ -0,0 +1,855 @@
|
||||
"""video_filter_builder 单元测试 — FFmpeg 滤镜构建纯逻辑层。
|
||||
|
||||
覆盖:
|
||||
- ClipFilterChain 数据类
|
||||
- 常量与映射表
|
||||
- build_clip_filter:单片段滤镜链
|
||||
- chain_filters:滤镜串联工具
|
||||
- has_audio:音频流判断
|
||||
- build_concat_filter:concat 滤镜
|
||||
- build_xfade_filter:xfade 转场滤镜
|
||||
- build_filter_complex:策略选择(空/单片段/concat/xfade)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
from packages.domain.video_filter_builder import (
|
||||
DEFAULT_CLIP_DURATION,
|
||||
DEFAULT_FPS,
|
||||
DEFAULT_OUTPUT_HEIGHT,
|
||||
DEFAULT_OUTPUT_WIDTH,
|
||||
DEFAULT_TRANSITION_DURATION,
|
||||
XFADE_TRANSITION_MAP,
|
||||
ClipFilterChain,
|
||||
build_clip_filter,
|
||||
build_concat_filter,
|
||||
build_filter_complex,
|
||||
build_xfade_filter,
|
||||
chain_filters,
|
||||
has_audio,
|
||||
)
|
||||
|
||||
# ── 辅助:构造 EditPlanClip ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_clip(
|
||||
clip_id: str = "clip-1",
|
||||
duration: float = 5.0,
|
||||
start_time: float = 0.0,
|
||||
clip_type: str = "video",
|
||||
asset_id: str | None = "asset-1",
|
||||
) -> EditPlanClip:
|
||||
"""构造一个测试用 EditPlanClip。"""
|
||||
return EditPlanClip(
|
||||
id=clip_id,
|
||||
plan_id="plan-1",
|
||||
asset_id=asset_id,
|
||||
clip_type=clip_type,
|
||||
duration=duration,
|
||||
start_time=start_time,
|
||||
order=0,
|
||||
status=EditPlanClipStatus.READY,
|
||||
)
|
||||
|
||||
|
||||
# ── ClipFilterChain 数据类测试 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClipFilterChain(unittest.TestCase):
|
||||
"""ClipFilterChain 数据类测试。"""
|
||||
|
||||
def test_immutable(self):
|
||||
"""ClipFilterChain 是 frozen dataclass,不可修改。"""
|
||||
chain = ClipFilterChain(
|
||||
clip_id="c1",
|
||||
input_index=0,
|
||||
video_label="v0",
|
||||
audio_label="a0",
|
||||
filters=["scale=1280:720"],
|
||||
duration=5.0,
|
||||
)
|
||||
with self.assertRaises(Exception):
|
||||
chain.duration = 10.0 # type: ignore[misc]
|
||||
|
||||
def test_fields(self):
|
||||
"""所有字段正确存储。"""
|
||||
chain = ClipFilterChain(
|
||||
clip_id="c1",
|
||||
input_index=2,
|
||||
video_label="v2",
|
||||
audio_label=None,
|
||||
filters=["fps=25", "trim=0:3"],
|
||||
duration=3.0,
|
||||
)
|
||||
self.assertEqual(chain.clip_id, "c1")
|
||||
self.assertEqual(chain.input_index, 2)
|
||||
self.assertEqual(chain.video_label, "v2")
|
||||
self.assertIsNone(chain.audio_label)
|
||||
self.assertEqual(chain.filters, ["fps=25", "trim=0:3"])
|
||||
self.assertEqual(chain.duration, 3.0)
|
||||
|
||||
|
||||
# ── 常量测试 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConstants(unittest.TestCase):
|
||||
"""常量与映射表测试。"""
|
||||
|
||||
def test_default_output_size(self):
|
||||
"""默认输出分辨率 1280x720。"""
|
||||
self.assertEqual(DEFAULT_OUTPUT_WIDTH, 1280)
|
||||
self.assertEqual(DEFAULT_OUTPUT_HEIGHT, 720)
|
||||
|
||||
def test_default_fps(self):
|
||||
"""默认帧率 25。"""
|
||||
self.assertEqual(DEFAULT_FPS, 25)
|
||||
|
||||
def test_default_transition_duration(self):
|
||||
"""默认转场时长 0.5 秒。"""
|
||||
self.assertEqual(DEFAULT_TRANSITION_DURATION, 0.5)
|
||||
|
||||
def test_default_clip_duration(self):
|
||||
"""默认片段时长 5 秒。"""
|
||||
self.assertEqual(DEFAULT_CLIP_DURATION, 5.0)
|
||||
|
||||
def test_xfade_transition_map_keys(self):
|
||||
"""xfade 映射包含所有转场类型。"""
|
||||
self.assertIn(TransitionEffect.FADE, XFADE_TRANSITION_MAP)
|
||||
self.assertIn(TransitionEffect.SLIDE_LEFT, XFADE_TRANSITION_MAP)
|
||||
self.assertIn(TransitionEffect.SLIDE_RIGHT, XFADE_TRANSITION_MAP)
|
||||
self.assertIn(TransitionEffect.DISSOLVE, XFADE_TRANSITION_MAP)
|
||||
self.assertIn(TransitionEffect.WIPE, XFADE_TRANSITION_MAP)
|
||||
|
||||
def test_xfade_transition_map_values(self):
|
||||
"""xfade 映射值为 FFmpeg 合法 transition 名称。"""
|
||||
self.assertEqual(XFADE_TRANSITION_MAP[TransitionEffect.FADE], "fade")
|
||||
self.assertEqual(XFADE_TRANSITION_MAP[TransitionEffect.SLIDE_LEFT], "slideleft")
|
||||
self.assertEqual(XFADE_TRANSITION_MAP[TransitionEffect.SLIDE_RIGHT], "slideright")
|
||||
self.assertEqual(XFADE_TRANSITION_MAP[TransitionEffect.DISSOLVE], "dissolve")
|
||||
self.assertEqual(XFADE_TRANSITION_MAP[TransitionEffect.WIPE], "wipeleft")
|
||||
|
||||
|
||||
# ── chain_filters 测试 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestChainFilters(unittest.TestCase):
|
||||
"""chain_filters 滤镜串联工具测试。"""
|
||||
|
||||
def test_single_filter(self):
|
||||
"""单个滤镜。"""
|
||||
result = chain_filters(["scale=1280:720"], "v0")
|
||||
self.assertEqual(result, "[0:v]scale=1280:720[v0]")
|
||||
|
||||
def test_multiple_filters(self):
|
||||
"""多个滤镜用逗号串联。"""
|
||||
result = chain_filters(["scale=1280:720", "fps=25", "trim=0:5"], "v1")
|
||||
self.assertEqual(result, "[0:v]scale=1280:720,fps=25,trim=0:5[v1]")
|
||||
|
||||
def test_empty_filters(self):
|
||||
"""空滤镜列表。"""
|
||||
result = chain_filters([], "v0")
|
||||
self.assertEqual(result, "[0:v][v0]")
|
||||
|
||||
def test_custom_input_label(self):
|
||||
"""自定义输入标签。"""
|
||||
result = chain_filters(["fps=30"], "out", input_label="v0")
|
||||
self.assertEqual(result, "[v0]fps=30[out]")
|
||||
|
||||
|
||||
# ── has_audio 测试 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestHasAudio(unittest.TestCase):
|
||||
"""has_audio 音频流判断测试。"""
|
||||
|
||||
def test_all_have_audio(self):
|
||||
"""所有片段都有音频。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", "a1", [], 3.0),
|
||||
]
|
||||
self.assertTrue(has_audio(chains))
|
||||
|
||||
def test_some_have_audio(self):
|
||||
"""部分片段有音频。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
]
|
||||
self.assertTrue(has_audio(chains))
|
||||
|
||||
def test_none_have_audio(self):
|
||||
"""没有片段有音频。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
]
|
||||
self.assertFalse(has_audio(chains))
|
||||
|
||||
def test_empty_list(self):
|
||||
"""空列表返回 False。"""
|
||||
self.assertFalse(has_audio([]))
|
||||
|
||||
|
||||
# ── build_clip_filter 测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildClipFilter(unittest.TestCase):
|
||||
"""build_clip_filter 单片段滤镜链测试。"""
|
||||
|
||||
def test_basic_video_clip(self):
|
||||
"""普通视频片段生成完整滤镜链。"""
|
||||
clip = _make_clip(duration=5.0, clip_type="video")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
|
||||
self.assertEqual(chain.clip_id, "clip-1")
|
||||
self.assertEqual(chain.input_index, 0)
|
||||
self.assertEqual(chain.video_label, "v0")
|
||||
self.assertEqual(chain.audio_label, "a0")
|
||||
self.assertEqual(chain.duration, 5.0)
|
||||
# 应有 7 个滤镜:scale, pad, format, fps, setpts, trim, setpts
|
||||
self.assertEqual(len(chain.filters), 7)
|
||||
|
||||
def test_filter_order(self):
|
||||
"""滤镜顺序:scale → pad → format → fps → setpts → trim → setpts。"""
|
||||
clip = _make_clip(duration=3.0, clip_type="video")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
|
||||
self.assertTrue(chain.filters[0].startswith("scale="))
|
||||
self.assertTrue(chain.filters[1].startswith("pad="))
|
||||
self.assertEqual(chain.filters[2], "format=yuv420p")
|
||||
self.assertTrue(chain.filters[3].startswith("fps="))
|
||||
self.assertTrue(chain.filters[4].startswith("setpts="))
|
||||
self.assertTrue(chain.filters[5].startswith("trim="))
|
||||
self.assertEqual(chain.filters[6], "setpts=PTS-STARTPTS")
|
||||
|
||||
def test_scale_force_original_aspect_ratio(self):
|
||||
"""scale 使用 decrease 保持比例。"""
|
||||
clip = _make_clip(duration=5.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIn("force_original_aspect_ratio=decrease", chain.filters[0])
|
||||
|
||||
def test_pad_centered_black(self):
|
||||
"""pad 居中 + 黑边。"""
|
||||
clip = _make_clip(duration=5.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIn("(ow-iw)/2:(oh-ih)/2:black", chain.filters[1])
|
||||
|
||||
def test_format_yuv420p(self):
|
||||
"""像素格式统一为 yuv420p。"""
|
||||
clip = _make_clip(duration=5.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertEqual(chain.filters[2], "format=yuv420p")
|
||||
|
||||
def test_custom_resolution(self):
|
||||
"""自定义输出分辨率。"""
|
||||
clip = _make_clip(duration=5.0)
|
||||
chain = build_clip_filter(clip, 0, 1920, 1080, 30)
|
||||
self.assertIn("scale=1920:1080", chain.filters[0])
|
||||
self.assertIn("pad=1920:1080", chain.filters[1])
|
||||
self.assertEqual(chain.filters[3], "fps=30")
|
||||
|
||||
def test_zero_duration_uses_default(self):
|
||||
"""duration <= 0 时使用默认时长。"""
|
||||
clip = _make_clip(duration=0.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertEqual(chain.duration, DEFAULT_CLIP_DURATION)
|
||||
self.assertIn(f"trim=0:{DEFAULT_CLIP_DURATION}", chain.filters[5])
|
||||
|
||||
def test_negative_duration_uses_default(self):
|
||||
"""负时长也使用默认时长。"""
|
||||
clip = _make_clip(duration=-1.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertEqual(chain.duration, DEFAULT_CLIP_DURATION)
|
||||
|
||||
def test_start_time_offset(self):
|
||||
"""start_time > 0 时 setpts 带偏移。"""
|
||||
clip = _make_clip(duration=3.0, start_time=2.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIn("PTS-STARTPTS+2.0/TB", chain.filters[4])
|
||||
|
||||
def test_zero_start_time_no_offset(self):
|
||||
"""start_time = 0 时 setpts 不带偏移。"""
|
||||
clip = _make_clip(duration=3.0, start_time=0.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertEqual(chain.filters[4], "setpts=PTS-STARTPTS")
|
||||
|
||||
def test_title_clip_no_audio(self):
|
||||
"""title 类型片段没有音频。"""
|
||||
clip = _make_clip(clip_type="title")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIsNone(chain.audio_label)
|
||||
|
||||
def test_subtitle_clip_no_audio(self):
|
||||
"""subtitle 类型片段没有音频。"""
|
||||
clip = _make_clip(clip_type="subtitle")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIsNone(chain.audio_label)
|
||||
|
||||
def test_video_clip_has_audio(self):
|
||||
"""video 类型片段有音频。"""
|
||||
clip = _make_clip(clip_type="video")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertEqual(chain.audio_label, "a0")
|
||||
|
||||
def test_image_clip_has_audio(self):
|
||||
"""image 类型片段有音频标签(可能有BGM)。"""
|
||||
clip = _make_clip(clip_type="image")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIsNotNone(chain.audio_label)
|
||||
|
||||
def test_clip_type_case_insensitive(self):
|
||||
"""clip_type 大小写不敏感。"""
|
||||
clip = _make_clip(clip_type="TITLE")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIsNone(chain.audio_label)
|
||||
|
||||
def test_empty_clip_type_has_audio(self):
|
||||
"""空 clip_type 默认有音频。"""
|
||||
clip = _make_clip(clip_type="")
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIsNotNone(chain.audio_label)
|
||||
|
||||
def test_input_index_reflected_in_labels(self):
|
||||
"""input_index 反映在 video_label 和 audio_label 中。"""
|
||||
clip = _make_clip()
|
||||
chain = build_clip_filter(clip, 3, 1280, 720, 25)
|
||||
self.assertEqual(chain.video_label, "v3")
|
||||
self.assertEqual(chain.audio_label, "a3")
|
||||
|
||||
def test_zero_fps_skipped(self):
|
||||
"""fps = 0 时跳过 fps 滤镜。"""
|
||||
clip = _make_clip(duration=5.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 0)
|
||||
# 少了 fps 滤镜:scale, pad, format, setpts, trim, setpts = 6个
|
||||
self.assertEqual(len(chain.filters), 6)
|
||||
self.assertFalse(any(f.startswith("fps=") for f in chain.filters))
|
||||
|
||||
def test_negative_fps_skipped(self):
|
||||
"""fps < 0 时也跳过 fps 滤镜。"""
|
||||
clip = _make_clip(duration=5.0)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, -1)
|
||||
self.assertEqual(len(chain.filters), 6)
|
||||
|
||||
def test_trim_uses_duration(self):
|
||||
"""trim 时长等于 clip.duration。"""
|
||||
clip = _make_clip(duration=7.5)
|
||||
chain = build_clip_filter(clip, 0, 1280, 720, 25)
|
||||
self.assertIn("trim=0:7.5", chain.filters[5])
|
||||
|
||||
|
||||
# ── build_concat_filter 测试 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildConcatFilter(unittest.TestCase):
|
||||
"""build_concat_filter 拼接滤镜测试。"""
|
||||
|
||||
def test_empty_clips(self):
|
||||
"""空列表返回空字符串和 0 时长。"""
|
||||
filter_str, duration = build_concat_filter([])
|
||||
self.assertEqual(filter_str, "")
|
||||
self.assertEqual(duration, 0.0)
|
||||
|
||||
def test_single_clip_no_audio(self):
|
||||
"""单片段无音频:视频滤镜 + concat(n=1)。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, ["scale=1280:720"], 5.0),
|
||||
]
|
||||
filter_str, duration = build_concat_filter(chains)
|
||||
|
||||
self.assertIn("[0:v]scale=1280:720[v0]", filter_str)
|
||||
self.assertIn("[v0]concat=n=1:v=1:a=0[outv]", filter_str)
|
||||
self.assertEqual(duration, 5.0)
|
||||
# 没有音频相关
|
||||
self.assertNotIn("[outa]", filter_str)
|
||||
|
||||
def test_single_clip_with_audio(self):
|
||||
"""单片段有音频:视频 + 音频归一化 + concat。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", ["scale=1280:720"], 5.0),
|
||||
]
|
||||
filter_str, duration = build_concat_filter(chains)
|
||||
|
||||
self.assertIn("[0:v]scale=1280:720[v0]", filter_str)
|
||||
self.assertIn("[v0]concat=n=1:v=1:a=0[outv]", filter_str)
|
||||
# 音频归一化
|
||||
self.assertIn("[0:a]aformat=sample_rates=48000", filter_str)
|
||||
self.assertIn("stereo:sample_fmts=fltp", filter_str)
|
||||
self.assertIn("atrim=0:5.0", filter_str)
|
||||
self.assertIn("[a0]", filter_str)
|
||||
# 音频 concat(n=1)
|
||||
self.assertIn("[a0]concat=n=1:v=0:a=1[outa]", filter_str)
|
||||
self.assertEqual(duration, 5.0)
|
||||
|
||||
def test_two_clips_no_audio(self):
|
||||
"""两片段无音频:两个视频滤镜 + concat(n=2)。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, ["fps=25"], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, ["fps=25"], 3.0),
|
||||
]
|
||||
filter_str, duration = build_concat_filter(chains)
|
||||
|
||||
self.assertIn("[0:v]fps=25[v0]", filter_str)
|
||||
self.assertIn("[1:v]fps=25[v1]", filter_str)
|
||||
self.assertIn("[v0][v1]concat=n=2:v=1:a=0[outv]", filter_str)
|
||||
self.assertEqual(duration, 8.0)
|
||||
|
||||
def test_two_clips_with_audio(self):
|
||||
"""两片段都有音频:视频 concat + 音频归一化 + 音频 concat。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", ["fps=25"], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", "a1", ["fps=25"], 3.0),
|
||||
]
|
||||
filter_str, duration = build_concat_filter(chains)
|
||||
|
||||
# 视频
|
||||
self.assertIn("[v0][v1]concat=n=2:v=1:a=0[outv]", filter_str)
|
||||
# 音频归一化
|
||||
self.assertIn(
|
||||
"[0:a]aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp,atrim=0:5.0,asetpts=PTS-STARTPTS[a0]",
|
||||
filter_str,
|
||||
)
|
||||
self.assertIn(
|
||||
"[1:a]aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp,atrim=0:3.0,asetpts=PTS-STARTPTS[a1]",
|
||||
filter_str,
|
||||
)
|
||||
# 音频 concat
|
||||
self.assertIn("[a0][a1]concat=n=2:v=0:a=1[outa]", filter_str)
|
||||
self.assertEqual(duration, 8.0)
|
||||
|
||||
def test_mixed_audio_some_none(self):
|
||||
"""部分有音频部分没有:只有有音频的片段参与音频 concat。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", ["fps=25"], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, ["fps=25"], 3.0),
|
||||
ClipFilterChain("c3", 2, "v2", "a2", ["fps=25"], 4.0),
|
||||
]
|
||||
filter_str, duration = build_concat_filter(chains)
|
||||
|
||||
# 视频 concat 有 3 个输入
|
||||
self.assertIn("[v0][v1][v2]concat=n=3:v=1:a=0[outv]", filter_str)
|
||||
# 音频 concat 只有 2 个输入
|
||||
self.assertIn("[a0][a2]concat=n=2:v=0:a=1[outa]", filter_str)
|
||||
# 片段 1 没有音频归一化
|
||||
self.assertNotIn("[1:a]", filter_str)
|
||||
self.assertEqual(duration, 12.0)
|
||||
|
||||
def test_three_clips_total_duration(self):
|
||||
"""三片段总时长为各片段之和。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 2.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
ClipFilterChain("c3", 2, "v2", None, [], 4.0),
|
||||
]
|
||||
_, duration = build_concat_filter(chains)
|
||||
self.assertEqual(duration, 9.0)
|
||||
|
||||
def test_audio_format_normalization(self):
|
||||
"""音频归一化包含 aformat/atrim/asetpts。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", [], 5.0),
|
||||
]
|
||||
filter_str, _ = build_concat_filter(chains)
|
||||
|
||||
self.assertIn("aformat=sample_rates=48000", filter_str)
|
||||
self.assertIn("channel_layouts=stereo", filter_str)
|
||||
self.assertIn("sample_fmts=fltp", filter_str)
|
||||
self.assertIn("atrim=0:5.0", filter_str)
|
||||
self.assertIn("asetpts=PTS-STARTPTS", filter_str)
|
||||
|
||||
def test_filter_parts_separated_by_semicolon(self):
|
||||
"""各滤镜部分用分号分隔。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", ["fps=25"], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", "a1", ["fps=25"], 3.0),
|
||||
]
|
||||
filter_str, _ = build_concat_filter(chains)
|
||||
parts = filter_str.split(";")
|
||||
# 2 视频 + 2 音频归一化 + 1 视频 concat + 1 音频 concat = 6
|
||||
self.assertEqual(len(parts), 6)
|
||||
|
||||
|
||||
# ── build_xfade_filter 测试 ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildXfadeFilter(unittest.TestCase):
|
||||
"""build_xfade_filter 转场滤镜测试。"""
|
||||
|
||||
def test_empty_clips(self):
|
||||
"""空列表返回空字符串和 0 时长。"""
|
||||
filter_str, duration = build_xfade_filter([], 0.5, [])
|
||||
self.assertEqual(filter_str, "")
|
||||
self.assertEqual(duration, 0.0)
|
||||
|
||||
def test_single_clip_no_audio(self):
|
||||
"""单片段无音频:视频滤镜 + copy。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, ["scale=1280:720"], 5.0),
|
||||
]
|
||||
filter_str, duration = build_xfade_filter(chains, 0.5, ["fade"])
|
||||
|
||||
self.assertIn("[0:v]scale=1280:720[v0]", filter_str)
|
||||
self.assertIn("[v0]copy[outv]", filter_str)
|
||||
self.assertEqual(duration, 5.0)
|
||||
# 单片段 xfade 没有音频输出
|
||||
self.assertNotIn("[outa]", filter_str)
|
||||
|
||||
def test_single_clip_with_audio(self):
|
||||
"""单片段有音频:xfade 路径下单片段不输出音频(与原实现一致)。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", ["scale=1280:720"], 5.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["fade"])
|
||||
# 单片段 xfade 没有音频输出
|
||||
self.assertNotIn("[outa]", filter_str)
|
||||
self.assertNotIn("acopy", filter_str)
|
||||
|
||||
def test_two_clips_fade_transition(self):
|
||||
"""两片段 fade 转场:xfade 滤镜结构正确。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, ["fps=25"], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, ["fps=25"], 3.0),
|
||||
]
|
||||
filter_str, duration = build_xfade_filter(chains, 0.5, ["cut", "fade"])
|
||||
|
||||
# 两个视频滤镜链
|
||||
self.assertIn("[0:v]fps=25[v0]", filter_str)
|
||||
self.assertIn("[1:v]fps=25[v1]", filter_str)
|
||||
# xfade 转场
|
||||
self.assertIn("xfade=transition=fade", filter_str)
|
||||
self.assertIn(":duration=0.5", filter_str)
|
||||
self.assertIn("[outv]", filter_str)
|
||||
# 总时长 = 5 + 3 - 0.5 = 7.5
|
||||
self.assertAlmostEqual(duration, 7.5)
|
||||
|
||||
def test_two_clips_offset_calculation(self):
|
||||
"""转场 offset = 第一个片段时长 - 转场时长。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["cut", "fade"])
|
||||
|
||||
# offset = 5.0 - 0.5 * 1 = 4.5
|
||||
self.assertIn(":offset=4.500", filter_str)
|
||||
|
||||
def test_three_clips_chain(self):
|
||||
"""三片段链式转场:两个 xfade,中间用 xf1 标签。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 4.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
ClipFilterChain("c3", 2, "v2", None, [], 5.0),
|
||||
]
|
||||
filter_str, duration = build_xfade_filter(chains, 0.5, ["cut", "fade", "dissolve"])
|
||||
|
||||
# 第一个转场输出到 xf1
|
||||
self.assertIn("[xf1]", filter_str)
|
||||
# 第二个转场输出到 outv
|
||||
self.assertIn("[xf1][v2]xfade=transition=dissolve", filter_str)
|
||||
self.assertIn("[outv]", filter_str)
|
||||
# 总时长 = 4 + 3 + 5 - 0.5 * 2 = 11.0
|
||||
self.assertAlmostEqual(duration, 11.0)
|
||||
|
||||
def test_three_clips_offsets(self):
|
||||
"""三片段两个转场的 offset 计算正确。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 4.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
ClipFilterChain("c3", 2, "v2", None, [], 5.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["cut", "fade", "slideleft"])
|
||||
|
||||
# 第一个 offset = 4.0 - 0.5*1 = 3.5
|
||||
# 第二个 offset = (4.0+3.0) - 0.5*2 = 7.0 - 1.0 = 6.0
|
||||
self.assertIn(":offset=3.500", filter_str)
|
||||
self.assertIn(":offset=6.000", filter_str)
|
||||
|
||||
def test_all_transition_types(self):
|
||||
"""所有转场类型都能正确映射。"""
|
||||
transitions = [
|
||||
(TransitionEffect.FADE, "fade"),
|
||||
(TransitionEffect.SLIDE_LEFT, "slideleft"),
|
||||
(TransitionEffect.SLIDE_RIGHT, "slideright"),
|
||||
(TransitionEffect.DISSOLVE, "dissolve"),
|
||||
(TransitionEffect.WIPE, "wipeleft"),
|
||||
]
|
||||
for effect, expected_name in transitions:
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 3.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 2.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["cut", effect])
|
||||
self.assertIn(
|
||||
f"xfade=transition={expected_name}",
|
||||
filter_str,
|
||||
f"Transition {effect} should map to {expected_name}",
|
||||
)
|
||||
|
||||
def test_unknown_transition_defaults_to_fade(self):
|
||||
"""未知转场类型默认使用 fade。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 3.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 2.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["cut", "nonexistent"])
|
||||
self.assertIn("xfade=transition=fade", filter_str)
|
||||
|
||||
def test_cut_still_uses_fade(self):
|
||||
"""cut 类型在 xfade 路径下也映射为 fade(因为走了 xfade 分支)。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 3.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 2.0),
|
||||
]
|
||||
# 只要有一个非 cut 就走 xfade,cut 的那个也用 fade 作为默认
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["fade", "cut"])
|
||||
# 第二个转场是 cut,默认用 fade
|
||||
self.assertIn("xfade=transition=fade", filter_str)
|
||||
|
||||
def test_transitions_shorter_than_clips(self):
|
||||
"""transitions 列表比 clip 短时,超出部分默认 cut→fade。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 2.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 2.0),
|
||||
ClipFilterChain("c3", 2, "v2", None, [], 2.0),
|
||||
]
|
||||
# 只给 1 个 transition(索引0),索引1和2会越界
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["fade"])
|
||||
# 应该有两个 xfade,都用 fade(第二个是默认值)
|
||||
self.assertEqual(filter_str.count("xfade=transition=fade"), 2)
|
||||
|
||||
def test_zero_transition_duration(self):
|
||||
"""转场时长为 0 时不减时长。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
]
|
||||
_, duration = build_xfade_filter(chains, 0.0, ["cut", "fade"])
|
||||
self.assertAlmostEqual(duration, 8.0)
|
||||
|
||||
def test_total_duration_not_negative(self):
|
||||
"""总时长不会为负数。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 0.1),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 0.1),
|
||||
]
|
||||
_, duration = build_xfade_filter(chains, 10.0, ["cut", "fade"])
|
||||
self.assertGreaterEqual(duration, 0.0)
|
||||
|
||||
def test_two_clips_with_audio_normalize_and_concat(self):
|
||||
"""两片段都有音频:归一化 + concat。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", "a1", [], 3.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["cut", "fade"])
|
||||
|
||||
# 音频归一化(注意:xfade 路径用 audio_label 作为输入,与原实现一致)
|
||||
self.assertIn(
|
||||
"[a0]aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp,atrim=0:5.0,asetpts=PTS-STARTPTS[anorm_v0]",
|
||||
filter_str,
|
||||
)
|
||||
self.assertIn(
|
||||
"[a1]aformat=sample_rates=48000:channel_layouts=stereo:sample_fmts=fltp,atrim=0:3.0,asetpts=PTS-STARTPTS[anorm_v1]",
|
||||
filter_str,
|
||||
)
|
||||
# 音频 concat
|
||||
self.assertIn("[anorm_v0][anorm_v1]concat=n=2:v=0:a=1[outa]", filter_str)
|
||||
|
||||
def test_single_audio_in_xfade_acopy(self):
|
||||
"""xfade 路径下只有一个音频片段时直接 acopy。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", "a0", [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["cut", "fade"])
|
||||
|
||||
self.assertIn("[a0]acopy[outa]", filter_str)
|
||||
# 没有音频归一化
|
||||
self.assertNotIn("aformat", filter_str)
|
||||
self.assertNotIn("concat=n=", filter_str)
|
||||
|
||||
def test_no_audio_in_xfade(self):
|
||||
"""xfade 路径下都没有音频时没有 outa。"""
|
||||
chains = [
|
||||
ClipFilterChain("c1", 0, "v0", None, [], 5.0),
|
||||
ClipFilterChain("c2", 1, "v1", None, [], 3.0),
|
||||
]
|
||||
filter_str, _ = build_xfade_filter(chains, 0.5, ["cut", "fade"])
|
||||
self.assertNotIn("[outa]", filter_str)
|
||||
self.assertNotIn("acopy", filter_str)
|
||||
|
||||
|
||||
# ── build_filter_complex 测试 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildFilterComplex(unittest.TestCase):
|
||||
"""build_filter_complex 策略选择测试。"""
|
||||
|
||||
def _chain(self, idx: int, has_audio: bool = True) -> ClipFilterChain:
|
||||
return ClipFilterChain(
|
||||
clip_id=f"c{idx}",
|
||||
input_index=idx,
|
||||
video_label=f"v{idx}",
|
||||
audio_label=f"a{idx}" if has_audio else None,
|
||||
filters=["fps=25"],
|
||||
duration=3.0,
|
||||
)
|
||||
|
||||
def test_empty_clips(self):
|
||||
"""空列表返回空字符串和 0 时长。"""
|
||||
filter_str, duration = build_filter_complex([], 1280, 720, 0.5, [])
|
||||
self.assertEqual(filter_str, "")
|
||||
self.assertEqual(duration, 0.0)
|
||||
|
||||
def test_single_clip_direct_output(self):
|
||||
"""单片段:直接输出单链滤镜。"""
|
||||
chains = [self._chain(0)]
|
||||
filter_str, duration = build_filter_complex(chains, 1280, 720, 0.5, [])
|
||||
|
||||
self.assertIn("[0:v]fps=25[v0]", filter_str)
|
||||
self.assertNotIn("concat", filter_str)
|
||||
self.assertNotIn("xfade", filter_str)
|
||||
self.assertEqual(duration, 3.0)
|
||||
|
||||
def test_single_clip_audio_passthrough(self):
|
||||
"""单片段有音频:音频直通标签。"""
|
||||
chains = [self._chain(0, has_audio=True)]
|
||||
filter_str, _ = build_filter_complex(chains, 1280, 720, 0.5, [])
|
||||
# 单片段音频:[0:a]a0(直通标签)
|
||||
self.assertIn("[0:a]a0", filter_str)
|
||||
|
||||
def test_single_clip_no_audio(self):
|
||||
"""单片段无音频:没有音频部分。"""
|
||||
chains = [self._chain(0, has_audio=False)]
|
||||
filter_str, _ = build_filter_complex(chains, 1280, 720, 0.5, [])
|
||||
self.assertNotIn("[0:a]", filter_str)
|
||||
self.assertNotIn("[outa]", filter_str)
|
||||
|
||||
def test_multiple_all_cut_uses_concat(self):
|
||||
"""多片段 + 全 cut:使用 concat 滤镜。"""
|
||||
chains = [self._chain(0), self._chain(1), self._chain(2)]
|
||||
transitions = [TransitionEffect.CUT, TransitionEffect.CUT, TransitionEffect.CUT]
|
||||
filter_str, duration = build_filter_complex(chains, 1280, 720, 0.5, transitions)
|
||||
|
||||
self.assertIn("concat=n=3:v=1:a=0[outv]", filter_str)
|
||||
self.assertNotIn("xfade", filter_str)
|
||||
self.assertEqual(duration, 9.0)
|
||||
|
||||
def test_multiple_one_transition_uses_xfade(self):
|
||||
"""多片段 + 有一个非 cut 转场:使用 xfade。"""
|
||||
chains = [self._chain(0), self._chain(1)]
|
||||
transitions = [TransitionEffect.CUT, TransitionEffect.FADE]
|
||||
filter_str, duration = build_filter_complex(chains, 1280, 720, 0.5, transitions)
|
||||
|
||||
self.assertIn("xfade=transition=fade", filter_str)
|
||||
self.assertNotIn("concat=n=2:v=1:a=0", filter_str)
|
||||
self.assertAlmostEqual(duration, 5.5) # 3 + 3 - 0.5
|
||||
|
||||
def test_string_cut_value(self):
|
||||
"""字符串 'cut' 也被识别为无转场。"""
|
||||
chains = [self._chain(0), self._chain(1)]
|
||||
transitions = ["cut", "cut"]
|
||||
filter_str, _ = build_filter_complex(chains, 1280, 720, 0.5, transitions)
|
||||
self.assertIn("concat=n=2:v=1:a=0[outv]", filter_str)
|
||||
self.assertNotIn("xfade", filter_str)
|
||||
|
||||
def test_mixed_cut_and_transition(self):
|
||||
"""混合 cut 和转场:走 xfade 路径。"""
|
||||
chains = [self._chain(0), self._chain(1), self._chain(2)]
|
||||
transitions = ["cut", TransitionEffect.FADE, "cut"]
|
||||
filter_str, _ = build_filter_complex(chains, 1280, 720, 0.5, transitions)
|
||||
self.assertIn("xfade", filter_str)
|
||||
|
||||
def test_all_dissolve_transition(self):
|
||||
"""全部 dissolve 转场。"""
|
||||
chains = [self._chain(0), self._chain(1)]
|
||||
transitions = [TransitionEffect.CUT, TransitionEffect.DISSOLVE]
|
||||
filter_str, _ = build_filter_complex(chains, 1280, 720, 0.5, transitions)
|
||||
self.assertIn("xfade=transition=dissolve", filter_str)
|
||||
|
||||
|
||||
# ── 集成测试:build_clip_filter + build_filter_complex 端到端 ────────────────
|
||||
|
||||
|
||||
class TestEndToEndFilterBuilding(unittest.TestCase):
|
||||
"""端到端集成测试:从 EditPlanClip 到完整 filter_complex。"""
|
||||
|
||||
def test_two_video_clips_concat(self):
|
||||
"""两个视频片段走 concat 路径的完整流程。"""
|
||||
clip1 = _make_clip("c1", duration=5.0, clip_type="video")
|
||||
clip2 = _make_clip("c2", duration=3.0, clip_type="video")
|
||||
|
||||
chain1 = build_clip_filter(clip1, 0, 1280, 720, 25)
|
||||
chain2 = build_clip_filter(clip2, 1, 1280, 720, 25)
|
||||
|
||||
filter_str, duration = build_filter_complex(
|
||||
[chain1, chain2],
|
||||
1280,
|
||||
720,
|
||||
0.5,
|
||||
[TransitionEffect.CUT, TransitionEffect.CUT],
|
||||
)
|
||||
|
||||
# 有两个视频滤镜链
|
||||
self.assertIn("[0:v]", filter_str)
|
||||
self.assertIn("[1:v]", filter_str)
|
||||
# concat 输出
|
||||
self.assertIn("[outv]", filter_str)
|
||||
self.assertIn("[outa]", filter_str)
|
||||
# 总时长
|
||||
self.assertAlmostEqual(duration, 8.0)
|
||||
# 结构:2视频 + 2音频 + 1视频concat + 1音频concat = 6 段
|
||||
self.assertEqual(len(filter_str.split(";")), 6)
|
||||
|
||||
def test_two_clips_with_xfade(self):
|
||||
"""两个片段走 xfade 转场的完整流程。"""
|
||||
clip1 = _make_clip("c1", duration=5.0, clip_type="video")
|
||||
clip2 = _make_clip("c2", duration=4.0, clip_type="video")
|
||||
|
||||
chain1 = build_clip_filter(clip1, 0, 1920, 1080, 30)
|
||||
chain2 = build_clip_filter(clip2, 1, 1920, 1080, 30)
|
||||
|
||||
filter_str, duration = build_filter_complex(
|
||||
[chain1, chain2],
|
||||
1920,
|
||||
1080,
|
||||
0.5,
|
||||
[TransitionEffect.CUT, TransitionEffect.FADE],
|
||||
)
|
||||
|
||||
self.assertIn("xfade=transition=fade:duration=0.5", filter_str)
|
||||
self.assertIn("[outv]", filter_str)
|
||||
# 总时长减去转场
|
||||
self.assertAlmostEqual(duration, 8.5) # 5 + 4 - 0.5
|
||||
|
||||
def test_title_plus_video(self):
|
||||
"""title 片段(无音频)+ video 片段(有音频)。"""
|
||||
clip1 = _make_clip("c1", duration=2.0, clip_type="title")
|
||||
clip2 = _make_clip("c2", duration=5.0, clip_type="video")
|
||||
|
||||
chain1 = build_clip_filter(clip1, 0, 1280, 720, 25)
|
||||
chain2 = build_clip_filter(clip2, 1, 1280, 720, 25)
|
||||
|
||||
# title 无音频,video 有音频
|
||||
self.assertIsNone(chain1.audio_label)
|
||||
self.assertIsNotNone(chain2.audio_label)
|
||||
|
||||
# concat 路径
|
||||
filter_str, _ = build_filter_complex(
|
||||
[chain1, chain2],
|
||||
1280,
|
||||
720,
|
||||
0.5,
|
||||
[TransitionEffect.CUT, TransitionEffect.CUT],
|
||||
)
|
||||
# 音频 concat 只有 1 个输入(片段2)
|
||||
self.assertIn("[a1]concat=n=1:v=0:a=1[outa]", filter_str)
|
||||
self.assertNotIn("[a0]", filter_str)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user