diff --git a/tests/unit/domain/test_clip_operations.py b/tests/unit/domain/test_clip_operations.py new file mode 100755 index 000000000..14c9d0dd2 --- /dev/null +++ b/tests/unit/domain/test_clip_operations.py @@ -0,0 +1,518 @@ +"""片段操作工具单测. + +纯函数模块,覆盖:分割校验/计算、合并校验/计算、 +order重排、order偏移。 +""" + +from __future__ import annotations + +from dataclasses import dataclass + +from packages.domain.clip_operations import ( + DEFAULT_SPLIT_DURATION, + MergeResult, + SplitResult, + calculate_merge, + calculate_reorder_new_orders, + calculate_shift_orders, + calculate_split, + validate_merge_clips, + validate_split_time, +) + + +class TestConstants: + def test_default_split_duration(self): + assert DEFAULT_SPLIT_DURATION == 5.0 + + +class TestValidateSplitTime: + def test_valid_middle(self): + validate_split_time(5.0, 10.0) # 不抛异常就是通过 + + def test_valid_small(self): + validate_split_time(0.1, 10.0) + + def test_valid_near_end(self): + validate_split_time(9.9, 10.0) + + def test_zero_invalid(self): + try: + validate_split_time(0.0, 10.0) + except ValueError as e: + assert "分割时间" in str(e) + else: + raise AssertionError("expected ValueError") + + def test_negative_invalid(self): + try: + validate_split_time(-1.0, 10.0) + except ValueError: + pass + else: + raise AssertionError("expected ValueError") + + def test_equal_to_duration_invalid(self): + try: + validate_split_time(10.0, 10.0) + except ValueError: + pass + else: + raise AssertionError("expected ValueError") + + def test_greater_than_duration_invalid(self): + try: + validate_split_time(15.0, 10.0) + except ValueError: + pass + else: + raise AssertionError("expected ValueError") + + +class TestCalculateSplit: + def test_split_half(self): + result = calculate_split(10.0, 5.0) + assert isinstance(result, SplitResult) + assert result.left_duration == 5.0 + assert result.right_duration == 5.0 + assert result.right_start_time == 5.0 + assert result.left_trim_end == 5.0 + assert result.right_trim_start == 5.0 + + def test_split_one_third(self): + result = calculate_split(9.0, 3.0) + assert result.left_duration == 3.0 + assert result.right_duration == 6.0 + assert result.right_start_time == 3.0 + + def test_split_with_start_time(self): + result = calculate_split(10.0, 4.0, start_time=100.0) + assert result.left_duration == 4.0 + assert result.right_duration == 6.0 + assert result.right_start_time == 104.0 + + def test_split_precision_rounding(self): + result = calculate_split(1.0, 1 / 3, precision=3) + assert result.left_duration == round(1 / 3, 3) + assert result.right_duration == round(2 / 3, 3) + + def test_split_default_precision_is_3(self): + result = calculate_split(1.0, 0.123456) + # 默认精度3位 + assert result.left_duration == 0.123 + + def test_custom_precision(self): + result = calculate_split(1.0, 0.123456, precision=5) + assert result.left_duration == 0.12346 # 5位精度,四舍五入 + + def test_split_returns_frozen_dataclass(self): + result = calculate_split(10.0, 5.0) + try: + result.left_duration = 3.0 # type: ignore + except AttributeError: + pass # frozen,应该抛异常 + else: + raise AssertionError("SplitResult should be frozen") + + def test_invalid_split_time_raises(self): + try: + calculate_split(10.0, 0.0) + except ValueError: + pass + else: + raise AssertionError("expected ValueError") + + +@dataclass +class FakeClip: + """模拟 EditPlanClip 的最小数据类.""" + + id: str = "" + plan_id: str = "plan_1" + order: int = 0 + duration: float = 3.0 + clip_type: str = "main" + text_content: str = "" + config: dict | None = None + + +class TestValidateMergeClips: + def test_valid_two_clips(self): + clips = [ + FakeClip(id="c1", order=0), + FakeClip(id="c2", order=1), + ] + plan_id, first_order = validate_merge_clips(clips) + assert plan_id == "plan_1" + assert first_order == 0 + + def test_valid_three_clips(self): + clips = [ + FakeClip(id="c1", order=2), + FakeClip(id="c2", order=3), + FakeClip(id="c3", order=4), + ] + plan_id, first_order = validate_merge_clips(clips) + assert plan_id == "plan_1" + assert first_order == 2 + + def test_unordered_input_still_valid(self): + """输入顺序不影响,内部会排序.""" + clips = [ + FakeClip(id="c3", order=2), + FakeClip(id="c1", order=0), + FakeClip(id="c2", order=1), + ] + plan_id, first_order = validate_merge_clips(clips) + assert first_order == 0 + + def test_single_clip_invalid(self): + try: + validate_merge_clips([FakeClip()]) + except ValueError as e: + assert "至少需要 2 个" in str(e) + else: + raise AssertionError("expected ValueError") + + def test_empty_list_invalid(self): + try: + validate_merge_clips([]) + except ValueError: + pass + else: + raise AssertionError("expected ValueError") + + def test_different_plan_invalid(self): + clips = [ + FakeClip(id="c1", plan_id="plan_a", order=0), + FakeClip(id="c2", plan_id="plan_b", order=1), + ] + try: + validate_merge_clips(clips) + except ValueError as e: + assert "同一计划" in str(e) + else: + raise AssertionError("expected ValueError") + + def test_non_consecutive_order_invalid(self): + clips = [ + FakeClip(id="c1", order=0), + FakeClip(id="c2", order=2), # 跳过1 + ] + try: + validate_merge_clips(clips) + except ValueError as e: + assert "不连续" in str(e) + else: + raise AssertionError("expected ValueError") + + def test_different_clip_type_invalid(self): + clips = [ + FakeClip(id="c1", order=0, clip_type="main"), + FakeClip(id="c2", order=1, clip_type="title"), + ] + try: + validate_merge_clips(clips) + except ValueError as e: + assert "相同类型" in str(e) + else: + raise AssertionError("expected ValueError") + + +class TestCalculateMerge: + def test_merge_two_clips_duration(self): + clips = [ + FakeClip(id="c1", order=0, duration=3.0), + FakeClip(id="c2", order=1, duration=5.0), + ] + result = calculate_merge(clips) + assert isinstance(result, MergeResult) + assert result.total_duration == 8.0 + + def test_merge_three_clips_duration(self): + clips = [ + FakeClip(id="c1", order=0, duration=2.0), + FakeClip(id="c2", order=1, duration=3.0), + FakeClip(id="c3", order=2, duration=4.0), + ] + result = calculate_merge(clips) + assert result.total_duration == 9.0 + + def test_merge_text_concatenation(self): + clips = [ + FakeClip(id="c1", order=0, text_content="第一句"), + FakeClip(id="c2", order=1, text_content="第二句"), + ] + result = calculate_merge(clips) + assert result.merged_text == "第一句\n第二句" + + def test_merge_empty_text_skipped(self): + clips = [ + FakeClip(id="c1", order=0, text_content="hello"), + FakeClip(id="c2", order=1, text_content=""), + FakeClip(id="c3", order=2, text_content="world"), + ] + result = calculate_merge(clips) + assert result.merged_text == "hello\nworld" + + def test_merge_whitespace_text_skipped(self): + clips = [ + FakeClip(id="c1", order=0, text_content="a"), + FakeClip(id="c2", order=1, text_content=" "), + FakeClip(id="c3", order=2, text_content="b"), + ] + result = calculate_merge(clips) + assert result.merged_text == "a\nb" + + def test_merge_all_empty_text(self): + clips = [ + FakeClip(id="c1", order=0, text_content=""), + FakeClip(id="c2", order=1, text_content=""), + ] + result = calculate_merge(clips) + assert result.merged_text == "" + + def test_merge_config_later_overrides(self): + clips = [ + FakeClip(id="c1", order=0, config={"font_size": 20, "color": "red"}), + FakeClip(id="c2", order=1, config={"font_size": 24, "bold": True}), + ] + result = calculate_merge(clips) + assert result.merged_config["font_size"] == 24 # 后面的覆盖 + assert result.merged_config["color"] == "red" + assert result.merged_config["bold"] is True + + def test_merge_config_removes_trim_fields(self): + clips = [ + FakeClip(id="c1", order=0, config={"trim_start": 1.0, "a": 1}), + FakeClip(id="c2", order=1, config={"trim_end": 2.0, "b": 2}), + ] + result = calculate_merge(clips) + assert "trim_start" not in result.merged_config + assert "trim_end" not in result.merged_config + assert result.merged_config["a"] == 1 + assert result.merged_config["b"] == 2 + + def test_merge_none_config_handled(self): + clips = [ + FakeClip(id="c1", order=0, config=None), + FakeClip(id="c2", order=1, config={"key": "val"}), + ] + result = calculate_merge(clips) + assert result.merged_config == {"key": "val"} + + def test_merge_first_order_and_shift(self): + clips = [ + FakeClip(id="c1", order=5), + FakeClip(id="c2", order=6), + FakeClip(id="c3", order=7), + ] + result = calculate_merge(clips) + assert result.first_order == 5 + assert result.shift_amount == 2 # 3个合并成1个,前移2位 + + def test_merge_two_clips_shift(self): + clips = [FakeClip(id="c1", order=0), FakeClip(id="c2", order=1)] + result = calculate_merge(clips) + assert result.shift_amount == 1 + + def test_merge_unordered_input(self): + """输入乱序也能正确处理(内部排序).""" + clips = [ + FakeClip(id="c3", order=2, duration=4.0, text_content="C"), + FakeClip(id="c1", order=0, duration=2.0, text_content="A"), + FakeClip(id="c2", order=1, duration=3.0, text_content="B"), + ] + result = calculate_merge(clips) + assert result.total_duration == 9.0 + assert result.merged_text == "A\nB\nC" + assert result.first_order == 0 + + def test_merge_precision(self): + clips = [ + FakeClip(id="c1", order=0, duration=1 / 3), + FakeClip(id="c2", order=1, duration=1 / 3), + ] + result = calculate_merge(clips, precision=3) + assert result.total_duration == round(2 / 3, 3) + + def test_merge_empty_list_raises(self): + try: + calculate_merge([]) + except ValueError as e: + assert "不能为空" in str(e) + else: + raise AssertionError("expected ValueError") + + def test_merge_result_is_frozen(self): + clips = [FakeClip(id="c1", order=0), FakeClip(id="c2", order=1)] + result = calculate_merge(clips) + try: + result.total_duration = 10.0 # type: ignore + except AttributeError: + pass + else: + raise AssertionError("MergeResult should be frozen") + + +@dataclass +class FakeItem: + id: str + order: int = 0 + + +class TestCalculateReorderNewOrders: + def test_basic_reorder(self): + items = [ + FakeItem(id="a", order=0), + FakeItem(id="b", order=1), + FakeItem(id="c", order=2), + ] + new_order = ["c", "a", "b"] + result = calculate_reorder_new_orders(new_order, items) + assert result == {"c": 0, "a": 1, "b": 2} + + def test_reverse_order(self): + items = [FakeItem(id="a"), FakeItem(id="b"), FakeItem(id="c")] + new_order = ["c", "b", "a"] + result = calculate_reorder_new_orders(new_order, items) + assert result["c"] == 0 + assert result["b"] == 1 + assert result["a"] == 2 + + def test_same_order(self): + items = [FakeItem(id="a"), FakeItem(id="b")] + new_order = ["a", "b"] + result = calculate_reorder_new_orders(new_order, items) + assert result == {"a": 0, "b": 1} + + def test_mismatched_ids_raises(self): + items = [FakeItem(id="a"), FakeItem(id="b")] + try: + calculate_reorder_new_orders(["a", "c"], items) + except ValueError as e: + assert "不匹配" in str(e) + else: + raise AssertionError("expected ValueError") + + def test_extra_id_in_list_raises(self): + items = [FakeItem(id="a")] + try: + calculate_reorder_new_orders(["a", "b"], items) + except ValueError: + pass + else: + raise AssertionError("expected ValueError") + + def test_missing_id_raises(self): + items = [FakeItem(id="a"), FakeItem(id="b")] + try: + calculate_reorder_new_orders(["a"], items) + except ValueError: + pass + else: + raise AssertionError("expected ValueError") + + def test_custom_id_attr(self): + @dataclass + class CustomItem: + key: str + order: int = 0 + + items = [CustomItem(key="x"), CustomItem(key="y")] + result = calculate_reorder_new_orders(["y", "x"], items, id_attr="key") + assert result == {"y": 0, "x": 1} + + def test_custom_order_attr_does_not_affect_return(self): + """order_attr不影响返回值(返回的是索引),只影响参数校验的ID提取.""" + items = [FakeItem(id="a", order=10), FakeItem(id="b", order=20)] + result = calculate_reorder_new_orders(["b", "a"], items) + assert result == {"b": 0, "a": 1} # 新order是索引,不是原值 + + +class TestCalculateShiftOrders: + def test_shift_positive(self): + items = [ + FakeItem(id="a", order=0), + FakeItem(id="b", order=1), + FakeItem(id="c", order=2), + ] + result = calculate_shift_orders(items, threshold_order=0, shift=5) + # order > 0 的是 b(1) 和 c(2) + shifted = {item.id: new_order for item, new_order in result} + assert len(result) == 2 + assert shifted["b"] == 6 + assert shifted["c"] == 7 + + def test_shift_negative(self): + items = [ + FakeItem(id="a", order=0), + FakeItem(id="b", order=1), + FakeItem(id="c", order=2), + ] + result = calculate_shift_orders(items, threshold_order=0, shift=-1) + shifted = {item.id: new_order for item, new_order in result} + assert shifted["b"] == 0 + assert shifted["c"] == 1 + + def test_threshold_not_included(self): + """threshold_order本身不包含在内(严格大于).""" + items = [FakeItem(id="a", order=5)] + result = calculate_shift_orders(items, threshold_order=5, shift=1) + assert len(result) == 0 + + def test_excluded_ids_skipped(self): + items = [ + FakeItem(id="a", order=1), + FakeItem(id="b", order=2), + FakeItem(id="c", order=3), + ] + result = calculate_shift_orders(items, threshold_order=0, shift=10, excluded_ids={"b"}) + shifted = {item.id: new_order for item, new_order in result} + assert "b" not in shifted + assert shifted["a"] == 11 + assert shifted["c"] == 13 + + def test_none_excluded_ids(self): + items = [FakeItem(id="a", order=1)] + result = calculate_shift_orders(items, threshold_order=0, shift=1, excluded_ids=None) + assert len(result) == 1 + + def test_empty_excluded_ids(self): + items = [FakeItem(id="a", order=1)] + result = calculate_shift_orders(items, threshold_order=0, shift=1, excluded_ids=set()) + assert len(result) == 1 + + def test_no_items_above_threshold(self): + items = [ + FakeItem(id="a", order=0), + FakeItem(id="b", order=1), + ] + result = calculate_shift_orders(items, threshold_order=10, shift=5) + assert len(result) == 0 + + def test_custom_id_attr(self): + @dataclass + class CustomItem: + key: str + pos: int = 0 + + items = [CustomItem(key="x", pos=1), CustomItem(key="y", pos=2)] + result = calculate_shift_orders( + items, + threshold_order=0, + shift=3, + id_attr="key", + order_attr="pos", + ) + assert len(result) == 2 + assert result[0][1] == 4 + assert result[1][1] == 5 + + def test_preserves_item_reference(self): + item = FakeItem(id="a", order=5) + items = [item] + result = calculate_shift_orders(items, threshold_order=3, shift=2) + assert len(result) == 1 + assert result[0][0] is item # 是同一个对象引用 + assert result[0][1] == 7