"""片段操作工具单测. 纯函数模块,覆盖:分割校验/计算、合并校验/计算、 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