diff --git a/apps/api/app/services/edit_plan_service.py b/apps/api/app/services/edit_plan_service.py index 8b8ad6f87..80a4bdb48 100755 --- a/apps/api/app/services/edit_plan_service.py +++ b/apps/api/app/services/edit_plan_service.py @@ -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 diff --git a/packages/domain/clip_operations.py b/packages/domain/clip_operations.py new file mode 100755 index 000000000..ebfdddb66 --- /dev/null +++ b/packages/domain/clip_operations.py @@ -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 diff --git a/tests/unit/test_clip_operations.py b/tests/unit/test_clip_operations.py new file mode 100755 index 000000000..01be94b59 --- /dev/null +++ b/tests/unit/test_clip_operations.py @@ -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()