Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 18ef9529eb |
@@ -18,6 +18,13 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
)
|
||||
from packages.domain.edit_plan import EditPlan, EditPlanStatus
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
from packages.domain.clip_operations import (
|
||||
calculate_merge as _calc_merge,
|
||||
calculate_shift_orders as _calc_shift_orders,
|
||||
calculate_split as _calc_split,
|
||||
validate_merge_clips as _validate_merge,
|
||||
validate_split_time as _validate_split,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -384,36 +391,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 +440,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 +454,8 @@ class EditPlanService:
|
||||
clip_id,
|
||||
plan_id,
|
||||
split_time,
|
||||
left_duration,
|
||||
right_duration,
|
||||
split.left_duration,
|
||||
split.right_duration,
|
||||
)
|
||||
|
||||
return {
|
||||
@@ -468,70 +484,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
|
||||
|
||||
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
+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 (
|
||||
MergeResult,
|
||||
ROUND_PRECISION,
|
||||
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()
|
||||
Reference in New Issue
Block a user